Hanse2-100M-Base / train.py
Evicka's picture
Upload folder using huggingface_hub
0a83496 verified
Raw History Blame Contribute Delete
25.7 kB
# ruff: noqa: E402
"""Restartable three-phase base pretraining for Hanse-LM."""
import argparse
import hashlib
import json
import math
import os
import shutil
from dataclasses import dataclass
from pathlib import Path
# Must be set before importing torch.
os.environ.setdefault("PYTORCH_ALLOC_CONF", "expandable_segments:True")
os.environ.setdefault("CUDA_VISIBLE_DEVICES", "0")
import numpy as np
import torch
from datasets import load_dataset
from torch.utils.data import Dataset
from tqdm import tqdm
from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
LlamaConfig,
LlamaForCausalLM,
PreTrainedTokenizerFast,
Trainer,
TrainingArguments,
set_seed,
)
from transformers.trainer_utils import get_last_checkpoint
SCRIPT_DIR = Path(__file__).resolve().parent
PROJECT_DIR = SCRIPT_DIR.parent
RUN_NAME = "Hanse-100M-Base-v1"
RUN_DIR = PROJECT_DIR / RUN_NAME
FINAL_DIR = PROJECT_DIR / f"{RUN_NAME}-FINAL"
TOKENIZER_FILE = SCRIPT_DIR / "hanse_tokenizer.json"
DATA_ROOT = Path(os.environ.get("HANSE_DATA_DIR", r"C:\hanselm-data"))
VOCAB_SIZE = 32_000
BASE_SEQUENCE_LENGTH = 2048
SEED = 42
BATCH_TEXTS = 256
FLUSH_EVERY = 1_000_000
EVAL_TARGET_TOKENS = 2_000_000
CACHE_FORMAT_VERSION = 1
EVAL_HASH_MODULUS = 10_000
EVAL_HASH_BUCKETS = 100
MODEL_PARAMETERS = 99_144_320
TOKENS_PER_STEP = 131_072
SPECIAL_TOKEN_IDS = {
"pad_token_id": 0,
"bos_token_id": 1,
"eos_token_id": 2,
"unk_token_id": 3,
}
@dataclass(frozen=True)
class Source:
dataset: str
config: str
revision: str
@dataclass(frozen=True)
class Phase:
name: str
requested_tokens: int
sequence_length: int
micro_batch_size: int
gradient_accumulation_steps: int
mix: dict[str, float]
learning_rate: float
warmup_ratio: float
decay_type: str | None = None
min_lr_ratio: float | None = None
@property
def tokens_per_step(self) -> int:
return self.sequence_length * self.micro_batch_size * self.gradient_accumulation_steps
@property
def steps(self) -> int:
return self.requested_tokens // self.tokens_per_step
@property
def actual_tokens(self) -> int:
return self.steps * self.tokens_per_step
@property
def source_budgets(self) -> dict[str, int]:
return allocate_budgets(self.actual_tokens, self.mix)
SOURCES = {
"fineweb_de": Source(
"HuggingFaceFW/fineweb-2",
"deu_Latn",
"af9c13333eb981300149d5ca60a8e9d659b276b9",
),
"fineweb_edu_en": Source(
"HuggingFaceFW/fineweb-edu",
"sample-100BT",
"87f09149ef4734204d70ed1d046ddc9ca3f2b8f9",
),
"finewiki_de": Source(
"HuggingFaceFW/finewiki",
"de",
"8bd13e72e6a002407649b3e898535f42ceb1aeb9",
),
"finewiki_en": Source(
"HuggingFaceFW/finewiki",
"en",
"8bd13e72e6a002407649b3e898535f42ceb1aeb9",
),
}
PHASES = (
Phase(
"phase-1",
15_000_000_000,
2048,
2,
32,
{"fineweb_de": 0.45, "fineweb_edu_en": 0.40, "finewiki_de": 0.10, "finewiki_en": 0.05},
6e-4,
0.01,
),
Phase(
"phase-2",
4_000_000_000,
4096,
1,
32,
{"fineweb_de": 0.40, "fineweb_edu_en": 0.35, "finewiki_de": 0.20, "finewiki_en": 0.05},
3e-4,
0.01,
"1-sqrt",
0.10,
),
Phase(
"phase-3",
1_000_000_000,
8192,
1,
16,
{"fineweb_de": 0.35, "fineweb_edu_en": 0.40, "finewiki_de": 0.20, "finewiki_en": 0.05},
3e-5,
0.0,
"cosine",
1 / 3,
),
)
EVAL_MIX = PHASES[0].mix
assert all(phase.tokens_per_step == TOKENS_PER_STEP for phase in PHASES)
def allocate_budgets(total_tokens: int, mix: dict[str, float]) -> dict[str, int]:
if not math.isclose(sum(mix.values()), 1.0):
raise ValueError("Source fractions must add up to 1.0.")
budgets = {source: int(total_tokens * fraction) for source, fraction in mix.items()}
budgets[next(iter(budgets))] += total_tokens - sum(budgets.values())
return budgets
def sha256_file(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as file:
for block in iter(lambda: file.read(1024 * 1024), b""):
digest.update(block)
return digest.hexdigest()
def phase_dir(phase: Phase) -> Path:
return RUN_DIR / phase.name
def phase_final_dir(phase: Phase) -> Path:
return FINAL_DIR if phase.name == PHASES[-1].name else phase_dir(phase) / f"{phase.name}-final"
def cache_paths(name: str) -> tuple[Path, Path]:
return DATA_ROOT / f"{name}.bin", DATA_ROOT / f"{name}.json"
def create_tokenizer() -> PreTrainedTokenizerFast:
tokenizer = PreTrainedTokenizerFast(
tokenizer_file=str(TOKENIZER_FILE),
pad_token="<|pad|>",
bos_token="<|bos|>",
eos_token="<|eos|>",
unk_token="<|unk|>",
additional_special_tokens=[
"<|system|>",
"<|user|>",
"<|assistant|>",
"<|tool|>",
"<|tool_result|>",
"<|end_of_turn|>",
],
model_max_length=BASE_SEQUENCE_LENGTH,
)
assert len(tokenizer) == VOCAB_SIZE
assert all(getattr(tokenizer, name) == token_id for name, token_id in SPECIAL_TOKEN_IDS.items())
return tokenizer
def create_model(tokenizer: PreTrainedTokenizerFast) -> LlamaForCausalLM:
model = LlamaForCausalLM(
LlamaConfig(
vocab_size=len(tokenizer),
hidden_size=640,
intermediate_size=2048,
num_hidden_layers=16,
num_attention_heads=10,
num_key_value_heads=2,
max_position_embeddings=BASE_SEQUENCE_LENGTH,
rope_theta=10_000,
tie_word_embeddings=True,
use_cache=False,
pad_token_id=tokenizer.pad_token_id,
bos_token_id=tokenizer.bos_token_id,
eos_token_id=tokenizer.eos_token_id,
)
)
assert model.num_parameters() == MODEL_PARAMETERS
return model
def is_evaluation_document(example: dict, text: str) -> bool:
key = text
for field in ("id", "document_id", "doc_id", "url"):
value = example.get(field)
if isinstance(value, str) and value:
key = value
break
digest = hashlib.blake2b(key.encode(), digest_size=8).digest()
return int.from_bytes(digest, "big") % EVAL_HASH_MODULUS < EVAL_HASH_BUCKETS
def load_stream(source_name: str, shuffle_seed: int):
source = SOURCES[source_name]
return load_dataset(
source.dataset,
source.config,
split="train",
streaming=True,
revision=source.revision,
).shuffle(seed=shuffle_seed, buffer_size=10_000)
def cache_metadata(
*,
cache_name: str,
kind: str,
target_tokens: int,
source_budgets: dict[str, int],
tokenizer_hash: str,
shuffle_seed: int,
sequence_length: int,
) -> dict:
return {
"cache_format_version": CACHE_FORMAT_VERSION,
"cache_name": cache_name,
"kind": kind,
"target_tokens": target_tokens,
"written_tokens": target_tokens,
"dtype": "uint16",
"sequence_length": sequence_length,
"tokenizer_file": str(TOKENIZER_FILE),
"tokenizer_sha256": tokenizer_hash,
"vocab_size": VOCAB_SIZE,
**SPECIAL_TOKEN_IDS,
"sources": {
name: {
"dataset": source.dataset,
"config": source.config,
"revision": source.revision,
}
for name, source in SOURCES.items()
},
"source_token_budgets": source_budgets,
"random_seed": SEED,
"shuffle_seed": shuffle_seed,
"shuffle_buffer_size": 10_000,
"evaluation_partition": {
"hash": "blake2b-64",
"modulus": EVAL_HASH_MODULUS,
"evaluation_buckets": EVAL_HASH_BUCKETS,
},
}
def token_cache_is_valid(path: Path, metadata_path: Path, expected: dict) -> bool:
if not path.is_file() or not metadata_path.is_file():
return False
try:
metadata = json.loads(metadata_path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError):
return False
expected_bytes = expected["target_tokens"] * np.dtype(np.uint16).itemsize
return metadata == expected and path.stat().st_size == expected_bytes
def close_memmap(memmap: np.memmap) -> None:
memmap.flush()
mmap = getattr(memmap, "_mmap", None)
if mmap is not None and not mmap.closed:
mmap.close()
def append_tokenized(buffer: list[int], texts: list[str], tokenizer: PreTrainedTokenizerFast) -> None:
for encoding in tokenizer.backend_tokenizer.encode_batch(texts, add_special_tokens=False):
buffer.append(tokenizer.bos_token_id)
buffer.extend(encoding.ids)
buffer.append(tokenizer.eos_token_id)
texts.clear()
def build_token_cache(
*,
path: Path,
metadata_path: Path,
metadata: dict,
tokenizer: PreTrainedTokenizerFast,
include_evaluation_documents: bool,
) -> None:
if token_cache_is_valid(path, metadata_path, metadata):
print(f"[=] Reusing {metadata['cache_name']}: {path}")
return
partial_path = path.with_suffix(f"{path.suffix}.partial")
for old in (partial_path, path, metadata_path):
old.unlink(missing_ok=True)
target_tokens = metadata["target_tokens"]
print(
f"[*] Building {metadata['cache_name']}: {target_tokens:,} tokens "
f"({target_tokens * 2 / 1_000_000_000:.2f} GB)"
)
token_memmap = np.memmap(partial_path, dtype=np.uint16, mode="w+", shape=(target_tokens,))
progress = tqdm(total=target_tokens, desc=metadata["cache_name"], unit="tok")
written = 0
def flush(buffer: list[int], limit: int) -> int:
nonlocal written
amount = min(len(buffer), limit)
token_memmap[written : written + amount] = np.asarray(buffer[:amount], dtype=np.uint16)
del buffer[:amount]
written += amount
progress.update(amount)
return amount
try:
for source_index, (source_name, budget) in enumerate(metadata["source_token_budgets"].items()):
source_written = 0
token_buffer: list[int] = []
text_batch: list[str] = []
stream = load_stream(source_name, metadata["shuffle_seed"] + source_index)
for example in stream:
text = example.get("text")
if not isinstance(text, str) or not (text := text.strip()):
continue
if is_evaluation_document(example, text) != include_evaluation_documents:
continue
text_batch.append(text)
if len(text_batch) < BATCH_TEXTS:
continue
append_tokenized(token_buffer, text_batch, tokenizer)
remaining = budget - source_written
if len(token_buffer) >= min(FLUSH_EVERY, remaining):
source_written += flush(token_buffer, remaining)
if source_written == budget:
break
if source_written < budget and text_batch:
append_tokenized(token_buffer, text_batch, tokenizer)
if source_written < budget:
source_written += flush(token_buffer, budget - source_written)
if source_written != budget:
raise RuntimeError(
f"{source_name} ended after {source_written:,} tokens; expected {budget:,}."
)
if written != target_tokens:
raise RuntimeError(f"Wrote {written:,} tokens; expected {target_tokens:,}.")
close_memmap(token_memmap)
token_memmap = None
partial_path.replace(path)
metadata_path.write_text(json.dumps(metadata, indent=2), encoding="utf-8")
except BaseException:
if token_memmap is not None:
close_memmap(token_memmap)
raise
finally:
progress.close()
class MemmapDataset(Dataset):
def __init__(self, path: Path, total_tokens: int, sequence_length: int):
if total_tokens % sequence_length:
raise ValueError("Token cache must contain complete sequences.")
self.path = path
self.total_tokens = total_tokens
self.sequence_length = sequence_length
self._data: np.memmap | None = None
@property
def data(self) -> np.memmap:
if self._data is None:
self._data = np.memmap(
self.path,
dtype=np.uint16,
mode="r",
shape=(self.total_tokens,),
)
return self._data
def __len__(self) -> int:
return self.total_tokens // self.sequence_length
def __getitem__(self, index: int) -> dict[str, torch.Tensor]:
start = index * self.sequence_length
input_ids = torch.tensor(self.data[start : start + self.sequence_length], dtype=torch.long)
return {"input_ids": input_ids, "labels": input_ids.clone()}
def phase_plan(phase: Phase, tokenizer_hash: str) -> dict:
return {
"name": phase.name,
"requested_tokens": phase.requested_tokens,
"actual_tokens": phase.actual_tokens,
"steps": phase.steps,
"source_token_budgets": phase.source_budgets,
"learning_rate": phase.learning_rate,
"warmup_ratio": phase.warmup_ratio,
"decay_type": phase.decay_type,
"min_lr_ratio": phase.min_lr_ratio,
"tokenizer_sha256": tokenizer_hash,
"sequence_length": phase.sequence_length,
"micro_batch_size": phase.micro_batch_size,
"gradient_accumulation_steps": phase.gradient_accumulation_steps,
"tokens_per_step": phase.tokens_per_step,
}
def ensure_phase_manifest(phase: Phase, tokenizer_hash: str) -> dict:
plan = phase_plan(phase, tokenizer_hash)
directory = phase_dir(phase)
manifest_path = directory / "phase-manifest.json"
directory.mkdir(parents=True, exist_ok=True)
if not manifest_path.exists():
manifest_path.write_text(json.dumps(plan, indent=2), encoding="utf-8")
return plan
existing = json.loads(manifest_path.read_text(encoding="utf-8"))
legacy_fields = {"sequence_length", "micro_batch_size", "gradient_accumulation_steps"}
legacy_plan = {key: value for key, value in plan.items() if key not in legacy_fields}
if phase.name == "phase-1" and existing == legacy_plan:
if existing["tokens_per_step"] != TOKENS_PER_STEP:
raise RuntimeError("phase-1 legacy manifest has an invalid token batch size.")
manifest_path.write_text(json.dumps(plan, indent=2), encoding="utf-8")
elif existing != plan:
raise RuntimeError(f"{phase.name} manifest differs from this run plan; refusing resume.")
return plan
def completed_phase_summary(phase: Phase, plan: dict) -> dict | None:
marker = phase_dir(phase) / "completed.json"
if not marker.is_file() or not phase_final_dir(phase).is_dir():
return None
try:
summary = json.loads(marker.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError):
return None
return summary if summary.get("phase_plan") == plan else None
def scheduler_args(phase: Phase) -> dict:
if phase.decay_type is None:
return {
"lr_scheduler_type": "constant_with_warmup",
"warmup_ratio": phase.warmup_ratio,
}
warmup_steps = math.ceil(phase.steps * phase.warmup_ratio)
return {
"lr_scheduler_type": "warmup_stable_decay",
"warmup_ratio": phase.warmup_ratio,
"lr_scheduler_kwargs": {
"num_stable_steps": 0,
"num_decay_steps": phase.steps - warmup_steps,
"decay_type": phase.decay_type,
"min_lr_ratio": phase.min_lr_ratio,
},
}
def training_arguments(phase: Phase) -> TrainingArguments:
return TrainingArguments(
output_dir=str(phase_dir(phase)),
run_name=f"{RUN_NAME}-{phase.name}",
max_steps=phase.steps,
per_device_train_batch_size=phase.micro_batch_size,
gradient_accumulation_steps=phase.gradient_accumulation_steps,
per_device_eval_batch_size=phase.micro_batch_size,
learning_rate=phase.learning_rate,
weight_decay=0.1,
adam_beta1=0.9,
adam_beta2=0.95,
max_grad_norm=1.0,
bf16=True,
optim="adamw_torch",
logging_steps=10,
logging_first_step=True,
logging_nan_inf_filter=False,
include_num_input_tokens_seen="all",
eval_strategy="steps",
eval_steps=1000,
prediction_loss_only=True,
save_steps=1000,
save_total_limit=3,
report_to=["tensorboard"],
logging_dir=str(RUN_DIR / "runs" / phase.name),
seed=SEED,
data_seed=SEED,
**scheduler_args(phase),
)
def load_phase_model(phase_index: int, tokenizer: PreTrainedTokenizerFast) -> LlamaForCausalLM:
phase = PHASES[phase_index]
if phase_index == 0:
set_seed(SEED)
return create_model(tokenizer)
previous = PHASES[phase_index - 1]
previous_final = phase_final_dir(previous)
if not previous_final.is_dir():
raise RuntimeError(f"Missing completed model for {previous.name}: {previous_final}")
config = LlamaConfig.from_pretrained(previous_final)
config.max_position_embeddings = phase.sequence_length
config.use_cache = False
model = LlamaForCausalLM.from_pretrained(previous_final, config=config)
assert model.num_parameters() == MODEL_PARAMETERS
return model
def run_phase(
phase_index: int,
tokenizer: PreTrainedTokenizerFast,
tokenizer_hash: str,
eval_path: Path,
evaluation_tokens: int,
) -> dict:
phase = PHASES[phase_index]
plan = ensure_phase_manifest(phase, tokenizer_hash)
if completed := completed_phase_summary(phase, plan):
print(f"[=] {phase.name} already completed; skipping.")
return completed
tokenizer.model_max_length = phase.sequence_length
cache_path, metadata_path = cache_paths(phase.name)
metadata = cache_metadata(
cache_name=phase.name,
kind="training",
target_tokens=phase.actual_tokens,
source_budgets=phase.source_budgets,
tokenizer_hash=tokenizer_hash,
shuffle_seed=SEED + (phase_index + 1) * 10_000,
sequence_length=phase.sequence_length,
)
build_token_cache(
path=cache_path,
metadata_path=metadata_path,
metadata=metadata,
tokenizer=tokenizer,
include_evaluation_documents=False,
)
args = training_arguments(phase)
trainer = Trainer(
model=load_phase_model(phase_index, tokenizer),
args=args,
train_dataset=MemmapDataset(cache_path, phase.actual_tokens, phase.sequence_length),
eval_dataset=MemmapDataset(eval_path, evaluation_tokens, phase.sequence_length),
)
checkpoint = get_last_checkpoint(args.output_dir)
if checkpoint:
print(f"[*] Resuming {phase.name} from {checkpoint}")
else:
print(f"[*] Starting {phase.name} from step zero")
torch.cuda.reset_peak_memory_stats()
train_result = trainer.train(resume_from_checkpoint=checkpoint)
trainer.log_metrics("train", train_result.metrics)
trainer.save_metrics("train", train_result.metrics)
trainer.save_state()
eval_metrics = trainer.evaluate()
trainer.log_metrics("eval", eval_metrics)
trainer.save_metrics("eval", eval_metrics)
final_path = phase_final_dir(phase)
trainer.save_model(final_path)
tokenizer.save_pretrained(final_path)
summary = {
"phase_plan": plan,
"final_model_path": str(final_path),
"peak_vram_gib": torch.cuda.max_memory_allocated() / 1024**3,
"train_metrics": train_result.metrics,
"eval_metrics": eval_metrics,
}
(phase_dir(phase) / "completed.json").write_text(json.dumps(summary, indent=2), encoding="utf-8")
return summary
def evaluation_tokens() -> int:
return EVAL_TARGET_TOKENS // BASE_SEQUENCE_LENGTH * BASE_SEQUENCE_LENGTH
def preflight(tokenizer: PreTrainedTokenizerFast, tokenizer_hash: str) -> None:
DATA_ROOT.mkdir(parents=True, exist_ok=True)
if any(part.lower() == "onedrive" for part in DATA_ROOT.resolve().parts):
raise RuntimeError(f"HANSE_DATA_DIR must not be inside OneDrive: {DATA_ROOT}")
if not torch.cuda.is_available():
raise RuntimeError("No ROCm GPU detected through torch.cuda.")
if not torch.cuda.is_bf16_supported():
raise RuntimeError("The detected GPU does not support BF16 training.")
model = create_model(tokenizer)
print(f"[+] GPU: {torch.cuda.get_device_name(0)}")
print(f"[+] Tokenizer: {TOKENIZER_FILE} ({tokenizer_hash})")
print(f"[+] Model parameters: {model.num_parameters():,}")
del model
eval_tokens = evaluation_tokens()
assert all(eval_tokens % phase.sequence_length == 0 for phase in PHASES)
cache_sizes = {
"evaluation": eval_tokens * 2,
**{phase.name: phase.actual_tokens * 2 for phase in PHASES},
}
for name in cache_sizes:
path, _ = cache_paths(name)
path.with_suffix(f"{path.suffix}.partial").unlink(missing_ok=True)
cache_bytes = sum(cache_sizes.values())
current_cache_bytes = sum(
min(path.stat().st_size, expected_bytes)
for name, expected_bytes in cache_sizes.items()
if (path := cache_paths(name)[0]).is_file()
)
additional_bytes = cache_bytes - current_cache_bytes
free_bytes = shutil.disk_usage(DATA_ROOT).free
print(f"[+] Data root: {DATA_ROOT.resolve()}")
print(f"[+] Token caches: {cache_bytes / 1_000_000_000:.2f} GB total")
print(f"[+] Additional disk space needed: {additional_bytes / 1_000_000_000:.2f} GB")
print(f"[+] Available disk space: {free_bytes / 1_000_000_000:.2f} GB")
print("[*] Phase plan:")
for phase in PHASES:
print(
f" {phase.name}: {phase.requested_tokens // 1_000_000_000}B, "
f"context={phase.sequence_length}, micro={phase.micro_batch_size}, "
f"accum={phase.gradient_accumulation_steps}, tokens/step={phase.tokens_per_step}"
)
print(f" evaluation: {eval_tokens:,} tokens")
if free_bytes < additional_bytes:
raise RuntimeError(
"Insufficient disk space for token caches. "
"Set HANSE_DATA_DIR to a larger non-OneDrive volume."
)
def verify_final_model() -> None:
tokenizer = AutoTokenizer.from_pretrained(FINAL_DIR, local_files_only=True)
model, loading_info = AutoModelForCausalLM.from_pretrained(
FINAL_DIR,
local_files_only=True,
output_loading_info=True,
)
assert not any(
loading_info[key] for key in ("missing_keys", "unexpected_keys", "mismatched_keys")
)
assert len(tokenizer) == VOCAB_SIZE
assert model.num_parameters() == MODEL_PARAMETERS
assert model.config.max_position_embeddings == PHASES[-1].sequence_length
assert tokenizer.model_max_length == PHASES[-1].sequence_length
assert model.get_input_embeddings().weight.data_ptr() == model.get_output_embeddings().weight.data_ptr()
model.config.use_cache = True
inputs = tokenizer("Hanse", return_tensors="pt")
inputs.pop("token_type_ids", None)
with torch.inference_mode():
output = model.generate(**inputs, max_new_tokens=1, do_sample=False)
assert output.shape[1] == inputs["input_ids"].shape[1] + 1
def print_final_summary(summaries: list[dict]) -> None:
print("[+] Base pretraining complete:")
for summary in summaries:
plan = summary["phase_plan"]
print(
f" {plan['name']}: {plan['actual_tokens']:,} tokens, "
f"train_loss={summary['train_metrics'].get('train_loss')}, "
f"eval_loss={summary['eval_metrics'].get('eval_loss')}, "
f"peak_vram={summary['peak_vram_gib']:.2f} GiB"
)
print(f" total training tokens: {sum(s['phase_plan']['actual_tokens'] for s in summaries):,}")
print(f" final output: {FINAL_DIR}")
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--dry-run", action="store_true")
args = parser.parse_args()
if not TOKENIZER_FILE.is_file():
raise FileNotFoundError(f"Tokenizer not found: {TOKENIZER_FILE}")
tokenizer = create_tokenizer()
tokenizer_hash = sha256_file(TOKENIZER_FILE)
preflight(tokenizer, tokenizer_hash)
if args.dry_run:
print("[+] Dry run complete; no data was downloaded or training started.")
return
eval_tokens = evaluation_tokens()
eval_path, eval_metadata_path = cache_paths("evaluation")
build_token_cache(
path=eval_path,
metadata_path=eval_metadata_path,
metadata=cache_metadata(
cache_name="evaluation",
kind="evaluation",
target_tokens=eval_tokens,
source_budgets=allocate_budgets(eval_tokens, EVAL_MIX),
tokenizer_hash=tokenizer_hash,
shuffle_seed=SEED,
sequence_length=BASE_SEQUENCE_LENGTH,
),
tokenizer=tokenizer,
include_evaluation_documents=True,
)
summaries = [
run_phase(index, tokenizer, tokenizer_hash, eval_path, eval_tokens)
for index in range(len(PHASES))
]
verify_final_model()
print_final_summary(summaries)
if __name__ == "__main__":
main()