khazic's picture
Archive three-epoch run: logs and provenance part 3
e65937c verified
Raw History Blame Contribute Delete
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)