| |
| """ |
| 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() |
|
|