uv-scripts/classification
Classification Scripts Text classification on HF Jobs: label a dataset with a model that needs no training, or train your own classifier from labelled examples. If you have seen Jev and other "System One" models: the models these scripts train are small, open versions of the same idea. They read a piece of data and return a label with a probability, and you can train one on your own labels. For example, this demo suggests task tags for any Hub dataset; its model was fine-tuned… See the full description on the dataset page: https://huggingface.co/datasets/uv-scripts/classification.
18316
1# /// script2# requires-python = ">=3.11"3# dependencies = [4# "datasets>=4.0.0",5# "transformers>=5.12",6# "torch",7# "accelerate",8# "safetensors",9# "scikit-learn",10# "numpy",11# "huggingface-hub",12# ]13# ///14"""15Fine-tune a text-classification encoder on a Hub dataset and push the trained model to the Hub.16 17Defaults to LiquidAI's LFM2.5-Encoder-350M — a bidirectional encoder converted from an LFM218decoder backbone (blog: https://huggingface.co/blog/LiquidAI/lfm2-5-encoders). The 230M variant19beats ModernBERT-base on GLUE/SuperGLUE and both handle 8,192-token documents, so long inputs20(dataset cards, legal documents, support threads) fit without chunking. Any Hub encoder works21via --model: models with a standard sequence-classification head (BERT, ModernBERT, DeBERTa, …)22train through `AutoModelForSequenceClassification` and produce standard artifacts; models23without one (like the LFM2.5 encoders) get a generic mean-pooling + linear head that is pushed24as custom code, so the output still round-trips through25`AutoModelForSequenceClassification.from_pretrained(..., trust_remote_code=True)`.26 27Single-label vs multi-label is auto-detected from the label column:28 29- `ClassLabel` / string / int column -> single-label (cross-entropy)30- `Sequence(ClassLabel)` / list of strings -> multi-label (BCE + per-label threshold tuning)31 32Run on HF Jobs (l4x1 is enough for 512-token contexts; the model is downloaded, trained,33evaluated, pushed, and reload-verified in one job):34 35 hf jobs uv run --flavor l4x1 --secrets HF_TOKEN \\36 https://huggingface.co/datasets/uv-scripts/classification/raw/main/train-classifier.py \\37 fancyzhx/ag_news username/my-news-classifier \\38 --max-samples 2000 --epochs 139 40Multi-label example (go_emotions has a Sequence(ClassLabel) `labels` column):41 42 hf jobs uv run --flavor l4x1 --secrets HF_TOKEN \\43 https://huggingface.co/datasets/uv-scripts/classification/raw/main/train-classifier.py \\44 google-research-datasets/go_emotions username/my-emotion-classifier \\45 --label-column labels46 47Long documents: pair --max-length 8192 with --gradient-checkpointing and a small batch size48(--batch-size 2 --grad-accum 8) on a10g/a100 flavors.49 50Model: https://huggingface.co/LiquidAI/LFM2.5-Encoder-350M51 52Smoke-tested 2026-07-28 on a10g-small (transformers 5.14.1, torch 2.13.0): single-label53(ag_news, acc 0.757 on a 2k/1-epoch smoke), multi-label (go_emotions, threshold tuning54lifting micro-F1 0.00->0.24 on a 2k/1-epoch smoke), and the standard-architecture path55(ModernBERT-base on ag_news, acc 0.871, vanilla artifact); pushed models pass the in-job56reload check and a fresh local CPU reload.57"""58 59import argparse60import importlib.util61import json62import logging63import os64import shutil65import sys66import tempfile67from datetime import datetime, timezone68from typing import Optional69 70import numpy as np71import torch72from datasets import ClassLabel, Dataset, load_dataset73from huggingface_hub import HfApi, ModelCard, hf_hub_download, list_repo_files, login74from sklearn.metrics import accuracy_score, f1_score75from transformers import (76 AutoConfig,77 AutoModel,78 AutoModelForSequenceClassification,79 AutoTokenizer,80 DataCollatorWithPadding,81 Trainer,82 TrainingArguments,83)84 85logging.basicConfig(level=logging.INFO)86logger = logging.getLogger(__name__)87 88DEFAULT_MODEL = "LiquidAI/LFM2.5-Encoder-350M"89SCRIPT_URL = "https://huggingface.co/datasets/uv-scripts/classification/raw/main/train-classifier.py"90WRAPPER_MODULE = "modeling_encoder_seq_cls"91WRAPPER_CLASS = "EncoderForSequenceClassification"92 93# Generic sequence-classification wrapper for encoders whose remote code ships no94# AutoModelForSequenceClassification (e.g. the LFM2.5 encoders expose only AutoModel +95# AutoModelForMaskedLM). This exact file is used for training AND copied into the pushed96# repo with an auto_map entry, so the training class and the reload class can never drift.97MODELING_FILE = '''"""Generic sequence classification head: AutoModel backbone + mean pooling + linear.98 99Auto-generated by the uv-scripts `train-classifier.py` recipe. Loaded via100`AutoModelForSequenceClassification.from_pretrained(repo, trust_remote_code=True)`;101the backbone class is resolved from this repo's own `auto_map`/code files.102"""103 104import torch105from torch import nn106from transformers import AutoModel, PreTrainedModel107from transformers.modeling_outputs import SequenceClassifierOutput108 109 110class EncoderForSequenceClassification(PreTrainedModel):111 base_model_prefix = "model"112 supports_gradient_checkpointing = True113 114 def __init__(self, config):115 super().__init__(config)116 self.num_labels = config.num_labels117 self.model = AutoModel.from_config(config, trust_remote_code=True)118 dropout = getattr(config, "classifier_dropout", None)119 self.dropout = nn.Dropout(0.1 if dropout is None else dropout)120 self.classifier = nn.Linear(config.hidden_size, config.num_labels)121 self.post_init()122 123 def forward(self, input_ids=None, attention_mask=None, labels=None, **kwargs):124 outputs = self.model(input_ids=input_ids, attention_mask=attention_mask)125 hidden = outputs.last_hidden_state126 if attention_mask is None:127 pooled = hidden.mean(dim=1)128 else:129 mask = attention_mask.unsqueeze(-1).to(hidden.dtype)130 pooled = (hidden * mask).sum(dim=1) / mask.sum(dim=1).clamp(min=1e-9)131 logits = self.classifier(self.dropout(pooled))132 loss = None133 if labels is not None:134 if self.config.problem_type == "multi_label_classification":135 loss = nn.functional.binary_cross_entropy_with_logits(136 logits, labels.to(logits.dtype)137 )138 else:139 loss = nn.functional.cross_entropy(logits, labels.view(-1))140 return SequenceClassifierOutput(loss=loss, logits=logits)141 142 143# AutoModelForSequenceClassification.from_pretrained registers this class against the144# config class, and that requires config_class to be set (transformers v5 crashes on None).145try:146 __CONFIG_IMPORT__147 EncoderForSequenceClassification.config_class = __CONFIG_CLASS__148except ImportError: # flat import during training; the trainer sets config_class itself149 pass150'''151 152 153def render_modeling_file(config) -> str:154 """Fill the wrapper template with the backbone's concrete config class."""155 config_cls = type(config)156 name = config_cls.__name__157 if config_cls.__module__.startswith("transformers."):158 import_stmt = f"from transformers import {name}"159 else:160 # remote-code config: its module file is copied into the pushed repo alongside161 # this wrapper, where the dynamic-module loader supports relative imports162 module_file = config_cls.__module__.split(".")[-1]163 import_stmt = f"from .{module_file} import {name}"164 return MODELING_FILE.replace("__CONFIG_IMPORT__", import_stmt).replace(165 "__CONFIG_CLASS__", name166 )167 168 169def check_cuda_availability() -> None:170 if not torch.cuda.is_available():171 logger.error("CUDA is not available. This script requires a GPU.")172 logger.error("Run on Hugging Face Jobs with: hf jobs uv run --flavor l4x1 ...")173 sys.exit(1)174 logger.info(f"CUDA is available. GPU: {torch.cuda.get_device_name()}")175 176 177# ---------------------------------------------------------------------------178# Labels179# ---------------------------------------------------------------------------180 181 182def detect_task(dataset: Dataset, label_column: str) -> tuple[str, list[str]]:183 """Return (problem_type, label_names) from the label column's feature/values.184 185 single_label_classification: ClassLabel, string, or int column.186 multi_label_classification: Sequence(ClassLabel)/List(ClassLabel) or list-of-strings column.187 """188 feature = dataset.features[label_column]189 190 # Sequence / List / LargeList all expose .feature; ClassLabel and Value do not.191 inner = getattr(feature, "feature", None)192 193 if inner is not None:194 if isinstance(inner, ClassLabel):195 return "multi_label_classification", list(inner.names)196 values = {v for row in dataset[label_column] for v in (row or [])}197 if not values:198 logger.error(f"Label column '{label_column}' contains only empty lists.")199 sys.exit(1)200 return "multi_label_classification", sorted(str(v) for v in values)201 202 if isinstance(feature, ClassLabel):203 return "single_label_classification", list(feature.names)204 205 values = dataset.unique(label_column)206 if any(v is None for v in values):207 logger.error(f"Label column '{label_column}' contains nulls.")208 sys.exit(1)209 if all(isinstance(v, (int, np.integer)) for v in values):210 return "single_label_classification", [str(v) for v in sorted(values)]211 if all(isinstance(v, str) for v in values):212 return "single_label_classification", sorted(values)213 214 logger.error(215 f"Unsupported label column '{label_column}' "216 f"(feature: {feature}). Supported: ClassLabel, string, int, "217 f"Sequence(ClassLabel), or list-of-strings."218 )219 sys.exit(1)220 221 222def encode_labels(example, label_column, problem_type, label2id, num_labels, ints_are_indices):223 """ints_are_indices: True for ClassLabel columns, where raw ints already ARE the224 class indices. Plain int columns (e.g. values [10, 20]) map via label2id instead."""225 raw = example[label_column]226 if problem_type == "multi_label_classification":227 vec = [0.0] * num_labels228 for v in raw or []:229 if isinstance(v, str):230 idx = label2id[v]231 elif ints_are_indices:232 idx = int(v)233 else:234 idx = label2id[str(v)]235 vec[idx] = 1.0236 return {"encoded_labels": vec}237 if isinstance(raw, str):238 return {"encoded_labels": label2id[raw]}239 if ints_are_indices:240 return {"encoded_labels": int(raw)}241 return {"encoded_labels": label2id[str(raw)]}242 243 244# ---------------------------------------------------------------------------245# Model construction — ordered decision rule (order matters):246# 1. auto_map has AutoModelForSequenceClassification -> custom model ships its own head247# 2. auto_map exists without one (LFM2.5 encoders) -> our mean-pooling wrapper; never248# fall through to the built-in mapping: a future *causal* Lfm2ForSequenceClassification249# in transformers would silently load a causal-mask head onto bidirectional weights250# 3. vanilla model -> standard AutoModelForSequenceClassification (standard artifact,251# servable by vllm/classify-dataset.py)252# ---------------------------------------------------------------------------253 254 255def build_model(model_id, problem_type, label_names, work_dir):256 """Return (model, tokenizer, path) where path is 'custom-shipped'|'custom-wrapper'|'standard'."""257 num_labels = len(label_names)258 id2label = {i: name for i, name in enumerate(label_names)}259 label2id = {name: i for i, name in enumerate(label_names)}260 261 tokenizer = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True)262 config = AutoConfig.from_pretrained(model_id, trust_remote_code=True)263 auto_map = getattr(config, "auto_map", None) or {}264 265 label_kwargs = dict(266 num_labels=num_labels,267 id2label=id2label,268 label2id=label2id,269 problem_type=problem_type,270 )271 272 if "AutoModelForSequenceClassification" in auto_map:273 logger.info("Model ships its own sequence-classification head (auto_map) — using it.")274 model = AutoModelForSequenceClassification.from_pretrained(275 model_id, trust_remote_code=True, **label_kwargs276 )277 return model, tokenizer, "custom-shipped"278 279 if auto_map:280 logger.info(281 "Custom-code model without a sequence-classification head — "282 "using the generic mean-pooling wrapper."283 )284 for key, value in label_kwargs.items():285 setattr(config, key, value)286 wrapper_path = os.path.join(work_dir, f"{WRAPPER_MODULE}.py")287 with open(wrapper_path, "w") as f:288 f.write(render_modeling_file(config))289 spec = importlib.util.spec_from_file_location(WRAPPER_MODULE, wrapper_path)290 module = importlib.util.module_from_spec(spec)291 sys.modules[WRAPPER_MODULE] = module292 spec.loader.exec_module(module)293 wrapper_cls = getattr(module, WRAPPER_CLASS)294 wrapper_cls.config_class = type(config)295 model = wrapper_cls(config)296 # Replace the randomly-initialised backbone with the pretrained weights.297 model.model = AutoModel.from_pretrained(model_id, trust_remote_code=True)298 return model, tokenizer, "custom-wrapper"299 300 logger.info("Standard architecture — using AutoModelForSequenceClassification.")301 model = AutoModelForSequenceClassification.from_pretrained(model_id, **label_kwargs)302 return model, tokenizer, "standard"303 304 305# ---------------------------------------------------------------------------306# Metrics307# ---------------------------------------------------------------------------308 309 310def make_compute_metrics(problem_type):311 def compute(eval_pred):312 logits, labels = eval_pred.predictions, eval_pred.label_ids313 if problem_type == "multi_label_classification":314 probs = 1 / (1 + np.exp(-logits))315 preds = (probs >= 0.5).astype(int)316 return {317 "f1_micro": f1_score(labels, preds, average="micro", zero_division=0),318 "f1_macro": f1_score(labels, preds, average="macro", zero_division=0),319 }320 preds = logits.argmax(axis=-1)321 return {322 "accuracy": accuracy_score(labels, preds),323 "f1_macro": f1_score(labels, preds, average="macro", zero_division=0),324 }325 326 return compute327 328 329def tune_thresholds(logits: np.ndarray, labels: np.ndarray) -> list[float]:330 """Per-label threshold sweep (0.05–0.95) maximising per-label F1 on the eval set."""331 probs = 1 / (1 + np.exp(-logits))332 thresholds = []333 for i in range(labels.shape[1]):334 best_t, best_f1 = 0.5, -1.0335 for t in np.arange(0.05, 0.96, 0.05):336 f1 = f1_score(labels[:, i], (probs[:, i] >= t).astype(int), zero_division=0)337 if f1 > best_f1:338 best_t, best_f1 = round(float(t), 2), f1339 thresholds.append(best_t)340 return thresholds341 342 343# ---------------------------------------------------------------------------344# Push + verify345# ---------------------------------------------------------------------------346 347 348def assemble_output_repo(model, tokenizer, path_kind, model_id, out_dir, extra_config):349 """Fill out_dir with a self-contained, from_pretrained-able model."""350 from safetensors.torch import save_model351 352 tokenizer.save_pretrained(out_dir)353 354 if path_kind != "custom-wrapper":355 # Standard / custom-shipped heads: transformers handles the layout natively356 # (custom_object_save copies remote modules for custom-shipped models).357 for key, value in extra_config.items():358 setattr(model.config, key, value)359 model.save_pretrained(out_dir)360 return361 362 # Custom wrapper: copy the backbone's code files so the pushed repo is self-sufficient,363 # then write config + weights manually (save_pretrained on a dynamically-imported class364 # would try to copy this whole uv script as the modeling file).365 for fname in list_repo_files(model_id):366 if fname.endswith(".py"):367 local = hf_hub_download(model_id, fname)368 shutil.copy(local, os.path.join(out_dir, os.path.basename(fname)))369 logger.info(f"Copied backbone code file: {fname}")370 371 config = model.config372 for key, value in extra_config.items():373 setattr(config, key, value)374 backbone_auto_map = getattr(config, "auto_map", None) or {}375 config.auto_map = {376 **backbone_auto_map,377 "AutoModelForSequenceClassification": f"{WRAPPER_MODULE}.{WRAPPER_CLASS}",378 }379 config.architectures = [WRAPPER_CLASS]380 config.save_pretrained(out_dir)381 382 # Belt and braces: force plain module.Class refs in the saved JSON (transformers can383 # rewrite auto_map entries to 'origin-repo--module.Class', which would point reloads384 # at the origin repo instead of the pushed one).385 config_path = os.path.join(out_dir, "config.json")386 with open(config_path) as f:387 saved = json.load(f)388 saved["auto_map"] = {389 k: v.split("--", 1)[-1] for k, v in saved.get("auto_map", {}).items()390 }391 saved["auto_map"]["AutoModelForSequenceClassification"] = (392 f"{WRAPPER_MODULE}.{WRAPPER_CLASS}"393 )394 with open(config_path, "w") as f:395 json.dump(saved, f, indent=2, sort_keys=True)396 397 save_model(model, os.path.join(out_dir, "model.safetensors"))398 399 400def verify_reload(output_repo, eval_texts, reference_preds, problem_type, max_length, hf_token):401 """Reload the *pushed* repo fresh and check prediction agreement. Hard-fail on mismatch."""402 logger.info(f"RELOAD CHECK: loading {output_repo} back from the Hub...")403 tokenizer = AutoTokenizer.from_pretrained(output_repo, trust_remote_code=True, token=hf_token)404 model = AutoModelForSequenceClassification.from_pretrained(405 output_repo, trust_remote_code=True, token=hf_token406 )407 model.eval()408 enc = tokenizer(409 eval_texts, truncation=True, max_length=max_length, padding=True, return_tensors="pt"410 )411 with torch.no_grad():412 logits = model(**enc).logits413 preds = logits.argmax(dim=-1).tolist()414 if preds != reference_preds:415 logger.error("RELOAD CHECK: FAILED — pushed model disagrees with trained model.")416 logger.error(f" in-memory: {reference_preds}")417 logger.error(f" reloaded: {preds}")418 sys.exit(1)419 logger.info(f"RELOAD CHECK: OK ({len(preds)}/{len(preds)} predictions agree)")420 421 422# ---------------------------------------------------------------------------423# Card424# ---------------------------------------------------------------------------425 426 427def build_card(428 input_dataset, output_repo, model_id, problem_type, label_names, metrics,429 thresholds, path_kind, args_summary,430) -> str:431 on_jobs = os.environ.get("JOB_ID") is not None # set by HF Jobs in-container432 hw = os.environ.get("ACCELERATOR") or "" # e.g. "l4x1"; empty on CPU433 origin = (434 "Produced on [Hugging Face Jobs](https://huggingface.co/docs/huggingface_hub/guides/jobs)"435 + (f" (`{hw}`)" if hw else "")436 ) if on_jobs else "Generated"437 438 tags = ["uv-script", "text-classification"]439 if on_jobs:440 tags.append("hf-jobs")441 tag_lines = "\n".join(f"- {t}" for t in tags)442 443 metric_rows = "\n".join(f"| {k} | {v:.4f} |" for k, v in metrics.items())444 multi = problem_type == "multi_label_classification"445 446 label_list = ", ".join(f"`{name}`" for name in label_names[:30])447 if len(label_names) > 30:448 label_list += f", … ({len(label_names)} total)"449 450 if multi:451 snippet = f"""```python452import torch453from transformers import AutoModelForSequenceClassification, AutoTokenizer454 455model = AutoModelForSequenceClassification.from_pretrained("{output_repo}", trust_remote_code=True)456tokenizer = AutoTokenizer.from_pretrained("{output_repo}", trust_remote_code=True)457 458inputs = tokenizer("your text here", return_tensors="pt", truncation=True)459probs = torch.sigmoid(model(**inputs).logits)[0]460thresholds = torch.tensor(model.config.classifier_thresholds) # tuned on validation461labels = [model.config.id2label[i] for i in (probs >= thresholds).nonzero().flatten().tolist()]462print(labels)463```"""464 else:465 snippet = f"""```python466from transformers import AutoModelForSequenceClassification, AutoTokenizer467 468model = AutoModelForSequenceClassification.from_pretrained("{output_repo}", trust_remote_code=True)469tokenizer = AutoTokenizer.from_pretrained("{output_repo}", trust_remote_code=True)470 471inputs = tokenizer("your text here", return_tensors="pt", truncation=True)472print(model.config.id2label[model(**inputs).logits.argmax().item()])473```"""474 475 serving_note = ""476 if path_kind == "custom-wrapper":477 serving_note = (478 "\n> [!NOTE]\n"479 "> This model uses a custom classification head (mean pooling over a backbone "480 "without a native sequence-classification class), so loading requires "481 "`trust_remote_code=True`. vLLM serving requires a standard architecture.\n"482 )483 484 return f"""---485tags:486{tag_lines}487base_model: {model_id}488datasets:489- {input_dataset}490pipeline_tag: text-classification491library_name: transformers492---493 494# {output_repo.split("/")[-1]}495 496[{model_id}](https://huggingface.co/{model_id}) fine-tuned for497{"multi-label" if multi else "single-label"} text classification on498[{input_dataset}](https://huggingface.co/datasets/{input_dataset}).499 500- **Labels ({len(label_names)})**: {label_list}501- **Date**: {datetime.now(timezone.utc).strftime("%Y-%m-%d %H:%M UTC")}502{serving_note}503## Evaluation504 505| Metric | Value |506|--------|-------|507{metric_rows}508{'''509Per-label decision thresholds tuned on the eval split are stored in510`config.classifier_thresholds`.511 512**Choosing an operating point**: the stored thresholds maximise per-label F1. For513precision-first use (e.g. auto-applying labels), act only on predictions well above514their threshold — sigmoid probabilities are a usable confidence signal, and filtering515to high-confidence predictions trades coverage for precision. Route the rest to review.516''' if multi and thresholds else ""}517## Usage518 519{snippet}520 521## Reproduction522 523{origin} with the [`train-classifier.py`]({SCRIPT_URL}) recipe from [uv-scripts](https://huggingface.co/uv-scripts). Run it yourself:524 525```bash526hf jobs uv run --flavor {hw or "l4x1"} --secrets HF_TOKEN \\527 {SCRIPT_URL} \\528 {args_summary}529```530"""531 532 533# ---------------------------------------------------------------------------534# Main535# ---------------------------------------------------------------------------536 537 538def main(539 input_dataset: str,540 output_repo: str,541 model_id: str = DEFAULT_MODEL,542 dataset_config: Optional[str] = None,543 text_column: str = "text",544 label_column: str = "label",545 train_split: str = "train",546 eval_split: Optional[str] = None,547 eval_fraction: float = 0.1,548 max_samples: Optional[int] = None,549 seed: int = 42,550 max_length: int = 512,551 epochs: int = 3,552 lr: float = 2e-5,553 batch_size: int = 16,554 grad_accum: int = 1,555 warmup_ratio: float = 0.05,556 gradient_checkpointing: bool = False,557 no_bf16: bool = False,558 private: bool = False,559 hf_token: Optional[str] = None,560) -> None:561 import transformers562 563 logger.info(f"transformers {transformers.__version__} | torch {torch.__version__}")564 check_cuda_availability()565 566 HF_TOKEN = hf_token or os.environ.get("HF_TOKEN")567 if HF_TOKEN:568 login(token=HF_TOKEN)569 570 # ----- data -----571 logger.info(f"Loading dataset: {input_dataset} (config={dataset_config})")572 ds = load_dataset(input_dataset, dataset_config)573 if train_split not in ds:574 logger.error(f"Split '{train_split}' not found. Available: {list(ds)}")575 sys.exit(1)576 train_ds = ds[train_split]577 578 if eval_split:579 if eval_split not in ds:580 logger.error(f"Split '{eval_split}' not found. Available: {list(ds)}")581 sys.exit(1)582 eval_ds = ds[eval_split]583 elif "validation" in ds:584 eval_ds, eval_split = ds["validation"], "validation"585 elif "test" in ds:586 eval_ds, eval_split = ds["test"], "test"587 else:588 logger.info(f"No eval split found — holding out {eval_fraction:.0%} of train.")589 parts = train_ds.train_test_split(test_size=eval_fraction, seed=seed)590 train_ds, eval_ds, eval_split = parts["train"], parts["test"], "held-out"591 592 if label_column not in train_ds.column_names and label_column == "label" and "labels" in train_ds.column_names:593 logger.info("Column 'label' not found; falling back to 'labels'.")594 label_column = "labels"595 for col in (text_column, label_column):596 if col not in train_ds.column_names:597 logger.error(f"Column '{col}' not found. Columns: {train_ds.column_names}")598 sys.exit(1)599 600 if max_samples:601 train_ds = train_ds.shuffle(seed=seed).select(range(min(max_samples, len(train_ds))))602 eval_ds = eval_ds.shuffle(seed=seed).select(range(min(max_samples, len(eval_ds))))603 604 problem_type, label_names = detect_task(train_ds, label_column)605 num_labels = len(label_names)606 label2id = {name: i for i, name in enumerate(label_names)}607 label_feature = train_ds.features[label_column]608 ints_are_indices = isinstance(label_feature, ClassLabel) or isinstance(609 getattr(label_feature, "feature", None), ClassLabel610 )611 logger.info(f"Task: {problem_type} | {num_labels} labels | "612 f"train={len(train_ds)} eval={len(eval_ds)} ({eval_split})")613 614 # ----- model -----615 work_dir = tempfile.mkdtemp(prefix="train-classifier-")616 out_dir = os.path.join(work_dir, "model")617 os.makedirs(out_dir, exist_ok=True)618 model, tokenizer, path_kind = build_model(model_id, problem_type, label_names, out_dir)619 if gradient_checkpointing:620 model.gradient_checkpointing_enable()621 622 # ----- tokenize -----623 def tokenize(batch):624 return tokenizer(625 [str(t) for t in batch[text_column]], truncation=True, max_length=max_length626 )627 628 keep = {"input_ids", "attention_mask", "labels"}629 630 def prepare(split):631 # Encode into a TEMP column, drop the original, then rename to "labels".632 # Writing straight into the original column name makes datasets cast the633 # encoded values back to the original schema (e.g. multi-hot floats ->634 # list-of-strings -> the collator crashes with "excessive nesting").635 split = split.map(636 lambda ex: encode_labels(637 ex, label_column, problem_type, label2id, num_labels, ints_are_indices638 ),639 remove_columns=[label_column],640 )641 split = split.rename_column("encoded_labels", "labels")642 split = split.map(tokenize, batched=True)643 return split.remove_columns([c for c in split.column_names if c not in keep])644 645 train_tok, eval_tok = prepare(train_ds), prepare(eval_ds)646 647 # ----- train -----648 bf16 = not no_bf16 and torch.cuda.is_bf16_supported()649 if not bf16:650 logger.warning("bf16 unavailable or disabled — training in fp32.")651 # save_strategy stays "no": Trainer checkpointing on the dynamically-imported wrapper652 # would trigger custom_object_save, which copies this whole uv script as modeling code.653 # The final save is manual (assemble_output_repo).654 training_args = TrainingArguments(655 output_dir=os.path.join(work_dir, "trainer"),656 num_train_epochs=epochs,657 learning_rate=lr,658 per_device_train_batch_size=batch_size,659 per_device_eval_batch_size=batch_size * 2,660 gradient_accumulation_steps=grad_accum,661 warmup_ratio=warmup_ratio,662 weight_decay=0.01,663 bf16=bf16,664 eval_strategy="epoch",665 save_strategy="no",666 logging_steps=10,667 seed=seed,668 report_to="none",669 )670 trainer = Trainer(671 model=model,672 args=training_args,673 train_dataset=train_tok,674 eval_dataset=eval_tok,675 data_collator=DataCollatorWithPadding(tokenizer),676 compute_metrics=make_compute_metrics(problem_type),677 )678 trainer.train()679 680 # ----- final eval (+ threshold tuning for multi-label) -----681 predictions = trainer.predict(eval_tok)682 logits, labels = predictions.predictions, predictions.label_ids683 metrics, thresholds = {}, None684 if problem_type == "multi_label_classification":685 probs = 1 / (1 + np.exp(-logits))686 preds_05 = (probs >= 0.5).astype(int)687 thresholds = tune_thresholds(logits, labels)688 preds_tuned = (probs >= np.array(thresholds)).astype(int)689 metrics = {690 "f1_micro @ 0.5": f1_score(labels, preds_05, average="micro", zero_division=0),691 "f1_macro @ 0.5": f1_score(labels, preds_05, average="macro", zero_division=0),692 "f1_micro @ tuned": f1_score(labels, preds_tuned, average="micro", zero_division=0),693 "f1_macro @ tuned": f1_score(labels, preds_tuned, average="macro", zero_division=0),694 }695 else:696 preds = logits.argmax(axis=-1)697 metrics = {698 "accuracy": accuracy_score(labels, preds),699 "f1_macro": f1_score(labels, preds, average="macro", zero_division=0),700 }701 for k, v in metrics.items():702 logger.info(f"eval {k}: {v:.4f}")703 704 # ----- push -----705 extra_config = {"problem_type": problem_type}706 if thresholds:707 extra_config["classifier_thresholds"] = thresholds708 709 logger.info(f"Assembling output repo in {out_dir}")710 model = model.to("cpu").float()711 assemble_output_repo(model, tokenizer, path_kind, model_id, out_dir, extra_config)712 713 api = HfApi(token=HF_TOKEN)714 api.create_repo(output_repo, repo_type="model", private=private, exist_ok=True)715 logger.info(f"Uploading to {output_repo}")716 api.upload_folder(folder_path=out_dir, repo_id=output_repo, repo_type="model")717 718 args_summary = f"{input_dataset} {output_repo}"719 if model_id != DEFAULT_MODEL:720 args_summary += f" --model {model_id}"721 if label_column != "label":722 args_summary += f" --label-column {label_column}"723 card = build_card(724 input_dataset, output_repo, model_id, problem_type, label_names,725 metrics, thresholds, path_kind, args_summary,726 )727 try:728 ModelCard(card).push_to_hub(output_repo, token=HF_TOKEN)729 except Exception as e:730 logger.warning(f"Could not push model card: {e}")731 732 # ----- verify the pushed artifact round-trips -----733 n_check = min(8, len(eval_ds))734 check_texts = [str(t) for t in eval_ds[text_column][:n_check]]735 model.eval()736 enc = tokenizer(737 check_texts, truncation=True, max_length=max_length, padding=True, return_tensors="pt"738 )739 with torch.no_grad():740 reference_preds = model(**enc).logits.argmax(dim=-1).tolist()741 verify_reload(output_repo, check_texts, reference_preds, problem_type, max_length, HF_TOKEN)742 743 logger.info("Done!")744 logger.info(f"Model: https://huggingface.co/{output_repo}")745 746 747if __name__ == "__main__":748 if len(sys.argv) == 1:749 print("Fine-tune a text-classification encoder (default: LFM2.5-Encoder-350M)")750 print("\nUsage:")751 print(" uv run train-classifier.py INPUT_DATASET OUTPUT_MODEL_REPO [options]")752 print("\nExamples:")753 print(" # single-label (ClassLabel column)")754 print(" uv run train-classifier.py fancyzhx/ag_news username/news-classifier")755 print("\n # multi-label (list-of-labels column)")756 print(" uv run train-classifier.py google-research-datasets/go_emotions \\")757 print(" username/emotion-classifier --label-column labels")758 print("\nFor full help: uv run train-classifier.py --help")759 sys.exit(0)760 761 parser = argparse.ArgumentParser(762 description="Fine-tune a text-classification encoder on a Hub dataset and push to Hub",763 )764 parser.add_argument("input_dataset", help="Input dataset ID")765 parser.add_argument("output_repo", help="Output model repo ID (username/model-name)")766 parser.add_argument("--model", default=DEFAULT_MODEL, help=f"Base model (default: {DEFAULT_MODEL})")767 parser.add_argument("--dataset-config", help="Dataset config name")768 parser.add_argument("--text-column", default="text", help="Text column (default: text)")769 parser.add_argument("--label-column", default="label",770 help="Label column (default: label, falls back to labels)")771 parser.add_argument("--train-split", default="train", help="Train split (default: train)")772 parser.add_argument("--eval-split",773 help="Eval split (default: validation, then test, then a held-out fraction of train)")774 parser.add_argument("--eval-fraction", type=float, default=0.1,775 help="Held-out fraction when no eval split exists (default: 0.1)")776 parser.add_argument("--max-samples", type=int, help="Cap train/eval examples (shuffled first)")777 parser.add_argument("--seed", type=int, default=42, help="Seed (default: 42)")778 parser.add_argument("--max-length", type=int, default=512,779 help="Max sequence length (default: 512; LFM2.5 encoders support 8192)")780 parser.add_argument("--epochs", type=int, default=3, help="Epochs (default: 3)")781 parser.add_argument("--lr", type=float, default=2e-5, help="Learning rate (default: 2e-5)")782 parser.add_argument("--batch-size", type=int, default=16, help="Batch size (default: 16)")783 parser.add_argument("--grad-accum", type=int, default=1, help="Gradient accumulation (default: 1)")784 parser.add_argument("--warmup-ratio", type=float, default=0.05, help="Warmup ratio (default: 0.05)")785 parser.add_argument("--gradient-checkpointing", action="store_true",786 help="Enable gradient checkpointing (for long contexts)")787 parser.add_argument("--no-bf16", action="store_true", help="Disable bf16 (train in fp32)")788 parser.add_argument("--private", action="store_true", help="Make output model repo private")789 parser.add_argument("--hf-token", help="HF token (or set HF_TOKEN)")790 args = parser.parse_args()791 792 main(793 input_dataset=args.input_dataset,794 output_repo=args.output_repo,795 model_id=args.model,796 dataset_config=args.dataset_config,797 text_column=args.text_column,798 label_column=args.label_column,799 train_split=args.train_split,800 eval_split=args.eval_split,801 eval_fraction=args.eval_fraction,802 max_samples=args.max_samples,803 seed=args.seed,804 max_length=args.max_length,805 epochs=args.epochs,806 lr=args.lr,807 batch_size=args.batch_size,808 grad_accum=args.grad_accum,809 warmup_ratio=args.warmup_ratio,810 gradient_checkpointing=args.gradient_checkpointing,811 no_bf16=args.no_bf16,812 private=args.private,813 hf_token=args.hf_token,814 )815 