#!/usr/bin/env python3 """ AES Security IP LoRA Fine-Tuning Script Trains a fresh LoRA adapter on top of the fully merged model (Qwen2.5-7B + V1 + V2 + V3 + V4 + SRAM + I2CS all baked in) using aes_all.jsonl (AES cryptographic engine RTL/TB/docs). Key features: - Loads the all-merged model as base (no LoRA stacking needed) - Fresh LoRA adapter (r=32, alpha=64) for AES IP - Assistant-only loss masking via chat template {% generation %} tags - bf16 precision, gradient checkpointing for 7B on single GPU - TensorBoard logging, checkpoint saving per epoch - Auto-adjusts batch size based on max_seq_length Usage: python3 train_aes_lora.py python3 train_aes_lora.py --epochs 5 --lr 1e-4 """ import argparse import json import logging import sys from pathlib import Path import torch from datasets import Dataset from transformers import AutoModelForCausalLM, AutoTokenizer, TrainerCallback from peft import LoraConfig, get_peft_model, TaskType from trl import SFTTrainer, SFTConfig class FileLoggingCallback(TrainerCallback): def __init__(self, logger): self.logger = logger self.step_start_time = None def on_log(self, args, state, control, logs=None, **kwargs): if logs is None: return current_step = state.global_step metrics_str = ", ".join(f"{k}={v:.6f}" if isinstance(v, float) else f"{k}={v}" for k, v in logs.items()) self.logger.info(f" [step {current_step}] {metrics_str}") def on_epoch_begin(self, args, state, control, **kwargs): self.logger.info(f" --- Epoch {int(state.epoch) + 1}/{args.num_train_epochs} beginning (global_step={state.global_step}) ---") def on_epoch_end(self, args, state, control, **kwargs): self.logger.info(f" --- Epoch {int(state.epoch) + 1}/{args.num_train_epochs} ended (global_step={state.global_step}) ---") def on_step_begin(self, args, state, control, **kwargs): import time self.step_start_time = time.time() def on_step_end(self, args, state, control, **kwargs): import time if self.step_start_time is not None: elapsed = time.time() - self.step_start_time if state.global_step % 5 == 0: self.logger.debug(f" [step {state.global_step}/{state.max_steps}] step_time={elapsed:.2f}s") WORKSPACE = Path("/workspace/elinnos") DEFAULT_MERGED_BASE = WORKSPACE / "merged_models" / "elinnos_all_merged_final" DEFAULT_TRAIN_DATA = WORKSPACE / "aes_training" / "data" / "aes_train.jsonl" DEFAULT_VAL_DATA = WORKSPACE / "aes_training" / "data" / "aes_val.jsonl" DEFAULT_OUTPUT_DIR = WORKSPACE / "elinnos-qwen2.5-7b-aes-lora" CHAT_TEMPLATE_SRC = WORKSPACE / "elinnos-qwen2.5-7b-multi-ip-lora-v4" / "chat_template.jinja" SEQ_LENGTH_FILE = WORKSPACE / "aes_training" / "data" / "recommended_seq_length.txt" AES_LORA_R = 32 AES_LORA_ALPHA = 64 AES_LORA_DROPOUT = 0.05 AES_TARGET_MODULES = ["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"] def setup_logging(output_dir: Path) -> logging.Logger: output_dir.mkdir(parents=True, exist_ok=True) log_file = output_dir / "training_aes.log" logger = logging.getLogger("train_aes") logger.setLevel(logging.DEBUG) fh = logging.FileHandler(str(log_file), mode="w") fh.setLevel(logging.DEBUG) fh.setFormatter(logging.Formatter("%(asctime)s [%(levelname)s] %(message)s")) ch = logging.StreamHandler(sys.stdout) ch.setLevel(logging.INFO) ch.setFormatter(logging.Formatter("%(asctime)s [%(levelname)s] %(message)s")) logger.addHandler(fh) logger.addHandler(ch) return logger def read_recommended_seq_length() -> int: if not SEQ_LENGTH_FILE.exists(): print(f"WARNING: {SEQ_LENGTH_FILE} not found. Defaulting to 8192.") return 8192 return int(SEQ_LENGTH_FILE.read_text().strip()) def parse_args(): p = argparse.ArgumentParser(description="Train AES LoRA on aes_all.jsonl (fresh LoRA on all-merged base)") p.add_argument("--merged_base", type=str, default=str(DEFAULT_MERGED_BASE)) p.add_argument("--train_data", type=str, default=str(DEFAULT_TRAIN_DATA)) p.add_argument("--val_data", type=str, default=str(DEFAULT_VAL_DATA)) p.add_argument("--output_dir", type=str, default=str(DEFAULT_OUTPUT_DIR)) p.add_argument("--epochs", type=int, default=3) p.add_argument("--lr", type=float, default=1e-4) p.add_argument("--force_seq_length", type=int, default=None) p.add_argument("--lora_r", type=int, default=AES_LORA_R) p.add_argument("--lora_alpha", type=int, default=AES_LORA_ALPHA) p.add_argument("--warmup_ratio", type=float, default=0.03) p.add_argument("--weight_decay", type=float, default=0.01) return p.parse_args() def load_jsonl_dataset(path: Path) -> Dataset: records = [] with path.open("r", encoding="utf-8") as f: for line in f: line = line.strip() if not line: continue rec = json.loads(line) records.append({"messages": rec["messages"]}) return Dataset.from_list(records) def main(): args = parse_args() merged_base = Path(args.merged_base).resolve() train_data = Path(args.train_data).resolve() val_data = Path(args.val_data).resolve() output_dir = Path(args.output_dir).resolve() logger = setup_logging(output_dir) logger.info("=" * 70) logger.info(" AES SECURITY IP LORA TRAINING") logger.info(" Base: All-merged model (V1+V2+V3+V4+SRAM+I2CS baked in)") logger.info(" Fresh LoRA adapter on aes_all.jsonl (AES cryptographic engine)") logger.info("=" * 70) logger.info("Pre-flight checks:") checks_ok = True if not (merged_base / "config.json").is_file(): logger.error(f"Merged base model not found at {merged_base}") logger.error("Run merge_all_for_aes.py first!") checks_ok = False else: logger.info(f" [OK] Merged base: {merged_base}") if not train_data.is_file(): logger.error(f"Training data not found: {train_data}") logger.error("Run split_aes_dataset.py first!") checks_ok = False else: logger.info(f" [OK] Train data: {train_data}") if not val_data.is_file(): logger.error(f"Validation data not found: {val_data}") checks_ok = False else: logger.info(f" [OK] Val data: {val_data}") if not CHAT_TEMPLATE_SRC.is_file(): logger.error(f"Chat template not found: {CHAT_TEMPLATE_SRC}") checks_ok = False else: logger.info(f" [OK] Chat template: {CHAT_TEMPLATE_SRC}") if not checks_ok: logger.error("Pre-flight checks FAILED. Aborting.") sys.exit(1) max_seq_length = args.force_seq_length if args.force_seq_length else read_recommended_seq_length() logger.info(f" Max seq length: {max_seq_length}") if max_seq_length > 4096: batch_size = 1 grad_accum = 16 logger.warning(f"Seq length {max_seq_length} > 4096 -- batch_size=1, grad_accum=16 (effective=16)") else: batch_size = 2 grad_accum = 8 logger.info(f"Seq length {max_seq_length} <= 4096 -- batch_size=2, grad_accum=8 (effective=16)") if torch.cuda.is_available(): gpu_name = torch.cuda.get_device_name(0) gpu_mem = torch.cuda.get_device_properties(0).total_memory / (1024**3) logger.info(f" GPU: {gpu_name} ({gpu_mem:.1f} GB)") logger.info(f" CPU cores: {torch.get_num_threads()}") else: logger.error("CUDA not available! Cannot train on CPU. Aborting.") sys.exit(1) logger.info("Loading tokenizer from merged base...") tokenizer = AutoTokenizer.from_pretrained(str(merged_base), trust_remote_code=True) if tokenizer.pad_token is None: tokenizer.pad_token = tokenizer.eos_token logger.info(f"Tokenizer loaded. Pad token: {tokenizer.pad_token}") if CHAT_TEMPLATE_SRC.is_file(): tokenizer.chat_template = CHAT_TEMPLATE_SRC.read_text() logger.info(f" Chat template loaded from {CHAT_TEMPLATE_SRC}") logger.info("Loading datasets...") train_dataset = load_jsonl_dataset(train_data) eval_dataset = load_jsonl_dataset(val_data) logger.info(f"Train samples: {len(train_dataset)}") logger.info(f"Eval samples: {len(eval_dataset)}") logger.info("Loading all-merged base model (bfloat16)...") model = AutoModelForCausalLM.from_pretrained( str(merged_base), torch_dtype=torch.bfloat16, device_map="auto", trust_remote_code=True, low_cpu_mem_usage=True, ) model.config.use_cache = False logger.info(f"Model loaded. Device map: {getattr(model, 'hf_device_map', 'N/A')}") logger.info("Applying AES LoRA configuration:") logger.info(f" r={args.lora_r}, alpha={args.lora_alpha}, dropout={AES_LORA_DROPOUT}") logger.info(f" target_modules={AES_TARGET_MODULES}") logger.info(f" bias=none, task_type=CAUSAL_LM") lora_config = LoraConfig( r=args.lora_r, lora_alpha=args.lora_alpha, lora_dropout=AES_LORA_DROPOUT, target_modules=AES_TARGET_MODULES, bias="none", task_type=TaskType.CAUSAL_LM, ) model = get_peft_model(model, lora_config) model.print_trainable_parameters() logger.info("Building SFTConfig...") sft_config = SFTConfig( output_dir=str(output_dir), num_train_epochs=args.epochs, per_device_train_batch_size=batch_size, per_device_eval_batch_size=batch_size, gradient_accumulation_steps=grad_accum, learning_rate=args.lr, lr_scheduler_type="cosine", warmup_ratio=args.warmup_ratio, weight_decay=args.weight_decay, optim="adamw_torch_fused", bf16=True, fp16=False, gradient_checkpointing=True, gradient_checkpointing_kwargs={"use_reentrant": False}, logging_steps=5, eval_strategy="epoch", save_strategy="epoch", save_total_limit=3, max_length=max_seq_length, packing=False, assistant_only_loss=True, report_to="tensorboard", seed=42, data_seed=42, dataset_text_field=None, ) logger.info(f"SFTConfig summary:") logger.info(f" epochs={args.epochs}, lr={args.lr}, bf16=True") logger.info(f" batch_size={batch_size}, grad_accum={grad_accum}, effective_batch={batch_size * grad_accum}") logger.info(f" max_seq_length={max_seq_length}") logger.info(f" warmup_ratio={args.warmup_ratio}, weight_decay={args.weight_decay}") logger.info(f" gradient_checkpointing=True, assistant_only_loss=True") logger.info("Initializing SFTTrainer...") trainer = SFTTrainer( model=model, args=sft_config, train_dataset=train_dataset, eval_dataset=eval_dataset, processing_class=tokenizer, callbacks=[FileLoggingCallback(logger)], ) logger.info("=" * 70) logger.info(" STARTING AES TRAINING") logger.info("=" * 70) try: trainer.train() logger.info("=" * 70) logger.info(" TRAINING COMPLETE") logger.info("=" * 70) except Exception as e: logger.error(f"Training failed with error: {e}", exc_info=True) raise logger.info("Saving final AES LoRA adapter...") model.save_pretrained(str(output_dir)) tokenizer.save_pretrained(str(output_dir)) if CHAT_TEMPLATE_SRC.is_file(): import shutil shutil.copy2(str(CHAT_TEMPLATE_SRC), str(output_dir / "chat_template.jinja")) logger.info(f" Chat template copied to {output_dir / 'chat_template.jinja'}") logger.info(f"AES LoRA adapter saved to: {output_dir}") metadata = { "aes_lora_config": { "r": args.lora_r, "lora_alpha": args.lora_alpha, "lora_dropout": AES_LORA_DROPOUT, "target_modules": AES_TARGET_MODULES, "bias": "none", "task_type": "CAUSAL_LM", }, "training_config": { "epochs": args.epochs, "learning_rate": args.lr, "batch_size": batch_size, "grad_accum": grad_accum, "effective_batch": batch_size * grad_accum, "max_seq_length": max_seq_length, "bf16": True, "gradient_checkpointing": True, "assistant_only_loss": True, "warmup_ratio": args.warmup_ratio, "weight_decay": args.weight_decay, }, "paths": { "merged_base": str(merged_base), "train_data": str(train_data), "val_data": str(val_data), "output_dir": str(output_dir), }, "dataset_stats": { "train_samples": len(train_dataset), "eval_samples": len(eval_dataset), "total_samples": len(train_dataset) + len(eval_dataset), }, "lineage": { "base": "Qwen2.5-7B-Instruct", "v1_lora": "r=64, alpha=128, 10k I2C Master RTL samples", "v2_lora": "r=32, alpha=64, 920 Unified I2C Master+Slave samples", "v3_lora": "r=32, alpha=64, combined I2C dataset (incremental on V2 merged)", "v4_lora": "r=32, alpha=64, 11400 multi-IP peripheral samples (GPIO/SPI/UART/APB Timer/PWM/QSPI)", "sram_lora": "r=32, alpha=64, 540 AHB SRAM / memory IP samples (incremental on V4 merged)", "i2cs_lora": "r=32, alpha=64, I2C Slave RTL samples (incremental on V3 merged)", "aes_lora": f"r={args.lora_r}, alpha={args.lora_alpha}, {len(train_dataset)} AES security IP samples (fresh LoRA on all-merged base)", }, } metadata_path = output_dir / "aes_training_metadata.json" with metadata_path.open("w") as f: json.dump(metadata, f, indent=2) logger.info(f"Training metadata saved to: {metadata_path}") logger.info("=" * 70) logger.info(" AES TRAINING COMPLETE -- ALL DONE") logger.info("=" * 70) if __name__ == "__main__": main()