mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
ci: pin the DSv4 bshd/thd parity test to the miles impl (#3162)
This commit is contained in:
@@ -59,6 +59,7 @@ def _prepare_args() -> ScriptArgs:
|
||||
return ScriptArgs(
|
||||
run_id="prepare",
|
||||
model_name=_MODEL_NAME,
|
||||
dsv4_impl="miles",
|
||||
task="gsm8k",
|
||||
enable_eval=False,
|
||||
num_nodes=1,
|
||||
@@ -100,6 +101,7 @@ def _run_args(
|
||||
return ScriptArgs(
|
||||
run_id=run_id,
|
||||
model_name=_MODEL_NAME,
|
||||
dsv4_impl="miles",
|
||||
task="gsm8k",
|
||||
enable_eval=False,
|
||||
num_nodes=1,
|
||||
@@ -198,22 +200,21 @@ def _run_train(
|
||||
return debug_root / qkv_format / "dump_details"
|
||||
|
||||
|
||||
def _load_rank_zero_train_data(directory: Path) -> dict:
|
||||
dump_files = sorted(directory.glob("*.pt"))
|
||||
expected_names = [f"{_ROLLOUT_ID}_{rank}.pt" for rank in range(_NUM_GPUS)]
|
||||
assert [path.name for path in dump_files] == expected_names
|
||||
def _load_train_data_shard(directory: Path) -> dict:
|
||||
# Train data is dumped once per (dp, cp) shard; this run has dp=1, cp=1.
|
||||
dump_files = sorted(directory.glob(f"{_ROLLOUT_ID}_*.pt"))
|
||||
assert len(dump_files) == 1, [path.name for path in dump_files]
|
||||
|
||||
payload = torch.load(dump_files[0], map_location="cpu", weights_only=False)
|
||||
assert payload["rollout_id"] == 0
|
||||
assert payload["rank"] == 0
|
||||
rollout_data = payload["rollout_data"]
|
||||
assert rollout_data["sample_indices"] == list(range(_NUM_SAMPLES))
|
||||
return rollout_data
|
||||
|
||||
|
||||
def _compare_train_log_probs(bshd_dir: Path, thd_dir: Path) -> None:
|
||||
bshd = _load_rank_zero_train_data(bshd_dir)
|
||||
thd = _load_rank_zero_train_data(thd_dir)
|
||||
bshd = _load_train_data_shard(bshd_dir)
|
||||
thd = _load_train_data_shard(thd_dir)
|
||||
bshd_log_probs: list[torch.Tensor] = []
|
||||
thd_log_probs: list[torch.Tensor] = []
|
||||
|
||||
|
||||
Reference in New Issue
Block a user