Qarakal-Rabi-Feedback / ai_model.py
maskil's picture
Upload 4 files
1f4ffe1 verified
Raw
History Blame
6.78 kB
"""
AI Model for Rabi Oscillation Parameter Estimation.
Transformer-based model that predicts fit parameters (A, T, phi, C) from raw signal data.
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
import math
DEVICE = torch.device("mps") if torch.backends.mps.is_available() else torch.device("cpu")
FIXED_LEN = 256
A_MIN, A_MAX = 1e-5, 0.5
T_MIN, T_MAX = 0.02, 0.3
C_MIN, C_MAX = -0.1, 1.0
LOG_A_MIN, LOG_A_MAX = math.log(A_MIN), math.log(A_MAX)
LOG_T_MIN, LOG_T_MAX = math.log(T_MIN), math.log(T_MAX)
AUX_SCALES = np.array([0.1, 0.1, 0.1, 0.05, 0.1, 2.0], dtype=np.float32)
def preprocess_sample(x_raw, y_raw):
xu = np.linspace(x_raw[0], x_raw[-1], FIXED_LEN)
yr = np.interp(xu, x_raw, y_raw)
ymin, ymax = float(yr.min()), float(yr.max())
ymean, ystd = float(yr.mean()), float(yr.std())
xrange = float(x_raw[-1] - x_raw[0])
ydm = yr - ymean
zc = np.where(np.diff(np.sign(ydm)))[0]
nper = len(zc) / 2.0 if len(zc) >= 1 else 0.5
span = ymax - ymin
if span < 1e-15:
span = 1.0
yn = ((yr - ymin) / span).astype(np.float32)
fft = np.fft.rfft(yn)
fft_mag = np.abs(fft[:64])
fft_ang = np.angle(fft[:64])
if len(fft_mag) < 64:
fft_mag = np.pad(fft_mag, (0, 64 - len(fft_mag)))
fft_ang = np.pad(fft_ang, (0, 64 - len(fft_ang)))
fft_feats = np.concatenate([fft_mag, fft_ang]).astype(np.float32)
aux = np.array([ymin, ymax, ymean, ystd, xrange, nper], dtype=np.float32) / AUX_SCALES
return yn, fft_feats, aux
def target_to_params(t):
nA, nT, sp, cp, nC = t[..., 0], t[..., 1], t[..., 2], t[..., 3], t[..., 4]
A = np.exp(nA * (LOG_A_MAX - LOG_A_MIN) + LOG_A_MIN)
T = np.exp(nT * (LOG_T_MAX - LOG_T_MIN) + LOG_T_MIN)
phi = np.arctan2(sp, cp)
C = nC * (C_MAX - C_MIN) + C_MIN
if t.ndim == 1:
return np.array([A, T, phi, C], dtype=np.float64)
return np.stack([A, T, phi, C], -1).astype(np.float64)
class PositionalEncoding(nn.Module):
def __init__(self, d_model, max_len=512):
super().__init__()
pe = torch.zeros(max_len, d_model)
position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
self.register_buffer('pe', pe)
def forward(self, x):
return x + self.pe[:x.size(1), :]
class RabiEstimator(nn.Module):
def __init__(self):
super().__init__()
self.conv1 = nn.Sequential(nn.Conv1d(1, 64, kernel_size=7, stride=2, padding=3), nn.BatchNorm1d(64), nn.GELU())
self.conv2 = nn.Sequential(nn.Conv1d(64, 128, kernel_size=5, stride=2, padding=2), nn.BatchNorm1d(128), nn.GELU())
self.conv3 = nn.Sequential(nn.Conv1d(128, 256, kernel_size=5, stride=2, padding=2), nn.BatchNorm1d(256), nn.GELU())
self.d_model = 256
self.pos_enc = PositionalEncoding(self.d_model)
encoder_layer = nn.TransformerEncoderLayer(d_model=self.d_model, nhead=8, dim_feedforward=512, dropout=0.1, batch_first=True)
self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=6)
self.fft_net = nn.Sequential(nn.Linear(128, 256), nn.GELU(), nn.Dropout(0.1), nn.Linear(256, 128), nn.GELU())
self.aux_net = nn.Sequential(nn.Linear(6, 64), nn.GELU(), nn.Linear(64, 128), nn.GELU())
self.fusion = nn.Sequential(nn.Linear(256 + 128 + 128, 512), nn.BatchNorm1d(512), nn.GELU(), nn.Dropout(0.1))
self.head_atc = nn.Sequential(nn.Linear(512, 256), nn.GELU(), nn.Linear(256, 3))
self.head_ph = nn.Sequential(nn.Linear(512, 256), nn.GELU(), nn.Linear(256, 2))
def forward(self, wave, fft_f, aux):
x = wave.unsqueeze(1)
x = self.conv3(self.conv2(self.conv1(x)))
x = x.permute(0, 2, 1)
x = self.pos_enc(x)
x = self.transformer(x)
x_pool = x.mean(dim=1)
f = self.fft_net(fft_f)
a = self.aux_net(aux)
feat = self.fusion(torch.cat([x_pool, f, a], dim=1))
atc = torch.sigmoid(self.head_atc(feat))
ph = torch.tanh(self.head_ph(feat))
return torch.cat([atc[:, :2], ph, atc[:, 2:3]], 1)
# ═══════════════════════════════════════════════════════════════════
# Helper functions
# ═══════════════════════════════════════════════════════════════════
def rabi_formula(x, A, T, phi, C):
"""Compute Rabi oscillation: A * cos(2*pi/T * x + phi) + C"""
return A * np.cos((2 * np.pi / T) * x + phi) + C
def load_model(model_path):
"""Load trained RabiEstimator weights from disk."""
model = RabiEstimator().to(DEVICE)
checkpoint = torch.load(model_path, map_location=DEVICE, weights_only=False)
if isinstance(checkpoint, dict) and 'model_state_dict' in checkpoint:
model.load_state_dict(checkpoint['model_state_dict'])
else:
model.load_state_dict(checkpoint)
model.eval()
return model
def predict_ai_fit(x_raw, y_raw, model):
"""
Predict Rabi parameters using the AI model and return the absolute fit curve.
Args:
x_raw: np.array of x values (time/amplitude axis)
y_raw: np.array of y values (signal)
model: loaded RabiEstimator model
Returns:
y_fit: np.array of predicted fit values evaluated on x_raw
params: dict with keys 'amplitude', 'T', 'phase', 'offset'
"""
x_raw = np.asarray(x_raw, dtype=np.float64)
y_raw = np.asarray(y_raw, dtype=np.float64)
# Ensure eval mode (BatchNorm1d requires it for batch_size=1)
model.eval()
# Sort by x for interpolation
sort_idx = np.argsort(x_raw)
x_sorted = x_raw[sort_idx]
y_sorted = y_raw[sort_idx]
yn, fft_feats, aux = preprocess_sample(x_sorted, y_sorted)
wave_t = torch.tensor(yn, dtype=torch.float32).unsqueeze(0).to(DEVICE)
fft_t = torch.tensor(fft_feats, dtype=torch.float32).unsqueeze(0).to(DEVICE)
aux_t = torch.tensor(aux, dtype=torch.float32).unsqueeze(0).to(DEVICE)
with torch.no_grad():
pred = model(wave_t, fft_t, aux_t)
pred_np = pred.cpu().numpy()[0]
A, T, phi, C = target_to_params(pred_np)
# Evaluate fit on the original x values (unsorted order preserved)
y_fit = rabi_formula(x_raw, float(A), float(T), float(phi), float(C))
return y_fit, {
'amplitude': float(A),
'T': float(T),
'phase': float(phi),
'offset': float(C),
}