xuan-luo/temp / utils /save_init_checkpoint.py
xuan-luo's picture
download
raw
7.89 kB
#!/usr/bin/env python3
"""Save a Megatron checkpoint at iteration 1 with random initialization."""
from __future__ import annotations
import argparse
import os
import shlex
import subprocess
import sys
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
from utils.megatron_launcher import load_config
from utils.paths import DELTAKV_ROOT, resolve
def _build_command(config: dict, nproc: int, master_port: int) -> list[str]:
paths = config["paths"]
model = config["model"]
data = config["data"]
training = config["training"]
launcher = config["launcher"]
megatron_repo = resolve(paths["megatron_repo"])
entrypoint = megatron_repo / "pretrain_gpt.py"
tokenizer_dir = resolve(paths["tokenizer_dir"])
init_dir = resolve(str(paths["init_checkpoint"]))
init_dir.mkdir(parents=True, exist_ok=True)
command = [
sys.executable,
"-m",
"torch.distributed.run",
"--nnodes",
"1",
"--nproc-per-node",
str(nproc),
"--master-addr",
str(launcher["master_addr"]),
"--master-port",
str(master_port),
str(entrypoint),
"--num-layers",
str(model["num_layers"]),
"--hidden-size",
str(model["hidden_size"]),
"--ffn-hidden-size",
str(model["ffn_hidden_size"]),
"--num-attention-heads",
str(model["num_attention_heads"]),
"--num-query-groups",
str(model["num_query_groups"]),
"--kv-channels",
str(model["kv_channels"]),
"--seq-length",
str(data["seq_length"]),
"--max-position-embeddings",
str(model["max_position_embeddings"]),
"--position-embedding-type",
"rope",
"--rotary-percent",
"1.0",
"--rotary-base",
str(model["rope_theta"]),
"--normalization",
str(model["normalization"]),
"--norm-epsilon",
str(model["norm_epsilon"]),
"--swiglu",
"--disable-bias-linear",
"--attention-dropout",
str(model["attention_dropout"]),
"--hidden-dropout",
str(model["hidden_dropout"]),
"--tokenizer-type",
"GPT2BPETokenizer",
"--vocab-file",
str(tokenizer_dir / "vocab.json"),
"--merge-file",
str(tokenizer_dir / "merges.txt"),
"--make-vocab-size-divisible-by",
str(model["make_vocab_size_divisible_by"]),
"--untie-embeddings-and-output-weights",
"--mock-data",
"--micro-batch-size",
"1",
"--global-batch-size",
str(nproc),
"--train-iters",
"1",
"--lr",
"0",
"--min-lr",
"0",
"--lr-decay-style",
"constant",
"--lr-decay-iters",
"1",
"--lr-warmup-iters",
"0",
"--weight-decay",
"0",
"--clip-grad",
"0",
"--eval-interval",
"100000",
"--eval-iters",
"0",
"--log-interval",
"1",
"--save-interval",
"1",
"--save",
str(init_dir),
"--seed",
str(training["seed"]),
"--tensor-model-parallel-size",
str(training["tensor_model_parallel_size"]),
"--pipeline-model-parallel-size",
str(training["pipeline_model_parallel_size"]),
"--context-parallel-size",
str(training["context_parallel_size"]),
"--transformer-impl",
str(training["transformer_impl"]),
"--bf16",
"--use-flash-attn",
"--use-distributed-optimizer",
"--no-gradient-accumulation-fusion",
"--optimizer",
"adam",
"--adam-beta1",
"0.9",
"--adam-beta2",
"0.95",
"--adam-eps",
"1e-8",
"--distributed-backend",
"nccl",
"--no-create-attention-mask-in-dataloader",
]
# Megatron ignores --num-query-groups unless --group-query-attention is set.
if int(model["num_query_groups"]) != int(model["num_attention_heads"]):
command += ["--group-query-attention"]
# Default True: baseline/depth_delta configs never set this key and always
# want qk_layernorm on. gqa/kvpath explicitly set it to false (no
# QK-Norm, see docs/attn_formulas.md) -- this must match their actual
# training config, or the init checkpoint's shapes (q_layernorm/k_layernorm
# present or absent) won't match the real training run's.
if model.get("qk_layernorm", True):
command += ["--qk-layernorm"]
attention_variant = model.get("attention_variant")
if attention_variant is not None:
command += ["--experimental-attention-variant", str(attention_variant)]
if attention_variant in {"depth_delta", "depth_delta_down"}:
command += [
"--depth-delta-window-size",
str(model["depth_delta_window_size"]),
"--depth-delta-group-size",
str(model["depth_delta_group_size"]),
"--depth-delta-rank",
str(model["depth_delta_rank"]),
"--depth-delta-init-scale",
str(model["depth_delta_init_scale"]),
]
if attention_variant == "independent_kv":
command += [
"--independent-kv-group-size",
str(model["independent_kv_group_size"]),
]
follower_groups = model.get("independent_kv_follower_query_groups")
if follower_groups is not None:
command += [
"--independent-kv-follower-query-groups",
str(follower_groups),
]
if attention_variant == "kv_share":
command += [
"--kv-share-group-size",
str(model["kv_share_group_size"]),
]
if attention_variant in {"kvpath", "kvpath_branching"}:
command += [
"--kvpath-group-size",
str(model["kvpath_group_size"]),
"--kvpath-rank",
str(model["kvpath_rank"]),
"--kvpath-base-head-dim",
str(model["kvpath_base_head_dim"]),
]
if attention_variant == "kvpath_branching":
command += ["--kvpath-branch-head-dim", str(model["kvpath_branch_head_dim"])]
if model.get("kvpath_base_rmsnorm", False):
command += ["--kvpath-base-rmsnorm"]
if attention_variant == "kvpath_delta":
command += [
"--kvpath-group-size",
str(model["kvpath_group_size"]),
"--kvpath-rank",
str(model["kvpath_rank"]),
]
return command
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--config",
type=Path,
default=DELTAKV_ROOT / "training" / "baseline" / "config.toml",
)
parser.add_argument("--nproc-per-node", type=int, default=1)
parser.add_argument("--master-port", type=int, default=6101)
parser.add_argument("--dry-run", action="store_true")
args = parser.parse_args()
config = load_config(args.config.resolve())
command = _build_command(config, args.nproc_per_node, args.master_port)
print("Saving random-init checkpoint:\n" + shlex.join(command), flush=True)
if args.dry_run:
return
env = os.environ.copy()
env["CUDA_DEVICE_MAX_CONNECTIONS"] = "1"
env["PYTHONPATH"] = os.pathsep.join(
part for part in (str(DELTAKV_ROOT), env.get("PYTHONPATH")) if part
)
subprocess.run(command, cwd=DELTAKV_ROOT, check=True, env=env)
init_dir = resolve(str(config["paths"]["init_checkpoint"]))
marker = init_dir / "latest_checkpointed_iteration.txt"
if not marker.exists():
raise SystemExit(f"Checkpoint save failed: {marker} not found")
print(f"Random-init checkpoint saved to {init_dir}")
if __name__ == "__main__":
main()

Xet Storage Details

Size:
7.89 kB
·
Xet hash:
f1dcae116f6b8b55457d1cb17211ab9e929ea6d246a38a70af26068cc5c86b4e

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.