import datetime import importlib.metadata import logging import os import shlex import sys import warnings from pathlib import Path from speculators.data_generation.preprocessing import get_tokenizer, load_processor from speculators.provenance import atomic_write, find_repo_root, git_diff, git_sha logger = logging.getLogger("speculators") def resolve_mask_token_id( verifier_name_or_path: str, vocab_size: int, mask_token_id: int | None = None, *, trust_remote_code: bool = False, ) -> int: """Resolve mask_token_id from explicit value, tokenizer, or fallback. Resolution order: 1. Explicit mask_token_id if provided 2. Tokenizer's existing mask_token_id 3. Add <|MASK|> to tokenizer if embed_tokens has unused slots 4. Fallback to pad/eos/unk token """ if mask_token_id is not None: logger.info(f"Using explicit mask_token_id={mask_token_id}") return mask_token_id processor = load_processor( verifier_name_or_path, trust_remote_code=trust_remote_code, ) tokenizer = get_tokenizer(processor) if tokenizer.mask_token_id is not None: logger.info(f"Using tokenizer mask_token_id={tokenizer.mask_token_id}") return tokenizer.mask_token_id if len(tokenizer) < vocab_size: tokenizer.add_special_tokens({"mask_token": "<|MASK|>"}) added_id: int = tokenizer.mask_token_id # type: ignore[assignment] logger.warning( f"Added <|MASK|> to tokenizer, mask_token_id={added_id} " f"(tokenizer len={len(tokenizer)}, vocab_size={vocab_size})" ) return added_id for token_name in ("pad_token_id", "eos_token_id", "unk_token_id"): token_id = getattr(tokenizer, token_name, None) if token_id is not None: warnings.warn( f"Tokenizer does not have mask_token and no unused embedding slots. " f"Using {token_name}={token_id} as fallback.", stacklevel=2, ) return token_id raise ValueError( "Could not resolve mask_token_id: no --mask-token-id provided, tokenizer has " "no mask_token, no unused embedding slots, and no pad/eos/unk fallback tokens." ) def normalize_counted_metrics( metrics: dict[str, float], world_size: int = 1 ) -> dict[str, float]: """Normalize metrics after ReduceOp.SUM across ranks. For any key ending in '_total', finds the matching '_sum' key, computes sum / total, and stores the result under the prefix (e.g. 'loss_sum' / 'loss_total' -> 'loss'). The raw sum/total keys are removed. Any remaining metrics (not part of a sum/total pair) are divided by world_size to compute the average across ranks. """ normalized_keys: set[str] = set() for tk in [k for k in metrics if k.endswith("_total")]: prefix = tk.removesuffix("_total") sk = f"{prefix}_sum" if sk in metrics: total = metrics[tk] metrics[prefix] = metrics[sk] / total if total > 0 else 0.0 del metrics[sk] normalized_keys.add(prefix) del metrics[tk] if world_size > 1: for k in metrics: if k not in normalized_keys: metrics[k] /= world_size return metrics def _save_speculators_patch(save_dir: Path, repo_root: Path | None, sha: str) -> None: if repo_root is None: return try: diff = git_diff(repo_root) content = f"# repo: {repo_root} ({sha})\n{diff}" atomic_write(save_dir / "speculators.patch", content) except OSError: logger.warning("Failed to save speculators.patch", exc_info=True) def save_train_command(save_path: str, argv: list[str] | None = None) -> None: """Write train_command.txt and speculators.patch to *save_path*. ``argv`` is the exact command the run was resolved from (``TrainConfig`` records it during resolution); it falls back to the live ``sys.argv`` when a caller has no recorded argv, so a direct call is unchanged. """ repo_root = find_repo_root(Path(__file__)) sha = git_sha(repo_root) pkg_versions: list[str] = [] for pkg in ("speculators", "vllm", "transformers", "torch", "compressed-tensors"): try: ver = importlib.metadata.version(pkg) except importlib.metadata.PackageNotFoundError: ver = "not installed" pkg_versions.append(f"# {pkg}: {ver}") header = "\n".join( [ f"# Timestamp: {datetime.datetime.now(datetime.timezone.utc).isoformat()}", f"# Git SHA: {sha}", f"# World size: {os.environ.get('WORLD_SIZE', '1')}", *pkg_versions, ] ) command = shlex.join(argv or sys.argv) content = f"{header}\n{command}\n" path = Path(save_path) path.mkdir(parents=True, exist_ok=True) atomic_write(path / "train_command.txt", content) _save_speculators_patch(path, repo_root, sha)