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
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)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]
)import jmp
# Enable automatic mixed precision
amp_policy = jmp.Policy(
compute_dtype=jnp.float16,
param_dtype=jnp.float32,
output_dtype=jnp.float32
)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()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
)