ScottzillaSystems/self-healing-training
1
Self-Healing Training System (SHTS)
Fully autonomous debugging and error recovery for Hugging Face TRL trainers. Add one callback, wrap with `SelfHealingTrainer`, and cut debugging costs to near zero.
 
The Problem
ML training fails constantly:
- CUDA OOM kills jobs at step 847/1000 — restart from scratch
- NaN loss silently corrupts models — discovered hours later
- Loss spikes cascade into divergence — manual intervention required
- DPO plateau at 0.693 loss (= random chance) — wasted GPU hours
- No postmortem — "what step did it die on?"
Each failure costs developer time + GPU credits + schedule delay. At scale, this is millions in wasted compute.
The Solution
SHTS wraps any Hugging Face TRL trainer with four autonomous layers:
┌─────────────────────────────────────────┐
│ LAYER 4: ORCHESTRATION │
│ SelfHealingTrainer retry loop │
│ while not converged: try → recover │
├─────────────────────────────────────────┤
│ LAYER 3: RECOVERY │
│ HealingActions: rollback, halve LR, │
│ halve batch, reclip, clear cache │
├─────────────────────────────────────────┤
│ LAYER 2: DIAGNOSIS │
│ Root-cause classifier: NaN/divergence/ │
│ OOM/data/API — with literature refs │
├─────────────────────────────────────────┤
│ LAYER 1: DETECTION │
│ SelfHealingCallback: loss, gradients, │
│ memory, ZClip adaptive clipping │
└─────────────────────────────────────────┘Quick Start
pip install git+https://huggingface.co/ScottzillaSystems/self-healing-trainingfrom self_healing import SelfHealingTrainer, HealingConfig
from trl import SFTTrainer, SFTConfig
# Your normal training setup
trainer = SFTTrainer(
model=model,
args=SFTConfig(
output_dir="./output",
learning_rate=2e-5,
per_device_train_batch_size=4,
),
train_dataset=dataset,
tokenizer=tokenizer,
)
# Wrap with self-healing — that's it!
sh = SelfHealingTrainer(
trainer,
HealingConfig(
max_recovery_attempts=5,
zclip_enabled=True,
),
)
# Optional: dry-run to catch config errors before full training
sh.dry_run(num_steps=2)
# Train with full autonomy
result = sh.train()What Handles What
Crash Postmortem
Every training interruption produces a postmortem.json:
{
"exit_reason": "exception",
"exception_type": "OutOfMemoryError",
"last_step": 847,
"timestamp": "2026-04-30T15:26:04Z",
"final_metrics": {"loss": 2.15, "grad_norm": 42.3},
"recovery_actions": [
{
"failure": "oom",
"diagnosis": "CUDA Out of Memory. Batch size exceeds GPU capacity.",
"actions": ["halve_batch_size", "enable_gradient_checkpointing", "clear_cache"]
}
],
"running_time_seconds": 1847.3
}Trackio Integration
Set report_to="trackio" in your training args. SHTS emits:
- Alerts at every decision point (INFO/WARN/ERROR)
- Metrics:
healing/recovery_attempts,healing/nan_count,healing/loss_spike_ratio,healing/eval_gap - ZClip metrics:
zclip/raw_grad_norm,zclip/clipped_grad_norm,zclip/z_score,zclip/total_clips
Dashboard URL: https://huggingface.co/spaces/<username>/<trackio-space>
HealingConfig Presets
# Aggressive — for unstable training, low tolerance
config = HealingConfig.aggressive()
# nan_patience=1, zclip_z_threshold=2.0, max_recovery_attempts=10
# Conservative — only intervene on clear failures
config = HealingConfig.conservative()
# nan_patience=10, loss_spike_factor=10.0, zclip_z_threshold=4.0, max_recovery_attempts=2
# Custom
config = HealingConfig(
nan_patience=5,
loss_spike_factor=8.0,
divergence_patience=100,
max_recovery_attempts=3,
zclip_enabled=True,
zclip_z_threshold=3.0,
)Compatibility
Architecture
SelfHealingTrainer.train()
│
├── dry_run() ← Validate setup first
│
└── while not converged:
│
├── trainer.train() ← Run training
│ │
│ ├── on_step_end ← Detect NaN, spikes, divergence
│ ├── on_log ← Monitor gradients (ZClip)
│ ├── on_evaluate ← Check overfitting
│ └── on_exception ← Catch OOM, API, data errors
│
├── [recovery needed?]
│ ├── diagnose ← Classify failure type
│ ├── heal ← Apply recovery actions
│ └── retry ← resume_from_checkpoint=True
│
└── [converged] ← Done!References
License
MIT — use freely, attribution appreciated.
Built autonomously by ML Intern. Questions? Open an issue on the Hub.
