""" 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