Phase 6: train a second LoRA for real Penpot UI design capability #5
+39
-1
@@ -67,7 +67,13 @@ from pathlib import Path
|
|||||||
import torch
|
import torch
|
||||||
from datasets import Dataset
|
from datasets import Dataset
|
||||||
from peft import LoraConfig, get_peft_model
|
from peft import LoraConfig, get_peft_model
|
||||||
from transformers import AutoModelForCausalLM, AutoTokenizer, Trainer, TrainingArguments
|
from transformers import (
|
||||||
|
AutoModelForCausalLM,
|
||||||
|
AutoTokenizer,
|
||||||
|
Trainer,
|
||||||
|
TrainerCallback,
|
||||||
|
TrainingArguments,
|
||||||
|
)
|
||||||
|
|
||||||
REPO_ROOT = Path(__file__).resolve().parent.parent
|
REPO_ROOT = Path(__file__).resolve().parent.parent
|
||||||
|
|
||||||
@@ -95,6 +101,16 @@ LORA_DROPOUT = float(os.environ.get("LORA_DROPOUT", "0.05"))
|
|||||||
EVAL_STEPS = int(os.environ.get("EVAL_STEPS", "50"))
|
EVAL_STEPS = int(os.environ.get("EVAL_STEPS", "50"))
|
||||||
SAVE_STEPS = int(os.environ.get("SAVE_STEPS", "50"))
|
SAVE_STEPS = int(os.environ.get("SAVE_STEPS", "50"))
|
||||||
GRAD_ACCUM = int(os.environ.get("GRAD_ACCUM", "16"))
|
GRAD_ACCUM = int(os.environ.get("GRAD_ACCUM", "16"))
|
||||||
|
# El default de HF Trainer para per_device_eval_batch_size es 8, INDEPENDIENTE del batch de
|
||||||
|
# entrenamiento. Nunca se fijaba explicitamente. Con evaluation_loop corriendo sin gradient
|
||||||
|
# checkpointing (no hace falta, no hay backward) y sin el truco de recompute, un lote de eval de
|
||||||
|
# 8 secuencias largas se materializa entero -- logits de forma (8, seq_len, vocab_size) en fp32
|
||||||
|
# via convert_to_fp32 -- y con vocab grande eso son decenas de GB en una sola asignacion. Es la
|
||||||
|
# causa real de los dos OOM de esta fase (el smoke run del 32-mas-largos Y la calibracion de 8
|
||||||
|
# ejemplos): en ambos casos NO fue el forward/backward de entrenamiento (medido establemente en
|
||||||
|
# ~74GB para cualquier longitud de 2808 a 3265 tokens), fue el forward de evaluacion agrupando
|
||||||
|
# examples largos en un batch de 8.
|
||||||
|
EVAL_BATCH_SIZE = int(os.environ.get("EVAL_BATCH_SIZE", "1"))
|
||||||
MAX_TOKENS = int(os.environ["MAX_TOKENS"]) if os.environ.get("MAX_TOKENS") else None
|
MAX_TOKENS = int(os.environ["MAX_TOKENS"]) if os.environ.get("MAX_TOKENS") else None
|
||||||
PRESERVE_THINKING = os.environ.get("PRESERVE_THINKING", "").lower() in ("1", "true", "yes")
|
PRESERVE_THINKING = os.environ.get("PRESERVE_THINKING", "").lower() in ("1", "true", "yes")
|
||||||
ALLOW_OVERWRITE = os.environ.get("ALLOW_OVERWRITE", "").lower() in ("1", "true", "yes")
|
ALLOW_OVERWRITE = os.environ.get("ALLOW_OVERWRITE", "").lower() in ("1", "true", "yes")
|
||||||
@@ -129,6 +145,7 @@ def print_banner(args):
|
|||||||
("learning_rate", LEARNING_RATE),
|
("learning_rate", LEARNING_RATE),
|
||||||
("num_train_epochs", NUM_EPOCHS),
|
("num_train_epochs", NUM_EPOCHS),
|
||||||
("grad_accum", GRAD_ACCUM),
|
("grad_accum", GRAD_ACCUM),
|
||||||
|
("eval_batch_size", EVAL_BATCH_SIZE),
|
||||||
("lora r / alpha / dropout", f"{LORA_R} / {LORA_ALPHA} / {LORA_DROPOUT}"),
|
("lora r / alpha / dropout", f"{LORA_R} / {LORA_ALPHA} / {LORA_DROPOUT}"),
|
||||||
("lora scaling (alpha/r)", LORA_ALPHA / LORA_R),
|
("lora scaling (alpha/r)", LORA_ALPHA / LORA_R),
|
||||||
("eval_steps / save_steps", f"{EVAL_STEPS} / {SAVE_STEPS}"),
|
("eval_steps / save_steps", f"{EVAL_STEPS} / {SAVE_STEPS}"),
|
||||||
@@ -228,6 +245,21 @@ def load_examples(tokenizer, path, label):
|
|||||||
return Dataset.from_dict({"input_ids": input_ids_list, "labels": labels_list})
|
return Dataset.from_dict({"input_ids": input_ids_list, "labels": labels_list})
|
||||||
|
|
||||||
|
|
||||||
|
class PerStepMemoryCallback(TrainerCallback):
|
||||||
|
"""Imprime el pico de memoria CUDA de CADA paso (no del promedio de la corrida) y lo
|
||||||
|
resetea, para poder correlacionar un pico puntual con la longitud del ejemplo que lo
|
||||||
|
causo. Se activa con PER_STEP_MEMORY_LOG=1 -- no cambia nada del comportamiento normal
|
||||||
|
de entrenamiento, es solo diagnostico. Existe porque un OOM de la Fase 6 con
|
||||||
|
GRAD_ACCUM=16 no permitia saber CUAL de los 16 micro-batches del paso lo causo."""
|
||||||
|
|
||||||
|
def on_step_end(self, args, state, control, **kwargs):
|
||||||
|
peak_gb = torch.cuda.max_memory_allocated() / (1024 ** 3)
|
||||||
|
reservado_gb = torch.cuda.max_memory_reserved() / (1024 ** 3)
|
||||||
|
print(f"[MEM] paso {state.global_step}: pico_asignado={peak_gb:.2f}GB "
|
||||||
|
f"pico_reservado={reservado_gb:.2f}GB", flush=True)
|
||||||
|
torch.cuda.reset_peak_memory_stats()
|
||||||
|
|
||||||
|
|
||||||
class DataCollatorForCausalLMWithMasking:
|
class DataCollatorForCausalLMWithMasking:
|
||||||
def __init__(self, pad_token_id):
|
def __init__(self, pad_token_id):
|
||||||
self.pad_token_id = pad_token_id
|
self.pad_token_id = pad_token_id
|
||||||
@@ -306,6 +338,7 @@ def main():
|
|||||||
output_dir=str(OUTPUT_DIR),
|
output_dir=str(OUTPUT_DIR),
|
||||||
num_train_epochs=NUM_EPOCHS,
|
num_train_epochs=NUM_EPOCHS,
|
||||||
per_device_train_batch_size=1,
|
per_device_train_batch_size=1,
|
||||||
|
per_device_eval_batch_size=EVAL_BATCH_SIZE,
|
||||||
gradient_accumulation_steps=GRAD_ACCUM,
|
gradient_accumulation_steps=GRAD_ACCUM,
|
||||||
gradient_checkpointing=True,
|
gradient_checkpointing=True,
|
||||||
bf16=True,
|
bf16=True,
|
||||||
@@ -323,12 +356,17 @@ def main():
|
|||||||
report_to="none",
|
report_to="none",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
callbacks = []
|
||||||
|
if os.environ.get("PER_STEP_MEMORY_LOG", "").lower() in ("1", "true", "yes"):
|
||||||
|
callbacks.append(PerStepMemoryCallback())
|
||||||
|
|
||||||
trainer = Trainer(
|
trainer = Trainer(
|
||||||
model=model,
|
model=model,
|
||||||
args=training_args,
|
args=training_args,
|
||||||
train_dataset=train_dataset,
|
train_dataset=train_dataset,
|
||||||
eval_dataset=eval_dataset,
|
eval_dataset=eval_dataset,
|
||||||
data_collator=DataCollatorForCausalLMWithMasking(tokenizer.pad_token_id),
|
data_collator=DataCollatorForCausalLMWithMasking(tokenizer.pad_token_id),
|
||||||
|
callbacks=callbacks,
|
||||||
)
|
)
|
||||||
|
|
||||||
torch.cuda.reset_peak_memory_stats()
|
torch.cuda.reset_peak_memory_stats()
|
||||||
|
|||||||
Reference in New Issue
Block a user