Team Ai
Apppublic

ujjwalpardeshi/pytorch-training-debugger

sourceHugging Faceupdated 6mo agoView on Hugging Face
2likes
simulation.py126 linesDownload Raw Back to ml_training_debugger
1"""Training curve generation — real PyTorch mini-training.2 3All curves come from run_real_training() in pytorch_engine.py:4  - Real torch.nn.Module (SimpleCNN or SimpleMLP)5  - Real torch.autograd forward + backward passes6  - Real torch.optim optimizer steps7  - Real validation on held-out data8  - 20 epochs, cached per (task_id, seed, model_type)9 10Zero numpy. Zero parametric formulas. Zero synthetic curves.11"""12 13from __future__ import annotations14 15import torch16 17from ml_training_debugger.scenarios import ScenarioParams18 19EPOCHS = 2020 21 22def _get_real_curves(scenario: ScenarioParams) -> dict[str, list[float]]:23    """Run real PyTorch training and return loss/accuracy curves.24 25    Calls pytorch_engine.run_real_training() which:26    - Creates a real SimpleCNN or SimpleMLP model27    - Generates random CIFAR-10 style data (3x32x32)28    - Runs 20 epochs of real forward/backward passes29    - Injects the actual fault (wrong LR, eval mode, data leakage, etc.)30    - Returns real loss_history, val_loss_history, val_acc_history31 32    Results are cached per (task_id, seed, model_type) for instant resets.33    """34    from ml_training_debugger.pytorch_engine import run_real_training35 36    return run_real_training(scenario)37 38 39def gen_loss_history(scenario: ScenarioParams) -> list[float]:40    """Generate training loss history (20 epochs) from real PyTorch training."""41    return _get_real_curves(scenario)["loss_history"]42 43 44def gen_val_accuracy_history(scenario: ScenarioParams) -> list[float]:45    """Generate validation accuracy history (20 epochs) from real PyTorch training."""46    return _get_real_curves(scenario)["val_acc_history"]47 48 49def gen_val_loss_history(scenario: ScenarioParams) -> list[float]:50    """Generate validation loss history (20 epochs) from real PyTorch training."""51    return _get_real_curves(scenario)["val_loss_history"]52 53 54def _gen_confusion_matrix(scenario: ScenarioParams) -> list[list[float]]:55    """Generate a 10x10 confusion matrix based on the fault type.56 57    Uses torch.Tensor operations on random data shaped by the fault scenario.58    """59    torch.manual_seed(scenario.seed + 10)60    root = scenario.root_cause.value61    n = 1062 63    if root == "data_leakage":64        # High diagonal but with leakage-induced off-diagonal noise65        base = torch.eye(n) * 0.866        noise = torch.rand(n, n) * scenario.leakage_pct * 0.367        cm = base + noise68    elif root == "overfitting":69        # Near-perfect diagonal (memorized)70        cm = torch.eye(n) * 0.95 + torch.rand(n, n) * 0.0271    else:72        # Normal confusion with moderate accuracy73        cm = torch.eye(n) * 0.6 + torch.rand(n, n) * 0.0874 75    # Normalize rows to sum to ~1.076    row_sums = cm.sum(dim=1, keepdim=True)77    cm = cm / row_sums78    return cm.tolist()79 80 81def gen_data_batch_stats(scenario: ScenarioParams) -> dict:82    """Generate data batch statistics for the scenario."""83    torch.manual_seed(scenario.seed + 3)84 85    root = scenario.root_cause.value86 87    cm = _gen_confusion_matrix(scenario)88 89    if root == "data_leakage":90        overlap = 0.5 + scenario.leakage_pct * 1.591        overlap = min(overlap, 0.92)92        return {93            "label_distribution": {i: 0.1 for i in range(10)},94            "feature_mean": 0.45 + torch.randn(1).item() * 0.05,95            "feature_std": 0.22 + torch.randn(1).item() * 0.02,96            "null_count": 0,97            "class_overlap_score": overlap,98            "batch_size": 64,99            "duplicate_ratio": scenario.leakage_pct,100            "confusion_matrix": cm,101        }102 103    if root == "overfitting":104        return {105            "label_distribution": {i: 0.1 for i in range(10)},106            "feature_mean": 0.48 + torch.randn(1).item() * 0.03,107            "feature_std": 0.25 + torch.randn(1).item() * 0.02,108            "null_count": 0,109            "class_overlap_score": 0.0,110            "batch_size": 64,111            "duplicate_ratio": 0.0,112            "confusion_matrix": cm,113        }114 115    # Default: normal data116    return {117        "label_distribution": {i: 0.1 for i in range(10)},118        "feature_mean": 0.47 + torch.randn(1).item() * 0.03,119        "feature_std": 0.24 + torch.randn(1).item() * 0.02,120        "null_count": 0,121        "class_overlap_score": 0.0 + torch.randn(1).abs().item() * 0.05,122        "batch_size": 64,123        "duplicate_ratio": 0.0,124        "confusion_matrix": cm,125    }126