ai-network-llms / training /train_switch1.py
JoeiBanana's picture
Upload batch 8/8
3453f3d verified
Raw
History Blame
3.09 kB
#!/usr/bin/env python3
"""
Train SWITCH-LLM1 for L2 switching incidents:
- MAC flapping
- STP topology change
- Port err-disable
Dataset: switch1_dataset_900.jsonl
Output: ONLY CLI FIX COMMANDS (no explanation)
"""
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/switch1_dataset_900.jsonl"
OUTPUT_DIR = "./switch_llm/lora_llm_switch1"
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"].train_test_split(test_size=0.1, seed=42)
train_data, eval_data = dataset["train"], dataset["test"]
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:
- Output ONLY valid Cisco IOS/IOS-XE CLI commands
- Do NOT include show/debug commands
- Prefer minimal-impact actions (avoid unnecessary global changes)
- If you shutdown a port, include 'no shutdown' only when appropriate for recovery
- 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 = train_data.map(format_sample)
eval_dataset = eval_data.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,
eval_dataset=eval_dataset
)
if __name__ == "__main__":
trainer.train()
model.save_pretrained(OUTPUT_DIR)
print("\n✅ Training complete. Model saved to:", OUTPUT_DIR)