Download source/src/speculators/train/utils.py from khazic/spec-b300: direct link, hf CLI and curl.
- Browser
- Download file 5.02 kB
-
https://huggingface.co/khazic/spec-b300/resolve/main/source/src/speculators/train/utils.py
- Command line
-
hf download hf://khazic/spec-b300/source/src/speculators/train/utils.py
-
curl -L -o utils.py https://huggingface.co/khazic/spec-b300/resolve/main/source/src/speculators/train/utils.py
5.02 kB
| 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) | |