unsloth/studio/backend/core/training/worker.py
2026-03-10 21:29:51 +00:00

769 lines
30 KiB
Python

# SPDX-License-Identifier: AGPL-3.0-only - See /studio/LICENSE.AGPL-3.0
# Copyright © 2025 Unsloth AI
"""
Training subprocess entry point.
Each training job runs in a fresh subprocess (mp.get_context("spawn")).
This gives us a clean Python interpreter with no stale module state —
solving the transformers version-switching problem completely.
Pattern follows core/data_recipe/jobs/worker.py.
"""
from __future__ import annotations
import logging
import os
import sys
import time
import traceback
from pathlib import Path
from typing import Any
logger = logging.getLogger(__name__)
def _activate_transformers_version(model_name: str, project_root: str) -> None:
"""Activate the correct transformers version BEFORE any ML imports.
If the model needs transformers 5.x, prepend the pre-installed .venv_t5/
directory to sys.path. Otherwise do nothing (default 4.57.x in .venv/).
"""
# Ensure backend is on path for utils imports
backend_path = os.path.join(project_root, "studio", "backend")
if backend_path not in sys.path:
sys.path.insert(0, backend_path)
from utils.transformers_version import needs_transformers_5, _resolve_base_model
resolved = _resolve_base_model(model_name)
if needs_transformers_5(resolved):
venv_t5 = os.path.join(project_root, ".venv_t5")
if os.path.isdir(venv_t5):
sys.path.insert(0, venv_t5)
logger.info("Activated transformers 5.x from %s", venv_t5)
else:
# Fallback: pip install at runtime (slower, ~10-15s)
logger.warning(".venv_t5 not found at %s — installing at runtime", venv_t5)
import subprocess as sp
os.makedirs(venv_t5, exist_ok=True)
r1 = sp.run(
[sys.executable, "-m", "pip", "install", "--target", venv_t5,
"--no-deps", "transformers==5.2.0"],
stdout=sp.PIPE, stderr=sp.STDOUT,
)
r2 = sp.run(
[sys.executable, "-m", "pip", "install", "--target", venv_t5,
"--no-deps", "huggingface_hub==1.3.0"],
stdout=sp.PIPE, stderr=sp.STDOUT,
)
if r1.returncode != 0 or r2.returncode != 0:
raise RuntimeError(
f"Failed to install transformers 5.x into {venv_t5}. "
f"pip returncode: transformers={r1.returncode}, huggingface_hub={r2.returncode}"
)
sys.path.insert(0, venv_t5)
# Propagate to child subprocesses (e.g. GGUF converter)
_pp = os.environ.get("PYTHONPATH", "")
os.environ["PYTHONPATH"] = venv_t5 + (os.pathsep + _pp if _pp else "")
else:
logger.info("Using default transformers (4.57.x) for %s", model_name)
def run_training_process(
*,
event_queue: Any,
stop_queue: Any,
config: dict,
) -> None:
"""Subprocess entrypoint. Fresh Python — no stale module state.
Args:
event_queue: mp.Queue for sending progress/status/error events to parent.
stop_queue: mp.Queue for receiving stop commands from parent.
config: Training configuration dict with all parameters.
"""
os.environ["TOKENIZERS_PARALLELISM"] = "false"
project_root = config["project_root"]
model_name = config["model_name"]
# ── 1. Activate correct transformers version BEFORE any ML imports ──
try:
_activate_transformers_version(model_name, project_root)
except Exception as exc:
event_queue.put({
"type": "error",
"error": f"Failed to activate transformers version: {exc}",
"stack": traceback.format_exc(limit=20),
"ts": time.time(),
})
return
# ── 1b. On Windows, check Triton availability (must be before import torch) ──
if sys.platform == "win32":
try:
import triton # noqa: F401
logger.info("Triton available — torch.compile enabled")
except ImportError:
os.environ["TORCHDYNAMO_DISABLE"] = "1"
logger.warning(
"Triton not found on Windows — torch.compile disabled. "
'Install for better performance: pip install "triton-windows<3.7"'
)
# ── 2. Now import ML libraries (fresh in this clean process) ──
try:
_send_status(event_queue, "Importing ML libraries...")
backend_path = os.path.join(project_root, "studio", "backend")
if backend_path not in sys.path:
sys.path.insert(0, backend_path)
from core.training.trainer import UnslothTrainer, TrainingProgress
import transformers
logger.info("Subprocess loaded transformers %s", transformers.__version__)
except Exception as exc:
event_queue.put({
"type": "error",
"error": f"Failed to import ML libraries: {exc}",
"stack": traceback.format_exc(limit=20),
"ts": time.time(),
})
return
# ── 2b. EMBEDDING MODEL FAST-PATH ──
# Embedding models use a completely different pipeline (FastSentenceTransformer
# + SentenceTransformerTrainer + MultipleNegativesRankingLoss) so we branch
# early and handle the entire flow in a self-contained function.
if config.get("is_embedding", False):
try:
_run_embedding_training(event_queue, stop_queue, config)
except Exception as exc:
event_queue.put({
"type": "error",
"error": str(exc),
"stack": traceback.format_exc(limit=20),
"ts": time.time(),
})
return
# ── 3. Create a fresh trainer instance ──
trainer = UnslothTrainer()
# Wire up progress callback → event_queue
def _on_progress(progress: TrainingProgress):
has_train_loss = progress.step >= 0 and progress.loss > 0
has_eval_loss = progress.eval_loss is not None
if has_train_loss or has_eval_loss:
event_queue.put({
"type": "progress",
"step": progress.step,
"epoch": progress.epoch,
"loss": progress.loss,
"learning_rate": progress.learning_rate,
"total_steps": progress.total_steps,
"elapsed_seconds": progress.elapsed_seconds,
"eta_seconds": progress.eta_seconds,
"grad_norm": progress.grad_norm,
"num_tokens": progress.num_tokens,
"eval_loss": progress.eval_loss,
"status_message": progress.status_message,
"ts": time.time(),
})
if progress.status_message:
_send_status(event_queue, progress.status_message)
trainer.add_progress_callback(_on_progress)
# Wire up stop_queue polling to trainer.should_stop
import threading
import queue as _queue
def _poll_stop():
while True:
try:
msg = stop_queue.get(timeout=1.0)
if msg and msg.get("type") == "stop":
save = msg.get("save", True)
trainer.should_stop = True
trainer.save_on_stop = save
logger.info("Stop signal received (save=%s)", save)
return
except _queue.Empty:
continue
except (EOFError, OSError):
return
stop_thread = threading.Thread(target=_poll_stop, daemon=True)
stop_thread.start()
# ── 4. Execute the training pipeline ──
try:
hf_token = config.get("hf_token", "")
hf_token = hf_token if hf_token and hf_token.strip() else None
# Load model
_send_status(event_queue, "Loading model...")
success = trainer.load_model(
model_name=model_name,
max_seq_length=config["max_seq_length"],
load_in_4bit=config["load_in_4bit"],
hf_token=hf_token,
is_dataset_image=config.get("is_dataset_image", False),
is_dataset_audio=config.get("is_dataset_audio", False),
trust_remote_code=config.get("trust_remote_code", False),
)
if not success or trainer.should_stop:
if trainer.should_stop:
event_queue.put({"type": "complete", "output_dir": None, "ts": time.time()})
else:
error_msg = trainer.training_progress.error or "Failed to load model"
event_queue.put({
"type": "error",
"error": error_msg,
"stack": "", "ts": time.time(),
})
return
# Prepare model (LoRA or full finetuning)
training_type = config.get("training_type", "LoRA/QLoRA")
use_lora = (training_type == "LoRA/QLoRA")
if use_lora:
_send_status(event_queue, "Configuring LoRA adapters...")
success = trainer.prepare_model_for_training(
use_lora=True,
finetune_vision_layers=config.get("finetune_vision_layers", True),
finetune_language_layers=config.get("finetune_language_layers", True),
finetune_attention_modules=config.get("finetune_attention_modules", True),
finetune_mlp_modules=config.get("finetune_mlp_modules", True),
target_modules=config.get("target_modules"),
lora_r=config.get("lora_r", 16),
lora_alpha=config.get("lora_alpha", 16),
lora_dropout=config.get("lora_dropout", 0.0),
use_gradient_checkpointing=config.get("gradient_checkpointing", "unsloth"),
use_rslora=config.get("use_rslora", False),
use_loftq=config.get("use_loftq", False),
)
else:
_send_status(event_queue, "Preparing model for full finetuning...")
success = trainer.prepare_model_for_training(use_lora=False)
if not success or trainer.should_stop:
if trainer.should_stop:
event_queue.put({"type": "complete", "output_dir": None, "ts": time.time()})
else:
event_queue.put({
"type": "error",
"error": trainer.training_progress.error or "Failed to prepare model",
"stack": "", "ts": time.time(),
})
return
# Load dataset
_send_status(event_queue, "Loading and formatting dataset...")
hf_dataset = config.get("hf_dataset", "")
dataset_result = trainer.load_and_format_dataset(
dataset_source=hf_dataset if hf_dataset and hf_dataset.strip() else None,
format_type=config.get("format_type", ""),
local_datasets=config.get("local_datasets") or None,
custom_format_mapping=config.get("custom_format_mapping"),
subset=config.get("subset"),
train_split=config.get("train_split", "train"),
eval_split=config.get("eval_split"),
eval_steps=config.get("eval_steps", 0.00),
dataset_slice_start=config.get("dataset_slice_start"),
dataset_slice_end=config.get("dataset_slice_end"),
)
if isinstance(dataset_result, tuple):
dataset, eval_dataset = dataset_result
else:
dataset = dataset_result
eval_dataset = None
# Disable eval if eval_steps <= 0
eval_steps = config.get("eval_steps", 0.00)
if eval_steps is not None and float(eval_steps) <= 0:
eval_dataset = None
# Tell the parent process that eval is configured so the frontend
# shows "Waiting for first evaluation step..." instead of "not configured"
if eval_dataset is not None:
event_queue.put({
"type": "eval_configured",
"ts": time.time(),
})
if dataset is None or trainer.should_stop:
if trainer.should_stop:
event_queue.put({"type": "complete", "output_dir": None, "ts": time.time()})
else:
event_queue.put({
"type": "error",
"error": trainer.training_progress.error or "Failed to load dataset",
"stack": "", "ts": time.time(),
})
return
# Convert learning rate
try:
lr_value = float(config.get("learning_rate", "2e-4"))
except ValueError:
event_queue.put({
"type": "error",
"error": f"Invalid learning rate: {config.get('learning_rate')}",
"stack": "", "ts": time.time(),
})
return
# Generate output dir
output_dir = config.get("output_dir")
if not output_dir:
output_dir = f"./outputs/{model_name.replace('/', '_')}_{int(time.time())}"
# Start training (directly — no inner thread, we ARE the subprocess)
_send_status(event_queue, "Starting training...")
max_steps = config.get("max_steps", 0)
save_steps = config.get("save_steps", 0)
trainer._train_worker(
dataset,
output_dir=output_dir,
num_epochs=config.get("num_epochs", 3),
learning_rate=lr_value,
batch_size=config.get("batch_size", 2),
gradient_accumulation_steps=config.get("gradient_accumulation_steps", 4),
warmup_steps=config.get("warmup_steps"),
warmup_ratio=config.get("warmup_ratio"),
max_steps=max_steps if max_steps and max_steps > 0 else 0,
save_steps=save_steps if save_steps and save_steps > 0 else 0,
weight_decay=config.get("weight_decay", 0.01),
random_seed=config.get("random_seed", 3407),
packing=config.get("packing", False),
train_on_completions=config.get("train_on_completions", False),
enable_wandb=config.get("enable_wandb", False),
wandb_project=config.get("wandb_project", "unsloth-training"),
wandb_token=config.get("wandb_token"),
enable_tensorboard=config.get("enable_tensorboard", False),
tensorboard_dir=config.get("tensorboard_dir", "runs"),
eval_dataset=eval_dataset,
eval_steps=eval_steps,
max_seq_length=config.get("max_seq_length", 2048),
optim=config.get("optim", "adamw_8bit"),
lr_scheduler_type=config.get("lr_scheduler_type", "linear"),
)
# Check final state
progress = trainer.get_training_progress()
if progress.error:
event_queue.put({
"type": "error",
"error": progress.error,
"stack": "",
"ts": time.time(),
})
else:
event_queue.put({
"type": "complete",
"output_dir": output_dir,
"status_message": progress.status_message or "Training completed",
"ts": time.time(),
})
except Exception as exc:
event_queue.put({
"type": "error",
"error": str(exc),
"stack": traceback.format_exc(limit=20),
"ts": time.time(),
})
def _send_status(event_queue: Any, message: str) -> None:
"""Send a status update to the parent process."""
event_queue.put({
"type": "status",
"message": message,
"ts": time.time(),
})
def _run_embedding_training(event_queue: Any, stop_queue: Any, config: dict) -> None:
"""Self-contained embedding model training pipeline.
Uses FastSentenceTransformer + SentenceTransformerTrainer +
MultipleNegativesRankingLoss — completely separate from the
LLM/VLM/audio paths in UnslothTrainer.
Mirrors the pattern from the reference embedding notebooks:
All_MiniLM_L6_v2.py, BGE_M3.py, EmbeddingGemma_300M.py,
ModernBert.py, Qwen3_Embedding_0_6B.py
"""
import math
import queue as _queue
import threading
model_name = config["model_name"]
training_start_time = time.time()
# ── 1. Import embedding-specific libraries ──
_send_status(event_queue, "Importing embedding libraries...")
try:
from unsloth import FastSentenceTransformer, is_bfloat16_supported
from sentence_transformers import (
SentenceTransformerTrainer,
SentenceTransformerTrainingArguments,
)
from sentence_transformers.losses import MultipleNegativesRankingLoss
from sentence_transformers.training_args import BatchSamplers
from datasets import load_dataset, Dataset
from transformers import TrainerCallback
except ImportError as e:
event_queue.put({
"type": "error",
"error": f"Failed to import embedding libraries: {e}. "
"Ensure 'sentence_transformers' and 'unsloth' are installed.",
"stack": traceback.format_exc(limit=20),
"ts": time.time(),
})
return
# ── Stop signal handling ──
_should_stop = False
_save_on_stop = True
def _poll_stop():
nonlocal _should_stop, _save_on_stop
while True:
try:
msg = stop_queue.get(timeout=1.0)
if msg and msg.get("type") == "stop":
_save_on_stop = msg.get("save", True)
_should_stop = True
logger.info("Embedding training: stop signal received (save=%s)", _save_on_stop)
return
except _queue.Empty:
continue
except (EOFError, OSError):
return
stop_thread = threading.Thread(target=_poll_stop, daemon=True)
stop_thread.start()
# ── 2. Load model ──
_send_status(event_queue, "Loading embedding model...")
try:
hf_token = config.get("hf_token", "")
hf_token = hf_token if hf_token and hf_token.strip() else None
max_seq_length = config.get("max_seq_length", 512)
training_type = config.get("training_type", "LoRA/QLoRA")
use_lora = (training_type == "LoRA/QLoRA")
model = FastSentenceTransformer.from_pretrained(
model_name=model_name,
max_seq_length=max_seq_length,
full_finetuning=not use_lora,
token=hf_token,
)
except Exception as e:
event_queue.put({
"type": "error",
"error": f"Failed to load embedding model '{model_name}': {e}",
"stack": traceback.format_exc(limit=20),
"ts": time.time(),
})
return
if _should_stop:
event_queue.put({"type": "complete", "output_dir": None, "ts": time.time()})
return
# ── 3. Apply LoRA ──
if use_lora:
_send_status(event_queue, "Configuring LoRA adapters (FEATURE_EXTRACTION)...")
try:
gradient_checkpointing = config.get("gradient_checkpointing", False)
# Normalize: "none" or empty → False
if gradient_checkpointing in ("none", "", None):
gradient_checkpointing = False
model = FastSentenceTransformer.get_peft_model(
model,
r=config.get("lora_r", 32),
target_modules=config.get("target_modules") or ["q_proj", "k_proj", "v_proj", "o_proj"],
lora_alpha=config.get("lora_alpha", 64),
lora_dropout=config.get("lora_dropout", 0.0),
bias="none",
use_gradient_checkpointing=gradient_checkpointing,
random_state=config.get("random_seed", 3407),
use_rslora=config.get("use_rslora", False),
loftq_config={"loftq_bits": 4, "loftq_iter": 1} if config.get("use_loftq") else None,
task_type="FEATURE_EXTRACTION",
)
except Exception as e:
event_queue.put({
"type": "error",
"error": f"Failed to configure LoRA for embedding model: {e}",
"stack": traceback.format_exc(limit=20),
"ts": time.time(),
})
return
if _should_stop:
event_queue.put({"type": "complete", "output_dir": None, "ts": time.time()})
return
# ── 4. Load dataset ──
_send_status(event_queue, "Loading dataset...")
try:
hf_dataset = config.get("hf_dataset", "")
local_datasets = config.get("local_datasets") or []
subset = config.get("subset") or None
train_split = config.get("train_split", "train") or "train"
if hf_dataset and hf_dataset.strip():
hf_token = config.get("hf_token", "")
hf_token = hf_token if hf_token and hf_token.strip() else None
dataset = load_dataset(
hf_dataset.strip(),
subset,
split=train_split,
token=hf_token,
)
elif local_datasets:
# Load from local file(s) — mirrors the non-embedding pipeline's
# directory handling so recipe outputs (parquet-files/) work.
all_files: list[str] = []
for dataset_file in local_datasets:
file_path = dataset_file if os.path.isabs(dataset_file) else os.path.join(
project_root, "studio", "backend", "assets", "datasets", dataset_file,
)
if os.path.isdir(file_path):
file_path_obj = Path(file_path)
parquet_dir = (
file_path_obj / "parquet-files"
if (file_path_obj / "parquet-files").exists()
else file_path_obj
)
parquet_files = sorted(parquet_dir.glob("*.parquet"))
if parquet_files:
all_files.extend(str(p) for p in parquet_files)
continue
candidates: list[Path] = []
for ext in (".json", ".jsonl", ".csv", ".parquet"):
candidates.extend(sorted(file_path_obj.glob(f"*{ext}")))
if candidates:
all_files.extend(str(c) for c in candidates)
continue
raise ValueError(f"No supported data files in directory: {file_path_obj}")
else:
all_files.append(file_path)
if all_files:
first_ext = Path(all_files[0]).suffix.lower()
if first_ext in (".json", ".jsonl"):
loader = "json"
elif first_ext == ".csv":
loader = "csv"
elif first_ext == ".parquet":
loader = "parquet"
else:
raise ValueError(f"Unsupported local dataset format: {all_files[0]}")
dataset = load_dataset(loader, data_files=all_files, split="train")
else:
event_queue.put({
"type": "error",
"error": "No dataset specified for embedding training.",
"stack": "", "ts": time.time(),
})
return
# Apply dataset slicing if specified
slice_start = config.get("dataset_slice_start")
slice_end = config.get("dataset_slice_end")
if slice_start is not None or slice_end is not None:
start = slice_start if slice_start is not None else 0
end = slice_end if slice_end is not None else len(dataset)
dataset = dataset.select(range(start, min(end + 1, len(dataset))))
logger.info(f"Embedding dataset loaded: {len(dataset)} samples")
except Exception as e:
event_queue.put({
"type": "error",
"error": f"Failed to load dataset: {e}",
"stack": traceback.format_exc(limit=20),
"ts": time.time(),
})
return
if _should_stop:
event_queue.put({"type": "complete", "output_dir": None, "ts": time.time()})
return
# ── 5. Create loss function ──
loss = MultipleNegativesRankingLoss(model)
# ── 6. Build training arguments ──
_send_status(event_queue, "Configuring training...")
try:
lr_value = float(config.get("learning_rate", "2e-4"))
except ValueError:
event_queue.put({
"type": "error",
"error": f"Invalid learning rate: {config.get('learning_rate')}",
"stack": "", "ts": time.time(),
})
return
output_dir = config.get("output_dir")
if not output_dir:
output_dir = f"./outputs/{model_name.replace('/', '_')}_{int(time.time())}"
num_epochs = config.get("num_epochs", 2)
batch_size = config.get("batch_size", 256)
gradient_accumulation_steps = config.get("gradient_accumulation_steps", 1)
max_steps_val = config.get("max_steps", 0)
save_steps_val = config.get("save_steps", 0)
warmup_ratio = config.get("warmup_ratio", 0.03)
warmup_steps_val = config.get("warmup_steps")
log_frequency = config.get("log_frequency", 50)
# Build args dict
training_args_kwargs = {
"output_dir": output_dir,
"per_device_train_batch_size": batch_size,
"gradient_accumulation_steps": gradient_accumulation_steps,
"learning_rate": lr_value,
"fp16": not is_bfloat16_supported(),
"bf16": is_bfloat16_supported(),
"logging_steps": 1,
"report_to": ["wandb"] if config.get("enable_wandb") else "none",
"lr_scheduler_type": config.get("lr_scheduler_type", "linear"),
"batch_sampler": BatchSamplers.NO_DUPLICATES,
"optim": config.get("optim", "adamw_8bit"),
"weight_decay": config.get("weight_decay", 0.01),
"seed": config.get("random_seed", 3407),
}
# max_steps vs epochs
if max_steps_val and max_steps_val > 0:
training_args_kwargs["max_steps"] = max_steps_val
else:
training_args_kwargs["num_train_epochs"] = num_epochs if num_epochs > 0 else 2
# warmup: prefer warmup_ratio (standard for embedding scripts), fallback to steps
if warmup_ratio is not None and warmup_ratio > 0:
training_args_kwargs["warmup_ratio"] = warmup_ratio
elif warmup_steps_val is not None and warmup_steps_val > 0:
training_args_kwargs["warmup_steps"] = warmup_steps_val
# save_steps
if save_steps_val and save_steps_val > 0:
training_args_kwargs["save_steps"] = save_steps_val
training_args_kwargs["save_strategy"] = "steps"
args = SentenceTransformerTrainingArguments(**training_args_kwargs)
# ── 7. Calculate total steps for progress tracking ──
if max_steps_val and max_steps_val > 0:
total_steps = max_steps_val
else:
effective_epochs = num_epochs if num_epochs > 0 else 2
len_dataloader = math.ceil(len(dataset) / batch_size)
steps_per_epoch = max(len_dataloader // gradient_accumulation_steps, 1)
total_steps = steps_per_epoch * effective_epochs
# ── 8. Create progress callback ──
class _EmbeddingProgressCallback(TrainerCallback):
"""Sends training progress events to the parent process via event_queue."""
def on_log(self, args, state, control, logs=None, **kwargs):
if not logs:
return
loss_value = logs.get("loss", logs.get("train_loss", 0.0))
current_step = state.global_step
elapsed = time.time() - training_start_time
eta = None
if current_step > 0 and total_steps > 0:
remaining = total_steps - current_step
if remaining > 0:
eta = (elapsed / current_step) * remaining
event_queue.put({
"type": "progress",
"step": current_step,
"epoch": round(state.epoch, 2) if state.epoch else 0,
"loss": loss_value,
"learning_rate": logs.get("learning_rate", 0.0),
"total_steps": total_steps,
"elapsed_seconds": elapsed,
"eta_seconds": eta,
"grad_norm": logs.get("grad_norm"),
"num_tokens": getattr(state, "num_input_tokens_seen", None),
"eval_loss": logs.get("eval_loss"),
"status_message": "",
"ts": time.time(),
})
def on_step_end(self, args, state, control, **kwargs):
if _should_stop:
logger.info("Embedding training: stop at step %d", state.global_step)
control.should_training_stop = True
return control
# ── 9. Create trainer and train ──
_send_status(event_queue, "Starting embedding training...")
try:
trainer = SentenceTransformerTrainer(
model=model,
train_dataset=dataset,
loss=loss,
args=args,
callbacks=[_EmbeddingProgressCallback()],
)
trainer.train()
except Exception as e:
event_queue.put({
"type": "error",
"error": f"Embedding training failed: {e}",
"stack": traceback.format_exc(limit=20),
"ts": time.time(),
})
return
# ── 10. Save model ──
if _should_stop and not _save_on_stop:
event_queue.put({
"type": "complete",
"output_dir": None,
"status_message": "Training cancelled",
"ts": time.time(),
})
return
_send_status(event_queue, "Saving model...")
try:
model.save_pretrained(output_dir)
model.tokenizer.save_pretrained(output_dir)
logger.info("Embedding model saved to %s", output_dir)
except Exception as e:
logger.error("Failed to save embedding model: %s", e)
event_queue.put({
"type": "error",
"error": f"Training completed but failed to save: {e}",
"stack": traceback.format_exc(limit=20),
"ts": time.time(),
})
return
# ── 11. Done ──
event_queue.put({
"type": "complete",
"output_dir": output_dir,
"status_message": "Embedding training completed",
"ts": time.time(),
})