Spaces:
Running
Running
| """ | |
| 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 | |