""" Central configuration for Rabi Oscillation Quality Classifier. V2: Calibrated to real data distributions. """ import torch import os # ─── Paths ─────────────────────────────────────────────────────────────────── PROJECT_DIR = os.path.dirname(os.path.abspath(__file__)) DATA_DIR = os.path.join(PROJECT_DIR, "data copy") UI_TEST_DIR = os.path.join(PROJECT_DIR, "ui_test_samples") MODEL_PATH = os.path.join(PROJECT_DIR, "best_model.pt") CORRECTIONS_DIR = os.path.join(PROJECT_DIR, "data_corrections") FAILED_PRED_DIR = os.path.join(PROJECT_DIR, "failed_predictions") # ─── Device (MPS > CUDA > CPU) ────────────────────────────────────────────── if torch.backends.mps.is_available(): DEVICE = torch.device("mps") elif torch.cuda.is_available(): DEVICE = torch.device("cuda") else: DEVICE = torch.device("cpu") # ─── Signal Preprocessing ─────────────────────────────────────────────────── SEQ_LEN = 256 # Increased from 128 for finer resolution NUM_FIT_PARAMS = 4 # amplitude, T, phase, offset # ─── Model ─────────────────────────────────────────────────────────────────── NUM_CLASSES = 3 # Classes per head: 0=Bad, 1=Acceptable, 2=Perfect IN_CHANNELS = 3 # raw signal + fit curve + residual (signal - fit) RESNET_DIM = 256 # ResNet backbone output dim DENSE_DIM = 64 # Dense branch output dim FUSED_DIM = RESNET_DIM + DENSE_DIM # 320 # ─── Training ──────────────────────────────────────────────────────────────── BATCH_SIZE = 256 LR = 3e-4 # Lower LR for better convergence WEIGHT_DECAY = 1e-4 EPOCHS = 80 # More epochs TRAIN_SIZE = 200_000 # Doubled for better generalization VAL_SIZE = 40_000 PATIENCE = 15 # More patience FOCAL_GAMMA = 2.0 CLASS_WEIGHTS = [1.2, 0.9, 0.9] # Stronger upweight for class 0 IGNORE_INDEX = -1 # ─── Synthetic Generator — Calibrated to Real Data ────────────────────────── # Real data statistics (from analysis of 2736 JSON files): # amplitude: mean=0.090, std=0.083, range=[~0, 0.40], p25=0.0002 # T: mean=0.101 but p75=0.062, most values 0.047-0.110 # phase: range=[-π, π] # offset: mean=0.185, std=0.166, range=[~0, 0.98] # x_range: mean=0.106, range=[0.006, 3.58] # n_points: mean=145, range=[80, 1000] # Data Quality SNR thresholds (dB) SNR_PERFECT_MIN = 15.0 # Loosened from 20 (real data is noisier) SNR_ACCEPTABLE_MIN = 5.0 # Loosened from 8 SNR_BORDERLINE_MIN = 3.0 # Fit Quality error thresholds (fraction) FIT_PERFECT_MAX_ERR = 0.03 FIT_ACCEPTABLE_MAX_ERR = 0.20 FIT_BAD_MIN_ERR = 0.20 FIT_BORDERLINE_ERR_RANGE = (0.15, 0.30) # Rabi oscillation parameter ranges — CALIBRATED TO REAL DATA AMPLITUDE_RANGE = (0.0001, 0.42) # Real: [~0, 0.40] PERIOD_RANGE = (0.025, 0.80) # Real: [0.025, 0.78] for good fits PHASE_RANGE = (-3.1416, 3.1416) # Real: [-π, π] OFFSET_RANGE = (0.0001, 0.98) # Real: [~0, 0.98] DECAY_RANGE = (0.0, 30.0) # Reduced max decay X_MAX_RANGE = (0.005, 0.50) # Real: [0.006, 3.58] but most < 0.2 # Extended ranges for diversity X_MAX_EXTENDED = (0.50, 3.60) # For occasional long sweeps PERIOD_BAD_RANGE = (0.002, 19.0) # Real bad fits have T up to 19 # Per-sample x-axis variation N_POINTS_RANGE = (80, 500) # Real: [80, 1000] # Proportion of borderline cases in Class 0 BORDERLINE_FRACTION = 0.5 # Increased from 0.4 # Proportion of near-zero amplitude samples (mimics ~25% of real data) NEAR_ZERO_AMP_FRACTION = 0.20