"""Fase 3: entrena el LoRA de Qwen3.6-35B-A3B sobre data/train.jsonl / data/eval.jsonl. Corre DENTRO del contenedor `qwen-lora-train` en spark (necesita transformers/peft/accelerate ya instalados ahi, y el checkpoint base en MODEL_PATH). Invocar via: docker exec qwen-lora-train python3 /workspace/ai-projects/qwen3-6-lora/scripts/10_train.py Tope de pasos para el dry-run via env var MAX_STEPS (o --max-steps N), sin tocar el resto de la config de TrainingArguments. Masking manual (no trl.SFTTrainer): usa data/chat_template_train.jinja (con tags {% generation %}) para que tokenizer.apply_chat_template devuelva assistant_masks, y arma labels = input_ids donde assistant_masks==1, -100 en el resto (nunca entrena sobre system/user/tool). """ import argparse import os os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") import sys from pathlib import Path import torch from datasets import Dataset from peft import LoraConfig, get_peft_model from transformers import AutoModelForCausalLM, AutoTokenizer, Trainer, TrainingArguments REPO_ROOT = Path(__file__).resolve().parent.parent MODEL_PATH = os.environ.get("MODEL_PATH", "/workspace/ft-models/Qwen--Qwen3.6-35B-A3B") TRAIN_CHAT_TEMPLATE_PATH = REPO_ROOT / "data" / "chat_template_train.jinja" TRAIN_FILE = REPO_ROOT / "data" / "train.jsonl" EVAL_FILE = REPO_ROOT / "data" / "eval.jsonl" OUTPUT_DIR = REPO_ROOT / "out" / "lora-adapter" TARGET_MODULES = [ "q_proj", "k_proj", "v_proj", "o_proj", "in_proj_qkv", "in_proj_z", "in_proj_a", "in_proj_b", "out_proj", "shared_expert.gate_proj", "shared_expert.up_proj", "shared_expert.down_proj", ] def parse_args(): parser = argparse.ArgumentParser() parser.add_argument("--max-steps", type=int, default=None) args = parser.parse_args() if args.max_steps is None: env_val = os.environ.get("MAX_STEPS") args.max_steps = int(env_val) if env_val else None return args def load_examples(tokenizer, path): import json input_ids_list = [] labels_list = [] with open(path, encoding="utf-8") as f: for line in f: line = line.strip() if not line: continue example = json.loads(line) rendered = tokenizer.apply_chat_template( example["messages"], tools=example.get("tools"), tokenize=True, return_assistant_tokens_mask=True, return_dict=True, add_generation_prompt=False, ) input_ids = rendered["input_ids"] assistant_masks = rendered["assistant_masks"] if sum(assistant_masks) == 0: raise AssertionError(f"assistant_masks vacia para un ejemplo de {path}") labels = [tok if mask == 1 else -100 for tok, mask in zip(input_ids, assistant_masks)] input_ids_list.append(input_ids) labels_list.append(labels) return Dataset.from_dict({"input_ids": input_ids_list, "labels": labels_list}) class DataCollatorForCausalLMWithMasking: def __init__(self, pad_token_id): self.pad_token_id = pad_token_id def __call__(self, features): max_len = max(len(f["input_ids"]) for f in features) input_ids = [] labels = [] attention_mask = [] for f in features: ids = f["input_ids"] lbl = f["labels"] pad_len = max_len - len(ids) input_ids.append(ids + [self.pad_token_id] * pad_len) labels.append(lbl + [-100] * pad_len) attention_mask.append([1] * len(ids) + [0] * pad_len) return { "input_ids": torch.tensor(input_ids, dtype=torch.long), "labels": torch.tensor(labels, dtype=torch.long), "attention_mask": torch.tensor(attention_mask, dtype=torch.long), } def main(): args = parse_args() print(f"[INFO] cargando tokenizer desde {MODEL_PATH}") tokenizer = AutoTokenizer.from_pretrained(MODEL_PATH) tokenizer.chat_template = TRAIN_CHAT_TEMPLATE_PATH.read_text(encoding="utf-8") if tokenizer.pad_token_id is None: tokenizer.pad_token = tokenizer.eos_token print(f"[INFO] tokenizando {TRAIN_FILE}") train_dataset = load_examples(tokenizer, TRAIN_FILE) print(f"[INFO] tokenizando {EVAL_FILE}") eval_dataset = load_examples(tokenizer, EVAL_FILE) print(f"[INFO] train={len(train_dataset)} eval={len(eval_dataset)}") print(f"[INFO] cargando modelo desde {MODEL_PATH}") model = AutoModelForCausalLM.from_pretrained( MODEL_PATH, dtype=torch.bfloat16, attn_implementation="flash_attention_2", ) lora_config = LoraConfig( target_modules=TARGET_MODULES, r=32, lora_alpha=64, lora_dropout=0.05, task_type="CAUSAL_LM", bias="none", ) model = get_peft_model(model, lora_config) model.print_trainable_parameters() trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad) total_params = sum(p.numel() for p in model.parameters()) trainable_pct = 100 * trainable_params / total_params print(f"[INFO] modulos entrenables por sufijo objetivo: {TARGET_MODULES}") if trainable_pct <= 0 or trainable_pct > 20: raise AssertionError( f"% entrenable fuera de rango razonable ({trainable_pct:.4f}%) — target_modules " "probablemente mal aplicado, abortando antes de entrenar" ) training_args = TrainingArguments( output_dir=str(OUTPUT_DIR), num_train_epochs=2, per_device_train_batch_size=1, gradient_accumulation_steps=16, gradient_checkpointing=True, bf16=True, optim="adamw_8bit", learning_rate=1e-4, lr_scheduler_type="cosine", warmup_ratio=0.03, eval_strategy="steps", eval_steps=50, save_strategy="steps", save_steps=50, save_total_limit=3, logging_steps=5, max_steps=args.max_steps if args.max_steps else -1, report_to="none", ) trainer = Trainer( model=model, args=training_args, train_dataset=train_dataset, eval_dataset=eval_dataset, data_collator=DataCollatorForCausalLMWithMasking(tokenizer.pad_token_id), ) torch.cuda.reset_peak_memory_stats() trainer.train() peak_mem_gb = torch.cuda.max_memory_allocated() / (1024 ** 3) print(f"[INFO] pico de memoria CUDA (max_memory_allocated): {peak_mem_gb:.2f} GB") if args.max_steps is None: trainer.save_model(str(OUTPUT_DIR)) tokenizer.save_pretrained(str(OUTPUT_DIR)) print(f"[INFO] adapter final guardado en {OUTPUT_DIR}") if __name__ == "__main__": main()