tek-wizard/devops-incident-responder
2
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 