Qiskit - Quantum Neural Networks训练:如何通过.fit方法绘制损失随epochs变化曲线及指定训练epochs数
Great question! Let's break this down into two parts: specifying the number of training epochs, and plotting the loss curve as training progresses using Qiskit's Quantum Neural Network (QNN) .fit() method. Here's a step-by-step guide with code examples:
1. Specifying the Number of Training Epochs
When using Qiskit's scikit-learn compatible models (like NeuralNetworkClassifier or NeuralNetworkRegressor), the .fit() method has a built-in epochs parameter that lets you directly set how many training iterations (epochs) you want to run.
An epoch represents one full pass over your training dataset. For optimizers like SPSA (common in quantum ML), each epoch may correspond to multiple gradient steps, but the epochs parameter simplifies controlling the total training cycles.
2. Collecting Loss Data & Plotting the Curve
To track loss over epochs, you have two straightforward options: using the returned training history object, or a custom callback function. Let's cover both:
Option 1: Use the Returned History Object
When you call .fit(), it returns a History object that stores the loss values for each epoch. You can extract these values and plot them with matplotlib.
Here's a complete code example:
# Import required libraries from qiskit import Aer from qiskit.utils import QuantumInstance, algorithm_globals from qiskit.circuit.library import TwoLocal from qiskit_machine_learning.neural_networks import CircuitQNN from qiskit_machine_learning.algorithms.classifiers import NeuralNetworkClassifier from qiskit.algorithms.optimizers import SPSA import matplotlib.pyplot as plt import numpy as np # Set random seed and quantum instance algorithm_globals.random_seed = 42 quantum_instance = QuantumInstance(Aer.get_backend('aer_simulator'), shots=1024) # Build a simple CircuitQNN feature_map = TwoLocal(2, 'ry', 'cz', reps=1, entanglement='linear') ansatz = TwoLocal(2, 'ry', 'cz', reps=1, entanglement='linear') qnn = CircuitQNN( circuit=feature_map.compose(ansatz).decompose(), input_params=feature_map.parameters, weight_params=ansatz.parameters, quantum_instance=quantum_instance, interpret=lambda x: x % 2, output_shape=2 ) # Initialize classifier with SPSA optimizer optimizer = SPSA(maxiter=100) classifier = NeuralNetworkClassifier(qnn, optimizer=optimizer) # Generate sample training data X = algorithm_globals.random.random((20, 2)) y = np.array([0 if x[0] + x[1] < 1 else 1 for x in X]) # Train with 50 epochs and save the history training_history = classifier.fit(X, y, epochs=50) # Extract loss values from history loss_values = training_history.losses # Plot the loss curve plt.figure(figsize=(10, 6)) plt.plot(range(1, len(loss_values)+1), loss_values, marker='o', color='#1f77b4') plt.xlabel('Epoch Number') plt.ylabel('Training Loss') plt.title('Loss Reduction Over Training Epochs') plt.grid(alpha=0.3) plt.show()
Option 2: Use a Custom Callback Function
If you want more control over tracking (e.g., logging additional metrics per step), you can define a custom callback class that captures loss values during training:
# Define a callback to track losses class LossTracker: def __init__(self): self.losses = [] def __call__(self, weights, current_loss): self.losses.append(current_loss) # Initialize the callback loss_tracker = LossTracker() # Pass the callback to the classifier classifier = NeuralNetworkClassifier(qnn, optimizer=optimizer, callback=loss_tracker) # Train with 50 epochs classifier.fit(X, y, epochs=50) # Plot using the tracked losses plt.figure(figsize=(10, 6)) plt.plot(range(1, len(loss_tracker.losses)+1), loss_tracker.losses, marker='s', color='#ff7f0e') plt.xlabel('Epoch/Step') plt.ylabel('Training Loss') plt.title('Loss Reduction (Custom Callback)') plt.grid(alpha=0.3) plt.show()
Key Notes
- The
epochsparameter works for both classification and regression models in Qiskit ML. - For optimizers like SPSA, each "epoch" may map to multiple optimization steps, but the
Historyobject aggregates loss per epoch automatically. - If you're using a PyTorch/TensorFlow wrapped QNN (via
TorchConnectororTFConnector), you'd use the framework's native training loops, but the.fit()method in Qiskit's scikit-learn interface simplifies this for most use cases.
内容的提问来源于stack exchange,提问作者Zohim Chandani

