#!/usr/bin/env python3 """ Training script for BGP-LLM3 (forwarding / AFI issues) Incidents: - bgp_missing_routes_in_rib - bgp_next_hop_self_issue - bgp_afi_safi_mismatch Output: ONLY CLI FIX COMMANDS (no explanation). Dataset format matches BGP1/BGP2 style: incident_type, instruction, rules, wazuh_alert, devices, cli_fix """ from transformers import AutoTokenizer, AutoModelForCausalLM, TrainingArguments, Trainer from peft import LoraConfig, get_peft_model from datasets import load_dataset import torch import json # ========================== # PATHS # ========================== BASE_MODEL = r"D:\dKorpesio\git_llm_wazuh\hermes\Hermes-3-Llama-3.1-8B" DATA_FILE = "datasets/bgp3_dataset_v2_1300.jsonl" OUTPUT_DIR = "./bpg_llm/lora_llm_bgp3" tokenizer = AutoTokenizer.from_pretrained(BASE_MODEL) base_model = AutoModelForCausalLM.from_pretrained( BASE_MODEL, torch_dtype=torch.float16, device_map="auto" ) # ========================== # LoRA # ========================== lora_cfg = LoraConfig( r=8, lora_alpha=32, lora_dropout=0.1, target_modules=["q_proj", "v_proj"], bias="none", task_type="CAUSAL_LM" ) model = get_peft_model(base_model, lora_cfg) # ========================== # LOAD DATASET # ========================== dataset = load_dataset("json", data_files=DATA_FILE)["train"] def format_sample(example): devices_section = "\n".join([ "- " + json.dumps(d, ensure_ascii=False) for d in example["devices"] ]) cli_fix = "\n".join(example["cli_fix"]) prompt = f""" ### Instruction: {example['instruction']} Rules: - Prefer minimal-impact actions - Use soft reset when changing policy or AFI/SAFI activation - Do not change BGP timers unless the incident is session stability related - Do NOT provide explanation ### Incident type: {example['incident_type']} ### Wazuh alert: {json.dumps(example['wazuh_alert'], indent=2, ensure_ascii=False)} ### Devices: {devices_section} ### Response (CLI FIX COMMANDS ONLY): {cli_fix} """.strip() tokens = tokenizer( prompt, truncation=True, max_length=1024, padding="max_length" ) tokens["labels"] = tokens["input_ids"].copy() return tokens train_dataset = dataset.map(format_sample) # ========================== # TRAINING ARGS # ========================== training_args = TrainingArguments( output_dir=OUTPUT_DIR, num_train_epochs=3, per_device_train_batch_size=1, gradient_accumulation_steps=4, learning_rate=2e-4, fp16=True, logging_steps=20, save_strategy="epoch", save_total_limit=2, report_to="none" ) trainer = Trainer( model=model, args=training_args, train_dataset=train_dataset ) if __name__ == "__main__": trainer.train() model.save_pretrained(OUTPUT_DIR) print("\n✅ Training complete. Model saved to:", OUTPUT_DIR)