unsloth/studio/backend/core/training/worker.py
Avaya Aggarwal 0c803242ef
feat(studio): add Continued Pretraining (CPT) as a training method (#4677)
* feat(studio): add Continued Pretraining (CPT) support

Implements CPT as a first-class training method in Unsloth Studio,
resolving feature request #4565.

Changes:
- frontend/src/types/training.ts: add 'cpt' to TrainingMethod union
- frontend/src/lib/vram.ts: add 'cpt' to VramTrainingMethod (fp16 footprint)
- frontend/src/features/export/constants.ts: add CPT to METHOD_LABELS
- frontend/src/features/training/api/mappers.ts: map 'cpt' -> 'Continued Pretraining',
  force packing=true and train_on_completions=false for CPT payloads
- frontend/src/features/studio/sections/model-section.tsx: add 'Continued Pretraining'
  option (purple dot) to Method selector; update tooltip
- frontend/src/features/onboarding/.../model-selection-step.tsx: add CPT to
  onboarding wizard method dropdown
- backend/models/training.py: update training_type field description
- backend/core/training/worker.py: detect is_cpt flag, force packing=True,
  train_on_completions=False, pass is_cpt to _train_worker
- backend/core/training/trainer.py: _train_worker reads is_cpt kwarg, forces
  packing on, skips train_on_responses_only for raw-text pretraining

CPT behaviour:
- Full model weights (no LoRA adapters), same as Full Finetuning
- Sequence packing always enabled for GPU efficiency
- Trains on every token (no chat-format masking)
- VRAM estimated at fp16 (2.0 bytes/param)

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Update mappers.ts

* Add CPT raw dataset support and UI fixes

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Add missing training methods module

* Handle invalid raw-text rows and expose raw in onboarding

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Roland Tannous <115670425+rolandtannous@users.noreply.github.com>
Co-authored-by: Etherll <61019402+Etherll@users.noreply.github.com>
Co-authored-by: Etherll <mrmrmidessam@gmail.com>
2026-05-06 13:38:35 +04:00

1985 lines
73 KiB
Python

# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""
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 structlog
from loggers import get_logger
import math
import os
import shutil
import sys
import time
import traceback
import subprocess as _sp
from pathlib import Path
from typing import Any, Callable
logger = get_logger(__name__)
from utils.hardware import apply_gpu_ids
from utils.wheel_utils import (
direct_wheel_url,
flash_attn_wheel_url,
install_wheel,
probe_torch_wheel_env,
url_exists,
)
def _output_dir_from_resume_checkpoint(
resume_from_checkpoint: str | None,
) -> str | None:
if not resume_from_checkpoint:
return None
path = Path(resume_from_checkpoint)
return str(path.parent if path.name.startswith("checkpoint-") else path)
_CAUSAL_CONV1D_RELEASE_TAG = "v1.6.1.post4"
_CAUSAL_CONV1D_PACKAGE_VERSION = "1.6.1"
_MAMBA_SSM_RELEASE_TAG = "v2.3.1"
_MAMBA_SSM_PACKAGE_VERSION = "2.3.1"
_FLASH_ATTN_RUNTIME_MIN_SEQ_LEN = 32768
_FLASH_ATTN_SKIP_ENV = "UNSLOTH_STUDIO_SKIP_FLASHATTN_INSTALL"
def _model_wants_causal_conv1d(model_name: str) -> bool:
name = model_name.lower()
return any(
key in name
for key in (
"qwen3.5",
"qwen3_5",
"qwen3.6",
"qwen3_6",
"qwen3-next",
"qwen3_next",
"nemotron_h",
"nemotron-h",
"nemotron-3-nano",
"falcon_h1",
"falcon-h1",
"granite-4.0-h",
"granitemoehybrid",
"lfm2",
)
)
def _install_package_wheel_first(
*,
event_queue: Any,
import_name: str,
display_name: str,
pypi_name: str,
pypi_version: str | None = None,
filename_prefix: str | None = None,
release_tag: str | None = None,
release_base_url: str | None = None,
wheel_url_builder: Callable[[dict[str, str] | None], str | None] | None = None,
pypi_spec: str | None = None,
pypi_status_message: str | None = None,
) -> bool:
try:
__import__(import_name)
logger.info("%s already installed", display_name)
return True
except ImportError:
pass
env = probe_torch_wheel_env(timeout = 30)
if wheel_url_builder is not None:
wheel_url = wheel_url_builder(env)
else:
wheel_url = direct_wheel_url(
filename_prefix = filename_prefix,
package_version = pypi_version,
release_tag = release_tag,
release_base_url = release_base_url,
env = env,
)
if wheel_url is None:
logger.info("No compatible %s wheel candidate", display_name)
elif url_exists(wheel_url):
_send_status(event_queue, f"Installing prebuilt {display_name} wheel...")
for installer, result in install_wheel(
wheel_url,
python_executable = sys.executable,
use_uv = bool(shutil.which("uv")),
run = _sp.run,
):
if result.returncode == 0:
logger.info("Installed prebuilt %s wheel successfully", display_name)
return True
logger.warning(
"%s failed to install %s wheel:\n%s",
installer,
display_name,
result.stdout,
)
else:
logger.info("No published %s wheel found: %s", display_name, wheel_url)
is_hip = env and env.get("hip_version")
if is_hip and not shutil.which("hipcc"):
logger.error(
"%s requires hipcc for source compilation on ROCm. "
"Install the ROCm HIP SDK: https://rocm.docs.amd.com",
display_name,
)
_send_status(
event_queue,
f"{display_name}: hipcc not found (ROCm HIP SDK required)",
)
return False
if pypi_spec is None:
pypi_spec = f"{pypi_name}=={pypi_version}"
if pypi_status_message is None:
if is_hip:
pypi_status_message = (
f"Compiling {display_name} from source for ROCm "
"(this may take several minutes)..."
)
else:
pypi_status_message = f"Installing {display_name} from PyPI..."
_send_status(event_queue, pypi_status_message)
# Prefer uv for faster dependency resolution when available
plain_pypi_install = pypi_version is None
if plain_pypi_install:
if shutil.which("uv"):
pypi_cmd = [
"uv",
"pip",
"install",
"--python",
sys.executable,
pypi_spec,
]
else:
pypi_cmd = [sys.executable, "-m", "pip", "install", pypi_spec]
else:
if shutil.which("uv"):
pypi_cmd = [
"uv",
"pip",
"install",
"--python",
sys.executable,
"--no-build-isolation",
"--no-deps",
]
# Avoid stale cache artifacts from partial HIP source builds
if is_hip:
pypi_cmd.append("--no-cache")
pypi_cmd.append(pypi_spec)
else:
pypi_cmd = [
sys.executable,
"-m",
"pip",
"install",
"--no-build-isolation",
"--no-deps",
"--no-cache-dir",
pypi_spec,
]
# Source compilation on ROCm can take 10-30 minutes; use a generous
# timeout. Non-HIP installs preserve the pre-existing "no timeout"
# behaviour so unrelated slow installs (e.g. causal-conv1d source
# build on Linux aarch64 or unsupported torch/CUDA combinations)
# are not aborted at 5 minutes by this PR.
_run_kwargs: dict[str, Any] = {
"stdout": _sp.PIPE,
"stderr": _sp.STDOUT,
"text": True,
}
if is_hip:
_run_kwargs["timeout"] = 1800
try:
result = _sp.run(pypi_cmd, **_run_kwargs)
except _sp.TimeoutExpired:
logger.error(
"%s installation timed out after %ds",
display_name,
_run_kwargs.get("timeout"),
)
_send_status(
event_queue,
f"{display_name} installation timed out after "
f"{_run_kwargs.get('timeout')}s",
)
return False
if result.returncode != 0:
if is_hip:
# Surface a clear error for ROCm source build failures
error_lines = (result.stdout or "").strip().splitlines()
snippet = "\n".join(error_lines[-5:]) if error_lines else "(no output)"
logger.error(
"Failed to compile %s for ROCm:\n%s",
display_name,
result.stdout,
)
_send_status(
event_queue,
f"Failed to compile {display_name} for ROCm. "
"Check that hipcc and ROCm development headers are installed.\n"
f"{snippet}",
)
else:
logger.error(
"Failed to install %s from PyPI:\n%s",
display_name,
result.stdout,
)
return False
if is_hip:
logger.info("Compiled and installed %s from source for ROCm", display_name)
else:
logger.info("Installed %s from PyPI", display_name)
return True
def _ensure_causal_conv1d_fast_path(event_queue: Any, model_name: str) -> None:
if not _model_wants_causal_conv1d(model_name):
return
_install_package_wheel_first(
event_queue = event_queue,
import_name = "causal_conv1d",
display_name = "causal-conv1d",
pypi_name = "causal-conv1d",
pypi_version = _CAUSAL_CONV1D_PACKAGE_VERSION,
filename_prefix = "causal_conv1d",
release_tag = _CAUSAL_CONV1D_RELEASE_TAG,
release_base_url = "https://github.com/Dao-AILab/causal-conv1d/releases/download",
)
_SSM_MODEL_SUBSTRINGS = (
"nemotron_h",
"nemotron-h",
"nemotron-3-nano",
"falcon_h1",
"falcon-h1",
"granite-4.0-h",
"granitemoehybrid",
)
def _ensure_mamba_ssm(event_queue: Any, model_name: str) -> None:
if not any(sub in model_name.lower() for sub in _SSM_MODEL_SUBSTRINGS):
return
logger.info("SSM model detected; setting up mamba-ssm after causal-conv1d")
_install_package_wheel_first(
event_queue = event_queue,
import_name = "mamba_ssm",
display_name = "mamba-ssm",
pypi_name = "mamba-ssm",
pypi_version = _MAMBA_SSM_PACKAGE_VERSION,
filename_prefix = "mamba_ssm",
release_tag = _MAMBA_SSM_RELEASE_TAG,
release_base_url = "https://github.com/state-spaces/mamba/releases/download",
)
def _should_try_runtime_flash_attn_install(max_seq_length: int) -> bool:
if os.getenv(_FLASH_ATTN_SKIP_ENV) == "1":
return False
if max_seq_length < _FLASH_ATTN_RUNTIME_MIN_SEQ_LEN:
return False
return sys.platform.startswith("linux")
def _ensure_flash_attn_for_long_context(event_queue: Any, max_seq_length: int) -> None:
if not _should_try_runtime_flash_attn_install(max_seq_length):
return
installed = _install_package_wheel_first(
event_queue = event_queue,
import_name = "flash_attn",
display_name = "flash-attn",
pypi_name = "flash-attn",
wheel_url_builder = flash_attn_wheel_url,
pypi_spec = "flash-attn",
pypi_status_message = "Installing flash-attn from PyPI for long-context training...",
)
if not installed:
_send_status(event_queue, "Continuing without flash-attn")
def _activate_transformers_version(model_name: str) -> None:
"""Activate the correct transformers version BEFORE any ML imports."""
# Ensure backend is on path for utils imports
backend_path = str(Path(__file__).resolve().parent.parent.parent)
if backend_path not in sys.path:
sys.path.insert(0, backend_path)
from utils.transformers_version import activate_transformers_for_subprocess
activate_transformers_for_subprocess(model_name)
def _adapt_for_mlx_vlm(items):
"""Adapt GPU-path VLM dataset output for mlx-vlm consumption.
The GPU path embeds PIL images inside messages content as
{"type": "image", "image": PIL_Image}. mlx-vlm's prepare_inputs
needs images at top-level to produce pixel_values — regardless of
model type. Extract them and leave bare {"type": "image"} placeholders.
"""
adapted = []
for item in items:
images = []
messages = []
for msg in item.get("messages", []):
content = msg.get("content", "")
if isinstance(content, list):
new_content = []
for part in content:
if isinstance(part, dict) and part.get("type") == "image":
img = part.get("image")
if img is not None:
images.append(img)
new_content.append({"type": "image"})
else:
new_content.append(part)
messages.append({"role": msg["role"], "content": new_content})
else:
messages.append(msg)
out = {"messages": messages}
if images:
out["image"] = images[0] if len(images) == 1 else images
elif "image" in item:
out["image"] = item["image"]
elif "images" in item:
out["images"] = item["images"]
adapted.append(out)
return adapted
_MLX_STUDIO_OPTIM_MAP = {
"adamw_8bit": "adamw",
"paged_adamw_8bit": "adamw",
"adamw_bnb_8bit": "adamw",
"paged_adamw_32bit": "adamw",
"adamw_torch": "adamw",
"adamw_torch_fused": "adamw",
"adamw": "adamw",
"adafactor": "adafactor",
"sgd": "sgd",
"adam": "adam",
"muon": "muon",
"lion": "lion",
}
_MLX_STUDIO_LR_SCHEDULERS = {"linear", "cosine", "constant"}
def _normalize_mlx_studio_optimizer(value):
raw = str(value or "adamw_8bit").strip().lower()
try:
return _MLX_STUDIO_OPTIM_MAP[raw]
except KeyError:
supported = ", ".join(sorted(_MLX_STUDIO_OPTIM_MAP))
raise ValueError(
f"Unsupported optimizer for MLX training: {value!r}. "
f"Supported values: {supported}."
)
def _normalize_mlx_studio_scheduler(value):
raw = str(value or "linear").strip().lower()
if raw not in _MLX_STUDIO_LR_SCHEDULERS:
supported = ", ".join(sorted(_MLX_STUDIO_LR_SCHEDULERS))
raise ValueError(
f"Unsupported LR scheduler for MLX training: {value!r}. "
f"Supported values: {supported}."
)
return raw
def _run_mlx_training(event_queue, stop_queue, config):
"""Self-contained MLX training path for Apple Silicon.
Uses MLXTrainer from unsloth_zoo directly -- no torch/SFTTrainer needed.
Mirrors the event_queue protocol so the parent process pump works unchanged.
"""
import time
import gc
import math
import threading
import queue as _queue
from pathlib import Path
def _send(event_type, **kwargs):
if event_type == "status" and "message" not in kwargs:
sm = kwargs.get("status_message")
if sm is not None:
kwargs["message"] = sm
event_queue.put({"type": event_type, "ts": time.time(), **kwargs})
_send("status", status_message = "Loading MLX libraries...")
import mlx.core as mx
try:
from unsloth_zoo.mlx_loader import FastMLXModel
from unsloth_zoo.mlx_trainer import (
MLXTrainer,
MLXTrainingConfig,
train_on_responses_only,
)
except ImportError as e:
raise ImportError(
"Unsloth: MLX training requires unsloth-zoo with the MLX modules "
"(unsloth_zoo.mlx_loader / unsloth_zoo.mlx_trainer). Reinstall via "
"install.sh on Apple Silicon."
) from e
from datasets import load_dataset
if mx.metal.is_available():
info = mx.device_info()
rec_bytes = info.get("max_recommended_working_set_size", 0) or 0
if rec_bytes > 0:
memory_cap = int(rec_bytes * 0.85)
wired_cap = min(int(rec_bytes), memory_cap)
mx.set_memory_limit(memory_cap)
mx.set_wired_limit(wired_cap)
model_name = config["model_name"]
hf_token = config.get("hf_token") or None
if hf_token:
os.environ["HF_TOKEN"] = hf_token
if config.get("use_loftq"):
message = "LoftQ is not supported for MLX training yet."
_send("error", error = message)
raise NotImplementedError(message)
optim_name = _normalize_mlx_studio_optimizer(config.get("optim", "adamw_8bit"))
lr_scheduler_type = _normalize_mlx_studio_scheduler(
config.get("lr_scheduler_type", "linear")
)
# ── 1. Load model ──
# Force text-only if the dataset is not an image dataset, even if the model
# has vision capabilities (e.g. Qwen3.5-VL trained on plain alpaca text).
_send("status", status_message = f"Loading {model_name}...")
is_dataset_image = bool(config.get("is_dataset_image", False))
training_type = config.get("training_type", "LoRA/QLoRA")
use_lora = training_type == "LoRA/QLoRA"
model, tokenizer = FastMLXModel.from_pretrained(
model_name,
load_in_4bit = config.get("load_in_4bit", True),
full_finetuning = not use_lora,
text_only = None if is_dataset_image else True,
token = hf_token,
trust_remote_code = bool(config.get("trust_remote_code", False)),
random_state = config.get("random_seed", 3407),
)
is_vlm = bool(is_dataset_image and getattr(model, "_is_vlm_model", False))
model._is_vlm_model = is_vlm
# ── 2. Apply LoRA / full FT ──
# Pass gradient_checkpointing as string ("mlx"/"unsloth"/"none"/etc.)
# get_peft_model and MLXTrainer both accept strings and handle them.
gc_setting = config.get("gradient_checkpointing", "mlx")
if isinstance(gc_setting, str):
use_grad_checkpoint = (
gc_setting if gc_setting.lower() not in ("false", "") else False
)
else:
use_grad_checkpoint = gc_setting
if use_lora:
_send("status", status_message = "Configuring LoRA adapters...")
peft_kwargs = dict(
r = config.get("lora_r", 16),
lora_alpha = config.get("lora_alpha", 16),
lora_dropout = config.get("lora_dropout", 0.0),
use_rslora = config.get("use_rslora", False),
init_lora_weights = config.get("init_lora_weights", True),
random_state = config.get("random_seed", 3407),
target_modules = config.get("target_modules")
or [
"q_proj",
"k_proj",
"v_proj",
"o_proj",
"gate_proj",
"up_proj",
"down_proj",
],
use_gradient_checkpointing = use_grad_checkpoint,
)
finetune_language = config.get("finetune_language_layers", True)
finetune_attention = config.get("finetune_attention_modules", True)
finetune_mlp = config.get("finetune_mlp_modules", True)
finetune_vision = (
config.get("finetune_vision_layers", False) if is_vlm else False
)
if (
(finetune_attention or finetune_mlp)
and not finetune_language
and not finetune_vision
):
finetune_language = True
peft_kwargs["finetune_language_layers"] = finetune_language
peft_kwargs["finetune_attention_modules"] = finetune_attention
peft_kwargs["finetune_mlp_modules"] = finetune_mlp
if is_vlm:
peft_kwargs["finetune_vision_layers"] = finetune_vision
model = FastMLXModel.get_peft_model(model, **peft_kwargs)
# ── 3. Load dataset ──
_send("status", status_message = "Loading dataset...")
hf_dataset = config.get("hf_dataset", "")
subset = config.get("subset")
train_split = config.get("train_split", "train") or "train"
eval_split = config.get("eval_split")
slice_start = config.get("dataset_slice_start")
slice_end = config.get("dataset_slice_end")
def _slice(ds):
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(ds) - 1
if end < start:
return ds.select([])
ds = ds.select(range(start, min(end + 1, len(ds))))
return ds
def _load_local(file_paths):
from core.training.trainer import UnslothTrainer
from datasets import load_from_disk
if len(file_paths) == 1:
p = Path(file_paths[0])
if p.is_dir() and (
(p / "dataset_info.json").exists() or (p / "state.json").exists()
):
return load_from_disk(str(p))
all_files = UnslothTrainer._resolve_local_files(file_paths)
if not all_files:
raise ValueError("No local dataset files found")
loader = UnslothTrainer._loader_for_files(all_files)
return load_dataset(loader, data_files = all_files, split = "train")
if hf_dataset:
load_kwargs = {"split": train_split, "token": hf_token}
if subset:
load_kwargs["name"] = subset
dataset = load_dataset(hf_dataset, **load_kwargs)
dataset = _slice(dataset)
elif config.get("local_datasets"):
dataset = _load_local(config["local_datasets"])
dataset = _slice(dataset)
else:
raise ValueError("No dataset specified")
# Eval dataset (separate split or local file)
eval_dataset = None
if eval_split and hf_dataset:
eval_kwargs = {"split": eval_split, "token": hf_token}
if subset:
eval_kwargs["name"] = subset
try:
eval_dataset = load_dataset(hf_dataset, **eval_kwargs)
except Exception as e:
_send("status", status_message = f"Eval split load failed: {e}")
eval_dataset = None
elif config.get("local_eval_datasets"):
eval_dataset = _load_local(config["local_eval_datasets"])
# ── 3b. Format dataset (VLM or text) ──
# Reuse the GPU path's format pipeline for both VLM (auto-detects OCR/caption/
# llava/sharegpt+images) and text (alpaca/sharegpt/chatml → "text" column).
format_type = config.get("format_type", "")
try:
from utils.datasets import format_and_template_dataset
def _fmt_progress(status_message = "", **_kw):
_send("status", status_message = status_message)
if is_vlm:
_send("status", status_message = "Formatting VLM dataset...")
vlm_info = format_and_template_dataset(
dataset,
model_name = model_name,
tokenizer = tokenizer,
is_vlm = True,
dataset_name = hf_dataset or "local",
progress_callback = _fmt_progress,
)
if vlm_info.get("success"):
dataset = _adapt_for_mlx_vlm(vlm_info["dataset"])
else:
errors = vlm_info.get("errors", [])
raise ValueError(
f"VLM dataset format conversion failed: {'; '.join(errors)}"
)
if eval_dataset is not None:
ev_info = format_and_template_dataset(
eval_dataset,
model_name = model_name,
tokenizer = tokenizer,
is_vlm = True,
dataset_name = hf_dataset or "local",
)
if ev_info.get("success"):
eval_dataset = _adapt_for_mlx_vlm(ev_info["dataset"])
elif format_type:
_send("status", status_message = f"Formatting dataset ({format_type})...")
info = format_and_template_dataset(
dataset,
model_name = model_name,
tokenizer = tokenizer,
is_vlm = False,
format_type = format_type,
dataset_name = hf_dataset or "local",
)
if info.get("success", True):
dataset = info.get("dataset", dataset)
if eval_dataset is not None:
ev = format_and_template_dataset(
eval_dataset,
model_name = model_name,
tokenizer = tokenizer,
is_vlm = False,
format_type = format_type,
dataset_name = hf_dataset or "local",
)
if ev.get("success", True):
eval_dataset = ev.get("dataset", eval_dataset)
except ImportError:
_send("status", status_message = "Format helper unavailable, using raw dataset")
# ── 4. Resolve training steps ──
max_steps = config.get("max_steps", 0) or 0
num_epochs = config.get("num_epochs", 3)
max_seq_length = config.get("max_seq_length", 2048)
batch_size = config.get("batch_size", 4)
grad_accum = config.get("gradient_accumulation_steps", 4)
if max_steps <= 0:
max_steps = max(
1,
math.ceil(len(dataset) / batch_size / grad_accum) * num_epochs,
)
lr_value = float(config.get("learning_rate", "2e-4"))
# Warmup: prefer warmup_steps; fall back to warmup_ratio
warmup_steps = config.get("warmup_steps")
warmup_ratio = config.get("warmup_ratio")
if warmup_steps is None and warmup_ratio is not None:
warmup_steps = int(round(warmup_ratio * max_steps))
if warmup_steps is None:
warmup_steps = 5
# ── 5. Build output dir ──
output_dir = config.get("output_dir", "")
if not output_dir:
output_dir = f"{model_name.replace('/', '_')}_{int(time.time())}"
# Resolve to ~/.unsloth/studio/outputs/ so the export page can find it
from utils.paths import resolve_output_dir, ensure_dir
output_dir = str(resolve_output_dir(output_dir))
ensure_dir(Path(output_dir))
# ── 6. Create trainer ──
eval_steps_val = config.get("eval_steps", 0) or 0
if isinstance(eval_steps_val, float) and 0 < eval_steps_val < 1:
# Studio sometimes sends fraction-of-total-steps
eval_steps_val = max(1, int(eval_steps_val * max_steps))
else:
eval_steps_val = int(eval_steps_val)
trainer = MLXTrainer(
model = model,
tokenizer = tokenizer,
train_dataset = dataset,
eval_dataset = eval_dataset,
args = MLXTrainingConfig(
per_device_train_batch_size = batch_size,
gradient_accumulation_steps = grad_accum,
max_steps = max_steps,
learning_rate = lr_value,
warmup_steps = warmup_steps,
lr_scheduler_type = lr_scheduler_type,
optim = optim_name,
weight_decay = float(config.get("weight_decay", 0.001) or 0.001),
logging_steps = 1,
max_seq_length = max_seq_length,
seed = config.get("random_seed", 3407),
use_cce = True,
compile = True,
gradient_checkpointing = use_grad_checkpoint,
streaming = is_vlm,
packing = bool(config.get("packing", False)),
output_dir = output_dir,
save_steps = int(config.get("save_steps", 0) or 0),
eval_steps = eval_steps_val,
),
)
# Tell the parent that eval is configured so the frontend shows the eval chart
if eval_dataset is not None and eval_steps_val > 0:
_send("eval_configured")
# ── 7. Apply train_on_responses_only if requested ──
if config.get("train_on_completions", False):
_send("status", status_message = "Configuring response-only training...")
try:
from utils.datasets import (
MODEL_TO_TEMPLATE_MAPPER,
TEMPLATE_TO_RESPONSES_MAPPER,
)
template_name = MODEL_TO_TEMPLATE_MAPPER.get(model_name.lower())
markers = (
TEMPLATE_TO_RESPONSES_MAPPER.get(template_name)
if template_name
else None
)
if markers:
trainer = train_on_responses_only(
trainer,
instruction_part = markers["instruction"],
response_part = markers["response"],
)
else:
_send(
"status",
status_message = f"train_on_completions skipped (no template for {model_name})",
)
except Exception as e:
_send("status", status_message = f"train_on_completions failed: {e}")
# ── 8. Setup wandb / tensorboard ──
wandb_run = None
tb_writer = None
if config.get("enable_wandb", False):
try:
import wandb as _wandb
wandb_token = config.get("wandb_token")
if wandb_token:
os.environ["WANDB_API_KEY"] = wandb_token
_wandb_sensitive = {"hf_token", "wandb_token"}
wandb_run = _wandb.init(
project = config.get("wandb_project") or "unsloth-mlx",
config = {k: v for k, v in config.items() if k not in _wandb_sensitive},
reinit = True,
)
except Exception as e:
_send("status", status_message = f"wandb init failed: {e}")
if config.get("enable_tensorboard", False):
try:
from tensorboardX import SummaryWriter
except ImportError:
try:
from torch.utils.tensorboard import SummaryWriter
except ImportError:
SummaryWriter = None
if SummaryWriter is not None:
try:
tb_dir = config.get("tensorboard_dir") or f"{output_dir}/runs"
tb_writer = SummaryWriter(log_dir = tb_dir)
except Exception as e:
_send("status", status_message = f"tensorboard init failed: {e}")
else:
_send(
"status",
status_message = "tensorboard unavailable (install tensorboardX)",
)
# ── 9. Real-time progress callback ──
_send("status", status_message = f"Training {model_name}...")
def _on_step(step, total, loss, lr, tok_s, peak_gb, elapsed, num_tokens):
eta = (elapsed / step * (total - step)) if step > 0 else 0
_send(
"progress",
step = step,
epoch = round(step / total * num_epochs, 2) if total > 0 else 0,
loss = loss,
learning_rate = lr,
total_steps = total,
elapsed_seconds = elapsed,
eta_seconds = max(0, eta),
grad_norm = None,
num_tokens = num_tokens,
eval_loss = None,
status_message = None,
peak_memory_gb = peak_gb,
)
if wandb_run is not None:
try:
wandb_run.log(
{
"train/loss": loss,
"train/learning_rate": lr,
"train/tokens_per_sec": tok_s,
"train/peak_gb": peak_gb,
"train/num_tokens": num_tokens,
},
step = step,
)
except Exception:
pass
if tb_writer is not None:
try:
tb_writer.add_scalar("train/loss", loss, step)
tb_writer.add_scalar("train/learning_rate", lr, step)
tb_writer.add_scalar("train/tokens_per_sec", tok_s, step)
tb_writer.add_scalar("train/peak_gb", peak_gb, step)
except Exception:
pass
trainer.add_step_callback(_on_step)
def _on_eval(step, eval_loss, perplexity):
_send("progress", step = step, eval_loss = eval_loss)
if wandb_run is not None:
try:
wandb_run.log(
{"eval/loss": eval_loss, "eval/perplexity": perplexity}, step = step
)
except Exception:
pass
if tb_writer is not None:
try:
tb_writer.add_scalar("eval/loss", eval_loss, step)
tb_writer.add_scalar("eval/perplexity", perplexity, step)
except Exception:
pass
trainer.add_eval_callback(_on_eval)
# ── 10. Stop signal polling ──
_stop_save = [True] # mutable so thread can update; [save_flag]
def _poll_stop():
while True:
try:
msg = stop_queue.get(timeout = 1.0)
if msg and msg.get("type") == "stop":
_stop_save[0] = msg.get("save", True)
trainer.stop_requested = True
return
except _queue.Empty:
continue
except (EOFError, OSError):
# why safe: pipe permanently broken, no further messages can arrive
return
stop_thread = threading.Thread(target = _poll_stop, daemon = True)
stop_thread.start()
# ── 11. Run training ──
gc.collect()
mx.synchronize()
trainer.train()
# ── 12. Save and finalize ──
if trainer.stop_requested and not _stop_save[0]:
# User clicked "Cancel" (save=False) — skip saving
_send("complete", output_dir = None, status_message = "Training cancelled")
else:
_send("status", status_message = "Saving model...")
mx.synchronize()
trainer.save_model(output_dir)
_send("complete", output_dir = output_dir, status_message = "Training completed")
if tb_writer is not None:
try:
tb_writer.close()
except Exception:
pass
if wandb_run is not None:
try:
wandb_run.finish()
except Exception:
pass
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"
os.environ["PYTHONWARNINGS"] = (
"ignore" # Suppress warnings at C-level before imports
)
import warnings
from loggers.config import LogConfig
if os.getenv("ENVIRONMENT_TYPE", "production") == "production":
warnings.filterwarnings("ignore")
LogConfig.setup_logging(
service_name = "unsloth-studio-training-worker",
env = os.getenv("ENVIRONMENT_TYPE", "production"),
)
apply_gpu_ids(config.get("resolved_gpu_ids"))
model_name = config["model_name"]
# ── 0. MLX FAST-PATH (must run before any torch/transformers imports) ──
# Apple Silicon uses MLXTrainer directly -- skip transformers version
# activation, causal-conv1d install, and torch imports entirely.
backend_path = str(Path(__file__).resolve().parent.parent.parent)
if backend_path not in sys.path:
sys.path.insert(0, backend_path)
from utils.hardware import hardware as _hw
_hw.detect_hardware()
if _hw.DEVICE == _hw.DeviceType.MLX:
if config.get("is_dataset_audio"):
event_queue.put(
{
"type": "error",
"error": "Audio dataset training is not yet supported on Apple Silicon.",
"stack": "",
"ts": time.time(),
}
)
return
# Activate correct transformers version (Gemma-4 needs 5.5.0, etc.)
# Must happen before any transformers/mlx-lm imports in _run_mlx_training.
try:
_activate_transformers_version(model_name)
except Exception:
pass # Non-fatal: fall through with whatever version is installed
try:
_run_mlx_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
# ── 1. Activate correct transformers version BEFORE any ML imports ──
try:
_activate_transformers_version(model_name)
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
# ── 1a. Auto-enable trust_remote_code for NemotronH/Nano models ──
# NemotronH has config parsing bugs in transformers that require
# trust_remote_code=True as a workaround. Other transformers 5.x models
# (Qwen3.5, Gemma 4, etc.) are native and do NOT need it — enabling it
# bypasses the compiler (disabling fused CE).
# NOTE: Must NOT match Llama-Nemotron (standard Llama architecture).
_NEMOTRON_TRUST_SUBSTRINGS = ("nemotron_h", "nemotron-h", "nemotron-3-nano")
_lowered = model_name.lower()
if (
any(sub in _lowered for sub in _NEMOTRON_TRUST_SUBSTRINGS)
and (_lowered.startswith("unsloth/") or _lowered.startswith("nvidia/"))
and not config.get("trust_remote_code", False)
):
config["trust_remote_code"] = True
logger.info(
"Auto-enabled trust_remote_code for Nemotron model: %s",
model_name,
)
# ── 1b. Set up causal-conv1d first, then install mamba-ssm if needed ──
try:
_ensure_causal_conv1d_fast_path(event_queue, model_name)
_ensure_mamba_ssm(event_queue, model_name)
_ensure_flash_attn_for_long_context(
event_queue,
int(config.get("max_seq_length", 2048)),
)
except Exception as exc:
event_queue.put(
{
"type": "error",
"error": (
f"Please choose another model to train, since "
f"causal-conv1d / mamba-ssm failed to install "
f"with error: {exc}"
),
"stack": traceback.format_exc(limit = 20),
"ts": time.time(),
}
)
return
# ── 1c. Set fork start method so dataset.map() can multiprocess ──
# The parent launched us via spawn (clean process), but the compiled
# SFTTrainer checks get_start_method() and disables num_proc if not "fork".
# Linux only: fork is the default start method and is safe here (no CUDA
# context exists yet). macOS defaults to spawn since Python 3.8 because
# fork is unsafe with macOS frameworks (Metal/MPS, CoreFoundation) --
# do NOT override on macOS. Windows has no fork at all.
if sys.platform == "linux":
import multiprocessing as _mp
try:
_mp.set_start_method("fork", force = True)
except RuntimeError:
pass # Already set
# ── 1c. 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 Unsloth...")
backend_path = str(Path(__file__).resolve().parent.parent.parent)
if backend_path not in sys.path:
sys.path.insert(0, backend_path)
from core.training.trainer import UnslothTrainer, TrainingProgress
from utils.paths import (
ensure_dir,
resolve_output_dir,
resolve_tensorboard_dir,
datasets_root,
)
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 is not None
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 ──
# Order: detect → dataset → model → prepare → train
# Dataset processing (including LLM-assisted detection) runs BEFORE model
# loading so both never occupy VRAM at the same time.
try:
hf_token = config.get("hf_token", "")
hf_token = hf_token if hf_token and hf_token.strip() else None
# ── 4a. Lightweight detection + tokenizer (no VRAM) ──
_send_status(event_queue, "Detecting model type...")
trainer.pre_detect_and_load_tokenizer(
model_name = model_name,
max_seq_length = config["max_seq_length"],
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 trainer.should_stop:
event_queue.put({"type": "complete", "output_dir": None, "ts": time.time()})
return
# ── 4b. Load and format dataset (LLM helper may use VRAM briefly) ──
_send_status(event_queue, "Loading and formatting dataset...")
hf_dataset = config.get("hf_dataset", "")
training_type = config.get("training_type", "LoRA/QLoRA")
_is_cpt_for_dataset = training_type == "Continued Pretraining"
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,
local_eval_datasets = config.get("local_eval_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"),
is_cpt = _is_cpt_for_dataset,
)
if isinstance(dataset_result, tuple):
dataset, eval_dataset = dataset_result
else:
dataset = dataset_result
eval_dataset = None
# [DEBUG] Print first sample before model is loaded
# dataset is a dict {"dataset": <Dataset>, "detected_format": ..., ...}
# or a raw Dataset for audio paths
# try:
# ds = dataset["dataset"] if isinstance(dataset, dict) else dataset
# print(
# f"\n[DEBUG] Dataset loaded BEFORE model. type={type(ds).__name__}, len={len(ds)}",
# flush = True,
# )
# print(f"[DEBUG] Columns: {ds.column_names}", flush = True)
# sample = ds[0]
# preview = {k: str(v)[:300] for k, v in sample.items()}
# print(f"[DEBUG] First sample: {preview}\n", flush = True)
# except Exception as e:
# print(
# f"[DEBUG] Could not preview first sample: {type(e).__name__}: {e}",
# flush = True,
# )
# 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
# ── Start tqdm monitor early so it captures download + tokenization bars ──
import threading as _th
_tqdm_stop = _th.Event()
def _monitor_tqdm():
from tqdm.auto import tqdm as _tqdm_cls
while not _tqdm_stop.is_set():
for bar in list(getattr(_tqdm_cls, "_instances", set())):
try:
n, total = bar.n or 0, bar.total or 0
desc = getattr(bar, "desc", "") or ""
if total > 0 and n > 0 and desc:
pct = min(int(n * 100 / total), 100)
_send_status(
event_queue, f"{desc.strip()} {pct}% ({n:,}/{total:,})"
)
except (AttributeError, ReferenceError):
pass
_tqdm_stop.wait(3)
_tqdm_thread = _th.Thread(target = _monitor_tqdm, daemon = True)
_tqdm_thread.start()
training_type = config.get("training_type", "LoRA/QLoRA")
is_cpt = training_type == "Continued Pretraining"
use_lora = training_type in ("LoRA/QLoRA", "Continued Pretraining")
cpt_trains_embeddings = False
# ── 4c. Load training model (uses VRAM — dataset already formatted) ──
_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"],
full_finetuning = not use_lora,
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),
gpu_ids = config.get("resolved_gpu_ids"),
)
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
# ── 4d. Prepare model (LoRA, full finetuning, or CPT) ──
if is_cpt:
_send_status(event_queue, "Configuring LoRA for continued pretraining...")
# embed_tokens (if the user included it) goes to modules_to_save —
# trained full-precision at embedding_learning_rate. lm_head stays as
# a LoRA target for merge compatibility (see unsloth PR #4106).
_user_modules = config.get("target_modules") or []
wants_embed = "embed_tokens" in _user_modules
cpt_trains_embeddings = wants_embed
cpt_target_modules = [m for m in _user_modules if m != "embed_tokens"]
if not cpt_target_modules:
cpt_target_modules = [
"q_proj",
"k_proj",
"v_proj",
"o_proj",
"gate_proj",
"up_proj",
"down_proj",
"lm_head",
]
success = trainer.prepare_model_for_training(
use_lora = True,
target_modules = cpt_target_modules,
modules_to_save = ["embed_tokens"] if wants_embed else None,
lora_r = config.get("lora_r", 128),
lora_alpha = config.get("lora_alpha", 32),
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),
)
elif 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
lr_default = "5e-5" if is_cpt else "2e-4"
try:
lr_value = float(config.get("learning_rate", lr_default))
except ValueError:
event_queue.put(
{
"type": "error",
"error": f"Invalid learning rate: {config.get('learning_rate')}",
"stack": "",
"ts": time.time(),
}
)
return
# embedding_learning_rate is validated by the Pydantic model (Optional[float],
# gt=0, lt=1.0); if present it is already a finite float in range.
embedding_lr_value = config.get("embedding_learning_rate")
if is_cpt:
if cpt_trains_embeddings:
if embedding_lr_value is None:
# Default embedding_learning_rate = lr/10 per Unsloth's CPT notebook.
embedding_lr_value = lr_value / 10.0
logger.info(
f"CPT: using default embedding_learning_rate={embedding_lr_value:.1e} "
f"(lr/10). Set explicitly to override.\n"
)
elif embedding_lr_value is not None:
logger.warning(
"CPT: embedding_learning_rate was provided but embed_tokens is "
"not being trained; ignoring the override.\n"
)
embedding_lr_value = None
# Generate output dir
resume_from_checkpoint = config.get("resume_from_checkpoint")
output_dir = config.get("output_dir") or _output_dir_from_resume_checkpoint(
resume_from_checkpoint
)
if not output_dir:
output_dir = f"{model_name.replace('/', '_')}_{int(time.time())}"
output_dir = str(resolve_output_dir(output_dir))
ensure_dir(Path(output_dir))
tensorboard_dir = config.get("tensorboard_dir")
if config.get("enable_tensorboard", False):
tensorboard_dir = str(resolve_tensorboard_dir(tensorboard_dir))
ensure_dir(Path(tensorboard_dir))
# Start training (directly — no inner thread, we ARE the subprocess)
dataset_display = (
config.get("hf_dataset", "") or config.get("uploaded_file", "") or ""
)
_send_status(
event_queue,
f'Training "{model_name}"'
+ (f"\nDataset = {dataset_display}" if dataset_display else ""),
)
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,
embedding_learning_rate = embedding_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.001),
random_seed = config.get("random_seed", 3407),
packing = config.get("packing", False),
train_on_completions = False
if is_cpt
else 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 = tensorboard_dir,
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"),
is_cpt = is_cpt,
resume_from_checkpoint = resume_from_checkpoint,
)
_tqdm_stop.set()
# Check final state
progress = trainer.get_training_progress()
if progress.error:
event_queue.put(
{
"type": "error",
"error": progress.error,
"stack": "",
"ts": time.time(),
}
)
else:
saved_output_dir = (
None if trainer.should_stop and not trainer.save_on_stop else output_dir
)
event_queue.put(
{
"type": "complete",
"output_dir": saved_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
from utils.paths import datasets_root, resolve_output_dir
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(
str(datasets_root()),
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
resume_from_checkpoint = config.get("resume_from_checkpoint")
output_dir = config.get("output_dir") or _output_dir_from_resume_checkpoint(
resume_from_checkpoint
)
if not output_dir:
output_dir = str(
resolve_output_dir(f"{model_name.replace('/', '_')}_{int(time.time())}")
)
output_dir = str(resolve_output_dir(output_dir))
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.001),
"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", None))
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", None),
"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(resume_from_checkpoint = resume_from_checkpoint)
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:
if _should_stop and _save_on_stop:
trainer._save_checkpoint(trainer.model, trial = None)
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(),
}
)