Skip to content

Latest commit

 

History

199 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

QjaxML - Quantum Machine Learning using JAX

A high-performance quantum machine learning library built on JAX, Flax, and PennyLane for training classical, quantum, and hybrid models with GPU acceleration.

Features:

  • JAX-Accelerated: Leverages JAX for fast, GPU-accelerated training with JIT compilation
  • Hybrid Models: Seamlessly combine classical neural networks with quantum circuits
  • High Performance: Mixed-precision training, distributed computing, and optimized gradients
  • Analysis Tools: Fisher Information Matrix & Fourier analysis computation for model expressivity analysis

Performance Benefits:

  • JIT Compilation: Automatic optimization of training loops
  • Vectorized Operations: Efficient batch processing with jax.vmap
  • GPU Acceleration: Native CUDA support through JAX
  • Memory Efficiency: Gradient checkpointing and mixed precision

Quick Start

Quantum Model

import optax
import jax.numpy as jnp

from qjaxml.core.modules import QuantumModule
from qjaxml.core.state import QuantumModuleState
from qjaxml.core.optimizer import ParametersOptimizer
from qjaxml.qcirc.ansatz import EfficientSU2
from qjaxml.qcirc.data_encoding import AngleEmbedding

# Define quantum circuit
ansatz = EfficientSU2(n_qubits=4, reps=2)
feature_map = AngleEmbeddins(n_qubits=4, rotation='Y')

# Create quantum module
q_module = QuantumModule(
    num_qubits=4,
    ansatz=ansatz,
    feature_map=feature_map,
    num_layers=2
)

# Initialize model state
q_state = QuantumModuleState.create(
    module=q_module,
    optimizer=optax.adam(0.01)
)

# Create optimizer and train
def mse_loss(predictions, targets):
    return jnp.mean((predictions - targets) ** 2)

optimizer = ParametersOptimizer(
    module=q_module,
    module_state=q_state,
    loss_fn=mse_loss,
    jitted=True
)

# Train the model
trained_state = optimizer.fit(
    train_dataloader, val_dataloader, epochs=100)

Hybrid Model

import flax.linen as nn
from qjaxml.core.modules import ClassicalModule, HybridModule

# Define classical component
class MLP(nn.Module):
    batch_norm: bool = False
    
    @nn.compact
    def __call__(self, x, train=False):
        x = nn.Dense(64)(x)
        x = nn.relu(x)
        x = nn.Dense(32)(x)
        return x

classical_module = ClassicalModule(
    input_shape=(1, 10),
    flax_module=MLP(batch_norm=True)
)

# Create hybrid model: Classical → Quantum
hybrid_module = HybridModule(
    components=[classical_module, q_module]
)

Advanced Features

Mixed Precision Training

import jmp

# Enable automatic mixed precision
amp_policy = jmp.Policy(
    compute_dtype=jnp.float16,
    param_dtype=jnp.float32,
    output_dtype=jnp.float32
)

Fisher Information Analysis

from qjaxml.analysis.fisher import FisherInformationMatrix

# Analyze model expressivity
fim = FisherInformationMatrix(
    module=q_module,
    data=sample_data
)

# Visualize Fisher Information Matrix
fim.plot_matrix()
fim.plot_spectrum()

Distributed Training

from jax.sharding import NamedSharding

# Multi-GPU training
sharding = NamedSharding(mesh, spec)
optimizer = ParametersOptimizer(
    module=hybrid_module,
    module_state=hybrid_state,
    loss_fn=loss_fn,
    jitted=True,
    sharding=sharding
)

About

Quantum Machine Learning using JAX

Resources

Stars

0 stars

Watchers

1 watching

Forks

Releases

Packages

Used by

Contributors

Languages