Files
qwen3-6-lora/scripts/10_train.py
T

191 lines
6.7 KiB
Python

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