mirror of https://github.com/razor-ai/soup.git
195 lines
7.2 KiB
Python
195 lines
7.2 KiB
Python
"""HF Trainer / TRL SFTTrainer subclass that swaps ``_get_train_sampler``.
|
|
|
|
The schema field ``training.multipack: bool`` shipped in v0.37.0 along
|
|
with the FFD bin-packing :class:`MultipackBatchSampler` and the
|
|
arch-allowlist validator. The live HF Trainer wiring was deliberately
|
|
deferred (mirrors the v0.27.0 MII / v0.38.0 quant-menu / v0.39.0 ReLoRA
|
|
stub-then-live pattern). v0.40.3 (#65) wires it up.
|
|
|
|
Usage from a trainer wrapper::
|
|
|
|
from soup_cli.utils.multipack_trainer import (
|
|
attach_multipack_state, lengths_from_dataset,
|
|
make_multipack_trainer_class,
|
|
)
|
|
if tcfg.multipack:
|
|
validate_multipack_architecture(arch)
|
|
TrainerCls = make_multipack_trainer_class(SFTTrainer)
|
|
trainer = TrainerCls(**trainer_kwargs)
|
|
attach_multipack_state(
|
|
trainer,
|
|
lengths=lengths_from_dataset(train_ds),
|
|
max_seq_len=cfg.data.max_length,
|
|
batch_size=batch_size,
|
|
seed=cfg.training.seed or 0,
|
|
)
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import functools
|
|
import logging
|
|
from typing import Any, Optional, Sequence
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
# Hidden attrs on the trainer instance — picked up by _get_train_sampler.
|
|
_LENGTHS_ATTR = "_soup_multipack_lengths"
|
|
_MAX_SEQ_ATTR = "_soup_multipack_max_seq_len"
|
|
_BATCH_SIZE_ATTR = "_soup_multipack_batch_size"
|
|
_SEED_ATTR = "_soup_multipack_seed"
|
|
|
|
|
|
def lengths_from_dataset(
|
|
dataset: Any,
|
|
*,
|
|
key: str = "input_ids",
|
|
) -> list[int]:
|
|
"""Extract per-sample token lengths from a tokenized dataset.
|
|
|
|
Falls back to a ``"length"`` column when ``input_ids`` is absent. A
|
|
sample with neither contributes ``0`` (which the sampler will reject
|
|
upstream — surfaces the misconfiguration loudly).
|
|
"""
|
|
if dataset is None:
|
|
return []
|
|
lengths: list[int] = []
|
|
try:
|
|
iterator = iter(dataset)
|
|
except TypeError:
|
|
return []
|
|
for row in iterator:
|
|
if not isinstance(row, dict):
|
|
lengths.append(0)
|
|
continue
|
|
value = row.get(key)
|
|
if value is not None:
|
|
try:
|
|
lengths.append(len(value))
|
|
continue
|
|
except TypeError:
|
|
pass
|
|
fallback = row.get("length")
|
|
if isinstance(fallback, int) and not isinstance(fallback, bool):
|
|
lengths.append(fallback)
|
|
else:
|
|
lengths.append(0)
|
|
|
|
# Surface the case where every length is 0 — the multipack sampler
|
|
# would otherwise silently produce phantom batches and the trainer
|
|
# would NaN-loss with no obvious cause. Loud-fail policy mirrors the
|
|
# v0.37.0 multipack-arch allowlist.
|
|
if lengths and not any(lengths):
|
|
logger.warning(
|
|
"lengths_from_dataset(key=%r) returned all zeros — dataset "
|
|
"rows have neither '%s' nor 'length' keys. Multipack sampling "
|
|
"will not work; check tokenisation pipeline.",
|
|
key, key,
|
|
)
|
|
return lengths
|
|
|
|
|
|
def detect_arch_name(model: Any) -> Optional[str]:
|
|
"""Return the model's architecture class name, or ``None``.
|
|
|
|
Probes ``model.config.architectures[0]`` first (HF convention), then
|
|
falls back to ``type(model).__name__``.
|
|
"""
|
|
if model is None:
|
|
return None
|
|
config = getattr(model, "config", None)
|
|
if config is not None:
|
|
archs = getattr(config, "architectures", None)
|
|
if archs:
|
|
try:
|
|
first = archs[0]
|
|
except (IndexError, TypeError):
|
|
first = None
|
|
if isinstance(first, str) and first:
|
|
return first
|
|
cls = type(model).__name__
|
|
return cls or None
|
|
|
|
|
|
@functools.lru_cache(maxsize=None)
|
|
def make_multipack_trainer_class(base_cls: type) -> type:
|
|
"""Return a subclass of ``base_cls`` overriding ``_get_train_sampler``.
|
|
|
|
The override returns a :class:`MultipackBatchSampler` built from
|
|
instance attrs set by :func:`attach_multipack_state`. If state is
|
|
missing the override delegates to the base implementation — so the
|
|
subclass is safe to instantiate even when multipack is later disabled.
|
|
|
|
Cached via ``functools.lru_cache`` so two calls with the same
|
|
``base_cls`` return the SAME subclass — this keeps ``isinstance``
|
|
checks consistent across sweep runs and avoids confusing pickle.
|
|
|
|
.. note::
|
|
v0.40.3 ships this factory but **does not** wire it into the SFT /
|
|
Pretrain trainer wrappers. Adversarial review surfaced that HF
|
|
Trainer's ``_get_train_sampler`` returns a ``Sampler[int]`` which
|
|
the DataLoader then consumes as scalar indices, while
|
|
:class:`MultipackBatchSampler` yields ``list[list[int]]``. Live
|
|
wiring requires a ``get_train_dataloader`` override (with the
|
|
sampler installed as ``batch_sampler=`` on the underlying
|
|
``DataLoader``) and lands in v0.40.4. The factory remains in code
|
|
as the stub end-point used by unit tests.
|
|
"""
|
|
from soup_cli.utils.multipack_sampler import MultipackBatchSampler
|
|
|
|
class MultipackTrainer(base_cls): # type: ignore[misc, valid-type]
|
|
soup_multipack: bool = True
|
|
|
|
def _get_train_sampler(self, *args: Any, **kwargs: Any) -> Any: # type: ignore[override]
|
|
# *args/**kwargs accept newer HF signature
|
|
# (transformers >=4.41 passes train_dataset as positional kwarg).
|
|
lengths = getattr(self, _LENGTHS_ATTR, None)
|
|
max_seq = getattr(self, _MAX_SEQ_ATTR, None)
|
|
batch_size = getattr(self, _BATCH_SIZE_ATTR, None)
|
|
seed = getattr(self, _SEED_ATTR, 0)
|
|
if not lengths or not max_seq or not batch_size:
|
|
return super()._get_train_sampler(*args, **kwargs)
|
|
return MultipackBatchSampler(
|
|
lengths=list(lengths),
|
|
batch_max_len=int(max_seq),
|
|
batch_size=int(batch_size),
|
|
real_batches=True,
|
|
seed=int(seed),
|
|
drop_last=False,
|
|
)
|
|
|
|
MultipackTrainer.__name__ = f"Multipack{base_cls.__name__}"
|
|
MultipackTrainer.__qualname__ = MultipackTrainer.__name__
|
|
return MultipackTrainer
|
|
|
|
|
|
def attach_multipack_state(
|
|
trainer: Any,
|
|
*,
|
|
lengths: Sequence[int],
|
|
max_seq_len: int,
|
|
batch_size: int,
|
|
seed: int = 0,
|
|
) -> None:
|
|
"""Stash the multipack sampler config on ``trainer``."""
|
|
if isinstance(max_seq_len, bool) or not isinstance(max_seq_len, int):
|
|
raise TypeError("max_seq_len must be int")
|
|
if isinstance(batch_size, bool) or not isinstance(batch_size, int):
|
|
raise TypeError("batch_size must be int")
|
|
if isinstance(seed, bool) or not isinstance(seed, int):
|
|
raise TypeError("seed must be int")
|
|
if max_seq_len <= 0:
|
|
raise ValueError(f"max_seq_len must be > 0, got {max_seq_len}")
|
|
if batch_size <= 0:
|
|
raise ValueError(f"batch_size must be > 0, got {batch_size}")
|
|
lengths_list = list(lengths)
|
|
if not lengths_list:
|
|
raise ValueError(
|
|
"lengths must not be empty — empty dataset cannot drive "
|
|
"MultipackBatchSampler. Check the tokenisation pipeline."
|
|
)
|
|
setattr(trainer, _LENGTHS_ATTR, lengths_list)
|
|
setattr(trainer, _MAX_SEQ_ATTR, max_seq_len)
|
|
setattr(trainer, _BATCH_SIZE_ATTR, batch_size)
|
|
setattr(trainer, _SEED_ATTR, seed)
|