Qarakal-Rabi-Feedback / preprocessor.py
maskil's picture
Upload 2 files
eb52e7f verified
Raw
History Blame
4.62 kB
"""
Preprocessor for Rabi oscillation JSON experiment files.
V3: Outputs 3 channels — signal + fit + residual (signal − fit).
"""
import json
import numpy as np
import torch
from config import SEQ_LEN
def load_json(path: str) -> dict:
"""Load a JSON experiment file."""
with open(path, 'r') as f:
return json.load(f)
def extract_fit_params(data: dict) -> dict:
"""Extract the 4 fit parameters from JSON data.
Returns dict with keys: amplitude, T, phase, offset.
Returns zeros if fitted_data is missing or None.
"""
defaults = {'amplitude': 0.0, 'T': 0.05, 'phase': 0.0, 'offset': 0.0}
fitted = data.get('fitted_data')
if fitted is None:
return defaults
params_list = fitted.get('parameters')
if params_list is None:
return defaults
params = dict(defaults)
for p in params_list:
if p.get('name') in params:
params[p['name']] = float(p.get('value', 0.0))
return params
def rabi_oscillation(x, amplitude, T, phase, offset):
"""Compute Rabi oscillation curve (without decay, matching JSON fit)."""
return amplitude * np.cos((2 * np.pi / T) * x + phase) + offset
def reconstruct_fit_curve(x: np.ndarray, params: dict) -> np.ndarray:
"""Reconstruct the fit curve on the original X-axis using JSON parameters."""
return rabi_oscillation(
x,
amplitude=params['amplitude'],
T=params['T'],
phase=params['phase'],
offset=params['offset']
)
def preprocess_sample(data: dict):
"""
Preprocess a single JSON experiment sample.
V3: Returns 3-channel tensor (signal, fit, residual).
Returns:
signal_tensor: (3, SEQ_LEN) — [signal, fit, residual]
params_tensor: (4,) fit parameters [amplitude, T, phase, offset]
question_id: str ('q1' or 'q2')
score: float (0.0 or 1.0)
"""
# Extract raw data
x_raw = np.array(data['measured_data']['x_values'], dtype=np.float64)
y_raw = np.array(data['measured_data']['y_values'], dtype=np.float64)
# Extract fit parameters
params = extract_fit_params(data)
# Reconstruct fit curve on original x-axis
y_fit = reconstruct_fit_curve(x_raw, params)
# Create uniform interpolation grid
x_uniform = np.linspace(x_raw.min(), x_raw.max(), SEQ_LEN)
# Interpolate both curves to fixed length
y_raw_interp = np.interp(x_uniform, x_raw, y_raw)
y_fit_interp = np.interp(x_uniform, x_raw, y_fit)
# Normalize both with same min-max scaling (from raw signal)
y_min = y_raw_interp.min()
y_max = y_raw_interp.max()
y_range = y_max - y_min
if y_range < 1e-10:
y_range = 1.0 # Avoid division by zero for flat signals
y_raw_norm = (y_raw_interp - y_min) / y_range
y_fit_norm = (y_fit_interp - y_min) / y_range
# Residual channel: signal − fit, normalized to [-1, 1]
residual = y_raw_norm - y_fit_norm
res_absmax = np.abs(residual).max()
if res_absmax < 1e-10:
res_absmax = 1.0
res_norm = residual / res_absmax
# Build 3-channel tensor
signal_tensor = torch.tensor(
np.stack([y_raw_norm, y_fit_norm, res_norm], axis=0),
dtype=torch.float32
)
params_tensor = torch.tensor(
[params['amplitude'], params['T'], params['phase'], params['offset']],
dtype=torch.float32
)
# Extract labels
question_id = data.get('question_id', 'q1')
score = float(data.get('score', 0.0))
return signal_tensor, params_tensor, question_id, score
def load_test_dataset(data_dir: str) -> list:
"""
Load and preprocess all JSON files from the test data directory.
Returns list of dicts with keys:
signal, params, question_id, score, filename
"""
import os
samples = []
json_files = sorted([f for f in os.listdir(data_dir) if f.endswith('.json')])
for fname in json_files:
path = os.path.join(data_dir, fname)
try:
data = load_json(path)
signal_3ch, params, qid, score = preprocess_sample(data)
samples.append({
'signal': signal_3ch, # (3, SEQ_LEN)
'params': params,
'question_id': qid,
'score': score,
'filename': fname,
'raw_data': data,
})
except Exception as e:
print(f"Warning: Failed to preprocess {fname}: {e}")
continue
print(f"Loaded {len(samples)} samples from {data_dir}")
return samples