Team Ai
Apppublic

tek-wizard/devops-incident-responder

sourceHugging Faceupdated 6mo agoView on Hugging Face
2likes
train.py96 linesDownload Raw Back to scripts
1from __future__ import annotations2 3import argparse4import json5import sys6from pathlib import Path7 8PROJECT_ROOT = Path(__file__).resolve().parents[1]9if str(PROJECT_ROOT) not in sys.path:10    sys.path.insert(0, str(PROJECT_ROOT))11 12 13def _render_prompt(observation: dict) -> str:14    return (15        "You are an incident commander.\n"16        f"Task: {observation['task_id']}\n"17        f"Incident: {observation['incident_brief']}\n"18        f"Global signals: {json.dumps(observation['global_signals'])}\n"19        f"Services: {json.dumps(observation['service_status'])}\n"20        f"Recent events: {json.dumps(observation['recent_events'])}\n"21        f"Discovered facts: {json.dumps(observation['discovered_facts'])}\n"22        "Return the next action as JSON with keys command and target."23    )24 25 26def build_sft_dataset(rollout_path: Path, output_path: Path) -> int:27    count = 028    output_path.parent.mkdir(parents=True, exist_ok=True)29    with rollout_path.open("r", encoding="utf-8") as source, output_path.open("w", encoding="utf-8") as sink:30        for line in source:31            episode = json.loads(line)32            for step in episode["steps"]:33                record = {34                    "prompt": _render_prompt(step["observation"]),35                    "completion": json.dumps(step["action"]),36                    "task_id": episode["task_id"],37                    "seed": episode["seed"],38                    "reward": step["reward"]["value"],39                }40                sink.write(json.dumps(record) + "\n")41                count += 142    return count43 44 45def maybe_run_trl(config: dict, dataset_path: Path) -> None:46    try:47        from datasets import load_dataset48        from transformers import AutoTokenizer49        from trl import SFTConfig, SFTTrainer50    except Exception as exc:  # pragma: no cover - only for hackathon runtime51        raise RuntimeError(52            "TRL / Transformers dependencies are not installed in this environment. "53            "Use this script onsite with the provided compute image."54        ) from exc55 56    dataset = load_dataset("json", data_files=str(dataset_path), split="train")57    tokenizer = AutoTokenizer.from_pretrained(config["model_name"])58    training_args = SFTConfig(59        output_dir=config["output_dir"],60        max_seq_length=config["max_seq_length"],61        per_device_train_batch_size=config["per_device_train_batch_size"],62        gradient_accumulation_steps=config["gradient_accumulation_steps"],63        learning_rate=config["learning_rate"],64        num_train_epochs=config["num_train_epochs"],65        logging_steps=config["logging_steps"],66        save_strategy="epoch",67        report_to=[],68    )69    trainer = SFTTrainer(70        model=config["model_name"],71        train_dataset=dataset,72        args=training_args,73        processing_class=tokenizer,74    )75    trainer.train()76    trainer.save_model(config["output_dir"])77 78 79def main() -> None:80    parser = argparse.ArgumentParser(description="Prepare an SFT dataset and optionally launch TRL training.")81    parser.add_argument("--config", type=Path, default=Path("configs/train.json"))82    parser.add_argument("--rollouts", type=Path, required=True)83    parser.add_argument("--dataset-output", type=Path, default=Path("artifacts/sft_dataset.jsonl"))84    parser.add_argument("--run-training", action="store_true")85    args = parser.parse_args()86 87    config = json.loads(args.config.read_text(encoding="utf-8"))88    count = build_sft_dataset(args.rollouts, args.dataset_output)89    print(f"Prepared {count} SFT records at {args.dataset_output}")90    if args.run_training:91        maybe_run_trl(config, args.dataset_output)92 93 94if __name__ == "__main__":95    main()96