Akiragpu / modules /gpu_trainer.py
Isaac Quarenta
fix(chat): chat nunca parte no Gradio (fallback GPU->sem GPU em 429 ZeroGPU) + NLP/classificadores forcados a CPU + probes de CUDA reais
e48cf8f
Raw History Blame Contribute Delete
15.8 kB
# type: ignore
"""
modules/gpu_trainer.py
================================================================================
TREINO LoRA/QLoRA REAL NA GPU — Space Akiragpu (NVIDIA L4 24GB)
================================================================================
- Base modelo 4-bit (bitsandbytes nf4) + LoRA (PEFT r=16) via TRL SFTTrainer
- Dataset: lista [{"text": ...}] (formato já produzido por
ModelTrainer.prepare_dataset() em modules/treinamento_modelo.py)
- Progresso persistido em <state_path> (JSON) → painel Gradio faz polling
- Artefactos: adapter LoRA + tokenizer + training_report.json
Env:
AKIRA_TRAIN_MODEL base model (default Qwen/Qwen2.5-7B-Instruct)
AKIRA_TRAIN_MAX_STEPS default 100
AKIRA_TRAIN_BATCH_SIZE default 1 (L4 24GB: 1 + grad accum)
AKIRA_TRAIN_GRAD_ACC default 4
AKIRA_TRAIN_LR default 2e-4
AKIRA_TRAIN_SEQ_LEN default 1024
AKIRA_TRAIN_OUT default /data/models/akira-tuned (ou ./models/...)
================================================================================
"""
import os
import json
import time
import threading
import traceback
from typing import Optional, List, Dict, Any
try:
from loguru import logger # type: ignore
except Exception:
class _D:
def info(self, *a, **k): pass
def success(self, *a, **k): pass
def warning(self, *a, **k): pass
def error(self, *a, **k): pass
def debug(self, *a, **k): pass
logger = _D() # type: ignore
DEFAULT_BASE_MODEL = os.getenv("AKIRA_TRAIN_MODEL", "Qwen/Qwen2.5-7B-Instruct")
_train_lock = threading.Lock() # um treino de cada vez
# ======================================================================
# Helpers de path / estado
# ======================================================================
def _data_root() -> str:
"""Diretório persistente da Space (/data) quando existe."""
if os.path.isdir("/data") and os.access("/data", os.W_OK):
return "/data"
return os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "models")
def default_output_dir() -> str:
out = os.getenv("AKIRA_TRAIN_OUT", "").strip()
if out:
return out
if _data_root() == "/data":
return "/data/models/akira-tuned"
return os.path.join(_data_root(), "akira-tuned")
def training_state_path() -> str:
if _data_root() == "/data":
return "/data/training_state.json"
return os.path.join(_data_root(), "training_state.json")
def write_training_state(state: Dict[str, Any]) -> None:
try:
path = training_state_path()
os.makedirs(os.path.dirname(path), exist_ok=True)
tmp = path + ".tmp"
with open(tmp, "w", encoding="utf-8") as f:
json.dump(state, f, ensure_ascii=False, indent=2)
os.replace(tmp, path)
except Exception as e:
logger.warning(f"[GPU-TRAIN] Falha ao gravar estado: {e}")
def read_training_state() -> Dict[str, Any]:
try:
with open(training_state_path(), "r", encoding="utf-8") as f:
return json.load(f)
except Exception:
return {"status": "idle"}
def gpu_info() -> Dict[str, Any]:
"""Info da GPU para a aba Status do app.py."""
info: Dict[str, Any] = {"cuda_available": False, "device": None, "torch": None}
try:
import torch # type: ignore
info["torch"] = torch.__version__
try:
if torch.cuda.is_available():
# ZeroGPU: is_available() é emulado fora dos decorators → prova real
torch.zeros(1, device="cuda")
info["cuda_available"] = True
except Exception as e:
info["zerogpu_fora_do_decorator"] = str(e)[:180]
if info["cuda_available"]:
info["device"] = torch.cuda.get_device_name(0)
props = torch.cuda.get_device_properties(0)
info["vram_total_gb"] = round(props.total_mem / (1024 ** 3), 1)
info["vram_allocated_gb"] = round(torch.cuda.memory_allocated(0) / (1024 ** 3), 2)
info["vram_reserved_gb"] = round(torch.cuda.memory_reserved(0) / (1024 ** 3), 2)
try:
info["cuda_version"] = torch.version.cuda
except Exception:
pass
except Exception as e:
info["error"] = str(e)
return info
def is_gpu_training_available() -> bool:
"""True se CUDA real + transformers + peft + trl + datasets estiverem prontos.
Em ZeroGPU, `torch.cuda.is_available()` devolve True FORA dos decorators
@spaces.GPU mas qualquer operação CUDA levanta erro → aqui prova-se com
uma operação mínima, para o painel dizer logo "precisa de L4 dedicada".
"""
try:
import torch # type: ignore
if not torch.cuda.is_available():
return False
torch.zeros(1, device="cuda")
import transformers # noqa: F401 # type: ignore
import peft # noqa: F401 # type: ignore
import trl # noqa: F401 # type: ignore
import datasets # noqa: F401 # type: ignore
return True
except Exception:
return False
# ======================================================================
# Construção de argumentos compatível com várias versões de TRL/Transformers
# ======================================================================
def _filter_kwargs(cls, kwargs: Dict[str, Any]) -> Dict[str, Any]:
try:
import dataclasses
names = {f.name for f in dataclasses.fields(cls)}
return {k: v for k, v in kwargs.items() if k in names}
except Exception:
return kwargs
def _build_args(out_dir: str, base_model: str, start_ts: float):
import os as _os
max_steps = int(_os.getenv("AKIRA_TRAIN_MAX_STEPS", "100"))
bs = int(_os.getenv("AKIRA_TRAIN_BATCH_SIZE", "1"))
ga = int(_os.getenv("AKIRA_TRAIN_GRAD_ACC", "4"))
lr = float(_os.getenv("AKIRA_TRAIN_LR", "2e-4"))
seq = int(_os.getenv("AKIRA_TRAIN_SEQ_LEN", "1024"))
common: Dict[str, Any] = dict(
output_dir=out_dir,
max_steps=max_steps,
per_device_train_batch_size=bs,
gradient_accumulation_steps=ga,
learning_rate=lr,
logging_steps=5,
save_strategy="steps",
save_steps=max(max_steps, 1),
warmup_ratio=0.03,
lr_scheduler_type="cosine",
bf16=True,
gradient_checkpointing=True,
report_to="none",
remove_unused_columns=False,
dataloader_num_workers=0,
seed=42,
)
# TRL SFTConfig (preferido — trata tokenização/formatting)
try:
from trl import SFTConfig # type: ignore
kwargs = dict(common)
fields = None
try:
import dataclasses
fields = {f.name for f in dataclasses.fields(SFTConfig)}
except Exception:
fields = set()
if "dataset_text_field" in fields:
kwargs["dataset_text_field"] = "text"
if "max_seq_length" in fields:
kwargs["max_seq_length"] = seq
elif "max_length" in fields:
kwargs["max_length"] = seq
if "packing" in fields:
kwargs["packing"] = False
return SFTConfig(**_filter_kwargs(SFTConfig, kwargs))
except Exception as e:
logger.warning(f"[GPU-TRAIN] SFTConfig indisponível ({e}); a usar TrainingArguments")
from transformers import TrainingArguments # type: ignore
return TrainingArguments(**_filter_kwargs(TrainingArguments, common))
def _build_trainer(model, tokenizer, dataset, args, peft_config=None):
from trl import SFTTrainer # type: ignore
try:
return SFTTrainer(model=model, args=args, train_dataset=dataset,
processing_class=tokenizer, peft_config=peft_config)
except TypeError:
# API antiga do TRL
return SFTTrainer(model=model, args=args, train_dataset=dataset,
tokenizer=tokenizer, peft_config=peft_config)
# ======================================================================
# Callback → estado JSON para o painel Gradio
# ======================================================================
def _make_state_callback(base_model: str, start_ts: float):
from transformers import TrainerCallback # type: ignore
class _StateCallback(TrainerCallback):
def on_log(self, args, state, control, logs=None, **kwargs):
if not logs:
return
write_training_state({
"status": "running",
"base_model": base_model,
"step": int(state.global_step or 0),
"max_steps": int(state.max_steps or 0),
"epoch": round(float(state.epoch or 0.0), 2),
"loss": logs.get("loss"),
"learning_rate": logs.get("learning_rate"),
"started_at": start_ts,
"elapsed_s": round(time.time() - start_ts, 1),
})
def on_train_end(self, args, state, control, **kwargs):
write_training_state({
"status": "finalizing",
"base_model": base_model,
"step": int(state.global_step or 0),
"max_steps": int(state.max_steps or 0),
"started_at": start_ts,
"elapsed_s": round(time.time() - start_ts, 1),
})
return _StateCallback()
# ======================================================================
# Treino principal (síncrono — correr em thread)
# ======================================================================
def train_lora(dataset: List[Dict[str, Any]],
base_model: Optional[str] = None,
out_dir: Optional[str] = None) -> Dict[str, Any]:
"""
Executa QLoRA sobre o dataset [{"text": ...}].
Devolve {"success": bool, ...} e mantém /data/training_state.json atualizado.
"""
base = base_model or DEFAULT_BASE_MODEL
out = out_dir or default_output_dir()
start_ts = time.time()
if not _train_lock.acquire(blocking=False):
return {"success": False, "error": "Treino já em execução."}
try:
# ---- Validações ----
rows = [r for r in (dataset or []) if r and (r.get("text") or r.get("messages"))]
if len(rows) < 5:
return {"success": False, "error": f"Dataset insuficiente: {len(rows)} < 5 exemplos."}
if not is_gpu_training_available():
return {"success": False, "error": "GPU/tr_dependencies indisponíveis (torch.cuda/trl/peft)."}
write_training_state({
"status": "starting",
"base_model": base,
"examples": len(rows),
"started_at": start_ts,
"out_dir": out,
})
import torch # type: ignore
from transformers import AutoTokenizer, AutoModelForCausalLM, BitsAndBytesConfig # type: ignore
from peft import LoraConfig, prepare_model_for_kbit_training # type: ignore
from datasets import Dataset # type: ignore
os.makedirs(out, exist_ok=True)
max_steps = int(os.getenv("AKIRA_TRAIN_MAX_STEPS", "100"))
logger.info(f"[GPU-TRAIN] A iniciar QLoRA: {base} | {len(rows)} exemplos | max_steps={max_steps}")
quant = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.bfloat16,
bnb_4bit_use_double_quant=True,
)
tokenizer = AutoTokenizer.from_pretrained(base, trust_remote_code=True)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
model = AutoModelForCausalLM.from_pretrained(
base,
quantization_config=quant,
device_map="auto",
dtype=torch.bfloat16,
trust_remote_code=True,
)
model.config.use_cache = False
# QLoRA: necessário (trl 1.14 não chama isto sozinho) — norm em fp32
# + enable_input_require_grads p/ gradient checkpointing com LoRA
model = prepare_model_for_kbit_training(model)
lora = LoraConfig(
r=16,
lora_alpha=32,
lora_dropout=0.05,
target_modules="all-linear",
bias="none",
task_type="CAUSAL_LM",
)
ds_texts: List[str] = []
for r in rows:
txt = ""
msgs = r.get("messages")
# Preferir chat-template do modelo base (Qwen/Llama/...) via messages
if msgs and getattr(tokenizer, "chat_template", None):
try:
txt = tokenizer.apply_chat_template(
msgs, tokenize=False, add_generation_prompt=False
)
except Exception:
txt = ""
if not txt:
txt = r.get("text") or ""
if txt and txt.strip():
ds_texts.append(txt)
if len(ds_texts) < 5:
return {"success": False, "error": f"Dataset após tokenização: {len(ds_texts)} < 5 exemplos."}
ds = Dataset.from_list([{"text": t} for t in ds_texts])
args = _build_args(out, base, start_ts)
trainer = _build_trainer(model, tokenizer, ds, args, peft_config=lora)
trainer.add_callback(_make_state_callback(base, start_ts))
result = trainer.train()
# ---- Guardar artefactos ----
trainer.save_model(out)
try:
tokenizer.save_pretrained(out)
except Exception:
pass
train_loss = None
try:
train_loss = float(result.training_loss)
except Exception:
pass
report = {
"success": True,
"base_model": base,
"out_dir": out,
"examples": len(rows),
"global_step": int(getattr(result, "global_step", 0) or 0),
"train_loss": train_loss,
"elapsed_s": round(time.time() - start_ts, 1),
"finished_at": time.time(),
}
with open(os.path.join(out, "training_report.json"), "w", encoding="utf-8") as f:
json.dump(report, f, ensure_ascii=False, indent=2)
write_training_state({
"status": "done",
"base_model": base,
"out_dir": out,
"step": report["global_step"],
"max_steps": max_steps,
"loss": train_loss,
"elapsed_s": report["elapsed_s"],
"started_at": start_ts,
})
logger.success(f"[GPU-TRAIN] Concluído em {report['elapsed_s']}s → {out}")
# ---- Libertar VRAM (modelo de treino não fica residente) ----
try:
del trainer
del model
import gc
gc.collect()
torch.cuda.empty_cache()
except Exception:
pass
return report
except Exception as e:
err = f"{e}"
logger.error(f"[GPU-TRAIN] Falha: {err}\n{traceback.format_exc()}")
write_training_state({
"status": "failed",
"base_model": base,
"error": err[:500],
"started_at": start_ts,
"elapsed_s": round(time.time() - start_ts, 1),
})
try:
import torch # type: ignore
import gc
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
except Exception:
pass
return {"success": False, "error": err}
finally:
_train_lock.release()
__all__ = [
"train_lora",
"is_gpu_training_available",
"gpu_info",
"read_training_state",
"write_training_state",
"training_state_path",
"default_output_dir",
]