aes-training-scripts / train_aes_lora.py
Manoj-2003's picture
Upload folder using huggingface_hub
66dee2d verified
Raw
History Blame Contribute Delete
14.1 kB
#!/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()