File size: 2,267 Bytes
6011e08
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

import json
from pathlib import Path
from types import SimpleNamespace
from typing import Annotated

import torch
import typer

from slime.ray.rollout import compute_metrics_from_samples
from slime.utils.types import Sample

_WHITELIST_KEYS = [
    "group_index",
    "index",
    "prompt",
    "response",
    "response_length",
    "label",
    "reward",
    "status",
    "metadata",
]


def main(
    # Deliberately make this name consistent with main training arguments
    load_debug_rollout_data: Annotated[str, typer.Option()],
    show_metrics: bool = True,
    show_samples: bool = True,
    category: list[str] = None,
):
    if category is None:
        category = ["train", "eval"]
    for rollout_id, path in _get_rollout_dump_paths(load_debug_rollout_data, category):
        print("-" * 80)
        print(f"{rollout_id=} {path=}")
        print("-" * 80)

        pack = torch.load(path)
        sample_dicts = pack["samples"]

        if show_metrics:
            # TODO read these configs from dumps
            args = SimpleNamespace(
                advantage_estimator="grpo",
                reward_key=None,
                log_reward_category=None,
            )
            sample_objects = [Sample.from_dict(s) for s in sample_dicts]
            metrics = compute_metrics_from_samples(args, sample_objects)
            print("metrics", metrics)

        if show_samples:
            for sample in sample_dicts:
                print(json.dumps({k: v for k, v in sample.items() if k in _WHITELIST_KEYS}))


def _get_rollout_dump_paths(load_debug_rollout_data: str, categories: list[str]):
    # may improve later
    for rollout_id in range(1000):
        for category in categories:
            prefix = {
                "train": "",
                "eval": "eval_",
            }[category]
            path = Path(load_debug_rollout_data.format(rollout_id=f"{prefix}{rollout_id}"))
            if path.exists():
                yield rollout_id, path


if __name__ == "__main__":
    """python -m slime.utils.debug_utils.display_debug_rollout_data --load-debug-rollout-data ..."""
    typer.run(main)