Files
myLightningOPD/slime/ray/rollout.py
ModelHub XC d4e0a1af66 初始化项目,由ModelHub XC社区提供模型
Model: ayh015/myLightningOPD
Source: Original Platform
2026-08-27 23:50:14 +08:00

689 lines
29 KiB
Python

# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import logging
import multiprocessing
import random
import time
from pathlib import Path
from typing import Any
import numpy as np
import ray
import torch
from ray.util.scheduling_strategies import PlacementGroupSchedulingStrategy
from slime.backends.sglang_utils.sglang_engine import SGLangEngine
from slime.rollout.base_types import call_rollout_fn
from slime.utils import tracking_utils
from slime.utils.health_monitor import RolloutHealthMonitor
from slime.utils.http_utils import _wrap_ipv6, find_available_port, get_host_info, init_http_client
from slime.utils.iter_utils import group_by
from slime.utils.logging_utils import configure_logger
from slime.utils.metric_checker import MetricChecker
from slime.utils.metric_utils import compute_pass_rate, compute_rollout_step, compute_statistics, dict_add_prefix
from slime.utils.misc import load_function
from slime.utils.ray_utils import Box
from slime.utils.seqlen_balancing import get_seqlen_balanced_partitions
from slime.utils.tracking_utils import init_tracking
from slime.utils.types import Sample
from ..utils.metric_utils import has_repetition
from .utils import NOSET_VISIBLE_DEVICES_ENV_VARS_LIST, Lock
logging.getLogger("httpx").setLevel(logging.WARNING)
logging.getLogger("httpcore").setLevel(logging.WARNING)
logger = logging.getLogger(__name__)
@ray.remote
class RolloutManager:
"""The class to run rollout and convert rollout data to training data."""
def __init__(self, args, pg):
configure_logger()
self.args = args
self.pg = pg
_start_router(args)
# TODO make args immutable
init_tracking(args, primary=False, router_addr=f"http://{args.sglang_router_ip}:{args.sglang_router_port}")
init_http_client(args)
data_source_cls = load_function(self.args.data_source_path)
self.data_source = data_source_cls(args)
self.generate_rollout = load_function(self.args.rollout_function_path)
self.eval_generate_rollout = load_function(self.args.eval_function_path)
self.custom_reward_post_process_func = None
if self.args.custom_reward_post_process_path is not None:
self.custom_reward_post_process_func = load_function(self.args.custom_reward_post_process_path)
self.custom_convert_samples_to_train_data_func = None
if self.args.custom_convert_samples_to_train_data_path is not None:
self.custom_convert_samples_to_train_data_func = load_function(
self.args.custom_convert_samples_to_train_data_path
)
logger.info(f"import {self.args.rollout_function_path} as generate_rollout function.")
logger.info(f"import {self.args.eval_function_path} as eval_generate_rollout function.")
if self.args.debug_train_only:
self.all_rollout_engines = []
else:
num_gpu_per_engine = min(args.rollout_num_gpus_per_engine, args.num_gpus_per_node)
num_engines = args.rollout_num_gpus // num_gpu_per_engine
self.all_rollout_engines = [None] * num_engines
self.num_new_engines = init_rollout_engines(args, pg, self.all_rollout_engines)
self.nodes_per_engine = max(1, args.rollout_num_gpus_per_engine // args.num_gpus_per_node)
self.rollout_engine_lock = Lock.options(num_cpus=1, num_gpus=0).remote()
self._metric_checker = MetricChecker.maybe_create(args)
if self.args.use_fault_tolerance:
self._health_monitor = RolloutHealthMonitor(self, args)
def dispose(self):
if self._metric_checker is not None:
self._metric_checker.dispose()
# TODO maybe rename "rollout_engines" and "all_rollout_engines" later
@property
def rollout_engines(self):
# when doing multi-node serving, we will only send request to node-0 for each engine.
return self.all_rollout_engines[:: self.nodes_per_engine]
def get_rollout_engines_and_lock(self):
return self.rollout_engines, self.rollout_engine_lock, self.num_new_engines
def get_num_rollout_per_epoch(self):
assert self.args.rollout_global_dataset
return len(self.data_source.dataset) // self.args.rollout_batch_size
def generate(self, rollout_id):
monitor_started = self.args.use_fault_tolerance and self._health_monitor.start()
start_time = time.time()
try:
data, metrics = self._get_rollout_data(rollout_id=rollout_id)
self._save_debug_rollout_data(data, rollout_id=rollout_id, evaluation=False)
_log_rollout_data(rollout_id, self.args, data, metrics, time.time() - start_time)
data = self._convert_samples_to_train_data(data)
return self._split_train_data_by_dp(data, self.train_parallel_config["dp_size"])
finally:
if monitor_started:
self._health_monitor.stop()
self.num_new_engines = init_rollout_engines(self.args, self.pg, self.all_rollout_engines)
else:
self.num_new_engines = 0
def eval(self, rollout_id):
if self.args.debug_train_only:
# if debug train only, we don't generate evaluation data
return
# TODO: add fault tolerance to eval
result = call_rollout_fn(self.eval_generate_rollout, self.args, rollout_id, self.data_source, evaluation=True)
data = result.data
self._save_debug_rollout_data(data, rollout_id=rollout_id, evaluation=True)
metrics = _log_eval_rollout_data(rollout_id, self.args, data, result.metrics)
if self._metric_checker is not None:
self._metric_checker.on_eval(metrics)
def save(self, rollout_id):
self.data_source.save(rollout_id)
def load(self, rollout_id=None):
self.data_source.load(rollout_id)
def offload(self):
return ray.get([engine.release_memory_occupation.remote() for engine in self.rollout_engines])
def onload(self, tags: list[str] = None):
return ray.get([engine.resume_memory_occupation.remote(tags=tags) for engine in self.rollout_engines])
def check_weights(self, action: str):
return ray.get([engine.check_weights.remote(action=action) for engine in self.rollout_engines])
def _get_rollout_data(self, rollout_id):
if self.args.load_debug_rollout_data:
data = torch.load(
open(self.args.load_debug_rollout_data.format(rollout_id=rollout_id), "rb"),
weights_only=False,
)["samples"]
data = [Sample.from_dict(sample) for sample in data]
if (ratio := self.args.load_debug_rollout_data_subsample) is not None:
original_num_rows = len(data)
rough_subsample_num_rows = int(original_num_rows * ratio)
data = data[: rough_subsample_num_rows // 2] + data[-rough_subsample_num_rows // 2 :]
logger.info(
f"Subsample loaded debug rollout data using {ratio=} and change num rows {original_num_rows} -> {len(data)}"
)
metrics = None
else:
data = call_rollout_fn(self.generate_rollout, self.args, rollout_id, self.data_source, evaluation=False)
metrics = data.metrics
data = data.samples
# flatten the data if it is a list of lists
while isinstance(data[0], list):
data = sum(data, [])
if self.args.disable_rollout_trim_samples:
logger.info(f"Collectd {len(data)} samples from rollout to train")
elif len(data) % self.args.global_batch_size != 0:
trim_len = (len(data) // self.args.global_batch_size) * self.args.global_batch_size
origin_data_length = len(data)
data = data[:trim_len]
logger.info(f"trim number of samples from {origin_data_length} to {trim_len}")
return data, metrics
def _save_debug_rollout_data(self, data, rollout_id, evaluation: bool):
# TODO to be refactored (originally Buffer._set_data)
if (path_template := self.args.save_debug_rollout_data) is not None:
path = Path(path_template.format(rollout_id=("eval_" if evaluation else "") + str(rollout_id)))
logger.info(f"Save debug rollout data to {path}")
path.parent.mkdir(parents=True, exist_ok=True)
# TODO may improve the format
if evaluation:
dump_data = dict(
samples=[sample.to_dict() for dataset_name, info in data.items() for sample in info["samples"]]
)
else:
dump_data = dict(
samples=[sample.to_dict() for sample in data],
)
torch.save(dict(rollout_id=rollout_id, **dump_data), path)
def _post_process_rewards(self, samples: list[Sample] | list[list[Sample]]):
if self.custom_reward_post_process_func is not None:
return self.custom_reward_post_process_func(self.args, samples)
raw_rewards = [sample.get_reward_value(self.args) for sample in samples]
if (
self.args.advantage_estimator in ["grpo", "gspo", "reinforce_plus_plus_baseline"]
and self.args.rewards_normalization
):
# group norm
rewards = torch.tensor(raw_rewards, dtype=torch.float)
if rewards.shape[-1] == self.args.n_samples_per_prompt * self.args.rollout_batch_size:
rewards = rewards.reshape(-1, self.args.n_samples_per_prompt)
else:
# when samples count are not equal in each group
rewards = rewards.view(-1, rewards.shape[-1])
mean = rewards.mean(dim=-1, keepdim=True)
rewards = rewards - mean
if self.args.advantage_estimator in ["grpo", "gspo"] and self.args.grpo_std_normalization:
std = rewards.std(dim=-1, keepdim=True)
rewards = rewards / (std + 1e-6)
return raw_rewards, rewards.flatten().tolist()
return raw_rewards, raw_rewards
def _convert_samples_to_train_data(self, samples: list[Sample] | list[list[Sample]]):
"""
Convert inference generated samples to training data.
"""
if self.custom_convert_samples_to_train_data_func is not None:
return self.custom_convert_samples_to_train_data_func(self.args, samples)
raw_rewards, rewards = self._post_process_rewards(samples)
assert len(raw_rewards) == len(samples)
assert len(rewards) == len(samples)
train_data = {
"tokens": [sample.tokens for sample in samples],
"response_lengths": [sample.response_length for sample in samples],
# some reward model, e.g. remote rm, may return multiple rewards,
# we could use key to select the reward.
"rewards": rewards,
"raw_reward": raw_rewards,
"truncated": [1 if sample.status == Sample.Status.TRUNCATED else 0 for sample in samples],
"sample_indices": [sample.index for sample in samples],
}
# loss mask
# TODO: compress the loss mask
loss_masks = []
for sample in samples:
# always instantiate loss_mask if not provided
if sample.loss_mask is None:
sample.loss_mask = [1] * sample.response_length
assert (
len(sample.loss_mask) == sample.response_length
), f"loss mask length {len(sample.loss_mask)} != response length {sample.response_length}"
if sample.remove_sample:
sample.loss_mask = [0] * sample.response_length
loss_masks.append(sample.loss_mask)
train_data["loss_masks"] = loss_masks
# overwriting the raw reward
if samples[0].metadata and "raw_reward" in samples[0].metadata:
train_data["raw_reward"] = [sample.metadata["raw_reward"] for sample in samples]
# For rollout buffer
if samples[0].metadata and "round_number" in samples[0].metadata:
train_data["round_number"] = [sample.metadata["round_number"] for sample in samples]
# Add rollout log probabilities for off-policy correction.
if all(s.rollout_log_probs is not None for s in samples):
train_data["rollout_log_probs"] = [sample.rollout_log_probs for sample in samples]
if all(s.rollout_routed_experts is not None for s in samples):
train_data["rollout_routed_experts"] = [sample.rollout_routed_experts for sample in samples]
if all(s.train_metadata is not None for s in samples):
train_data["metadata"] = [sample.train_metadata for sample in samples]
if all(s.multimodal_train_inputs is not None for s in samples):
train_data["multimodal_train_inputs"] = [sample.multimodal_train_inputs for sample in samples]
if "teacher_log_probs" in samples[0].__dict__:
train_data["teacher_log_probs"] = [sample.teacher_log_probs for sample in samples]
if "verifiable_rewards" in samples[0].__dict__:
train_data["verifiable_rewards"] = [
getattr(sample, "verifiable_rewards", None) for sample in samples
]
return train_data
def set_train_parallel_config(self, config: dict):
self.train_parallel_config = config
def _split_train_data_by_dp(self, data, dp_size):
"""Split the train data by data parallel size."""
rollout_data = {}
if "prompt" in data:
rollout_data["prompt"] = data["prompt"]
total_lengths = [len(t) for t in data["tokens"]]
data["total_lengths"] = total_lengths
if self.args.balance_data:
partitions = get_seqlen_balanced_partitions(total_lengths, dp_size, equal_size=True)
else:
partitions = [range(i, len(total_lengths), dp_size) for i in range(dp_size)]
rollout_data_refs = []
for i in range(dp_size):
rollout_data = {}
partition = partitions[i]
rollout_data["partition"] = partition
for key in [
"tokens",
"multimodal_train_inputs",
"response_lengths",
"rewards",
"truncated",
"loss_masks",
"round_number",
"sample_indices",
"rollout_log_probs",
"rollout_routed_experts",
"prompt",
"teacher_log_probs",
"verifiable_rewards",
]:
if key not in data:
continue
val = [data[key][j] for j in partition]
rollout_data[key] = val
# keys that need to be splited at train side
for key in [
"raw_reward",
"total_lengths",
]:
if key not in data:
continue
rollout_data[key] = data[key]
rollout_data_refs.append(Box(ray.put(rollout_data)))
return rollout_data_refs
def init_rollout_engines(args, pg, all_rollout_engines):
if args.debug_train_only:
return 0
num_gpu_per_engine = min(args.rollout_num_gpus_per_engine, args.num_gpus_per_node)
num_engines = args.rollout_num_gpus // num_gpu_per_engine
assert len(all_rollout_engines) == num_engines
if args.prefill_num_servers is not None:
prefill_num_servers = args.prefill_num_servers * args.rollout_num_gpus_per_engine // num_gpu_per_engine
assert (
num_engines > prefill_num_servers
), f"num_engines {num_engines} should be larger than prefill_num_servers {prefill_num_servers}"
pg, reordered_bundle_indices = pg
RolloutRayActor = ray.remote(SGLangEngine)
rollout_engines = []
for i in range(num_engines):
if all_rollout_engines[i] is not None:
continue
num_gpus = 0.2
num_cpus = num_gpus
scheduling_strategy = PlacementGroupSchedulingStrategy(
placement_group=pg,
placement_group_capture_child_tasks=True,
placement_group_bundle_index=reordered_bundle_indices[i * num_gpu_per_engine],
)
env_vars = {name: "1" for name in NOSET_VISIBLE_DEVICES_ENV_VARS_LIST} | {
"SGL_JIT_DEEPGEMM_PRECOMPILE": "false",
"SGLANG_JIT_DEEPGEMM_PRECOMPILE": "false",
"SGL_DISABLE_TP_MEMORY_INBALANCE_CHECK": "true",
"SGLANG_DISABLE_TP_MEMORY_INBALANCE_CHECK": "true",
"SGLANG_MEMORY_SAVER_CUDA_GRAPH": "true",
"SGLANG_BATCH_INVARIANT_OPS_ENABLE_MM_FALLBACK_VARIANT": "true",
"SGLANG_ENABLE_HEALTH_ENDPOINT_GENERATION": "false",
}
worker_type = "regular"
if args.prefill_num_servers is not None:
if i < prefill_num_servers:
worker_type = "prefill"
else:
worker_type = "decode"
rollout_engine = RolloutRayActor.options(
num_cpus=num_cpus,
num_gpus=num_gpus,
scheduling_strategy=scheduling_strategy,
runtime_env={
"env_vars": env_vars,
},
).remote(args, rank=i, worker_type=worker_type)
rollout_engines.append((i, rollout_engine))
all_rollout_engines[i] = rollout_engine
num_new_engines = len(rollout_engines)
if num_new_engines == 0:
return num_new_engines
if args.rollout_external:
addr_and_ports = _allocate_rollout_engine_addr_and_ports_external(args=args, rollout_engines=rollout_engines)
else:
addr_and_ports = _allocate_rollout_engine_addr_and_ports_normal(
args=args, num_engines=num_engines, rollout_engines=rollout_engines
)
# TODO: don't ray.get here to overlap train actor init with rollout engine init.
# somehow if we don't sync here, the --debug-rollout-only mode will crash.
init_handles = [engine.init.remote(**(addr_and_ports[rank])) for rank, engine in rollout_engines]
ray.get(init_handles)
return num_new_engines
def _allocate_rollout_engine_addr_and_ports_external(args, rollout_engines):
addr_and_ports = []
for rank, _ in rollout_engines:
[host, port] = args.rollout_external_engine_addrs[rank].split(":")
addr_and_ports.append(
dict(
dist_init_addr=None,
nccl_port=None,
host=host,
port=int(port),
)
)
return addr_and_ports
def _allocate_rollout_engine_addr_and_ports_normal(*, args, num_engines, rollout_engines):
# get ports
# there are 4 ports we need to allocate
# 1. server port
# 2. nccl port
# 3. dist_init_addr port
# 4. other ports for dp_attention, which is of size 4 + dp_size
num_engines_per_node = max(
1, min(args.num_gpus_per_node, args.rollout_num_gpus) // args.rollout_num_gpus_per_engine
)
addr_and_ports = [{} for _ in range(num_engines)]
# Calculate prefill limit to identify prefill engines
prefill_limit = 0
if args.prefill_num_servers is not None:
num_gpu_per_engine = min(args.rollout_num_gpus_per_engine, args.num_gpus_per_node)
prefill_limit = args.prefill_num_servers * args.rollout_num_gpus_per_engine // num_gpu_per_engine
visited_nodes = set()
for rank, engine in rollout_engines:
if rank // num_engines_per_node in visited_nodes:
continue
visited_nodes.add(rank // num_engines_per_node)
# TODO: currently when restarting engines, we will set port for all engines on this node starting with this rank.
# e.g. for 8 gpus, if we are restarting engine on gpu 3, we will set port for engine 3,4,5,6,7 on this node.
num_engines_on_this_node = num_engines_per_node - (rank % num_engines_per_node)
def get_addr_and_ports(engine):
# use small ports to prevent ephemeral port between 32768 and 65536.
# also, ray uses port 10002-19999, thus we avoid near-10002 to avoid racing condition
start_port = 15000
def port(consecutive=1):
nonlocal start_port
_, port = ray.get(
engine._get_current_node_ip_and_free_port.remote(
start_port=start_port,
consecutive=consecutive,
)
)
start_port = port + consecutive
return port
def addr():
addr, _ = ray.get(engine._get_current_node_ip_and_free_port.remote())
return addr
return addr, port
get_addr, get_port = get_addr_and_ports(engine)
for i in range(num_engines_on_this_node):
current_rank = rank + i
addr_and_ports[current_rank]["host"] = get_addr()
addr_and_ports[current_rank]["port"] = get_port()
addr_and_ports[current_rank]["nccl_port"] = get_port()
if args.prefill_num_servers is not None and current_rank < prefill_limit:
addr_and_ports[current_rank]["disaggregation_bootstrap_port"] = get_port()
if args.rollout_num_gpus_per_engine > args.num_gpus_per_node:
num_node_per_engine = args.rollout_num_gpus_per_engine // args.num_gpus_per_node
if rank % num_node_per_engine == 0:
# this is the first node in the engine, we need to allocate the dist_init_addr port
dist_init_addr = f"{get_addr()}:{get_port(30 + args.sglang_dp_size)}"
for i in range(num_node_per_engine):
addr_and_ports[rank + i]["dist_init_addr"] = dist_init_addr
else:
for i in range(num_engines_on_this_node):
addr_and_ports[rank + i]["dist_init_addr"] = f"{get_addr()}:{get_port(30 + args.sglang_dp_size)}"
for i, _ in rollout_engines:
for key in ["port", "nccl_port", "dist_init_addr"]:
assert key in addr_and_ports[i], f"Engine {i} {key} is not set."
logger.info(f"Ports for engine {i}: {addr_and_ports[i]}")
return addr_and_ports
def _start_router(args):
"""start sgl router and slime router"""
if not args.rollout_num_gpus:
# No rollout engines (e.g. Lightning OPD) — skip router entirely.
args.sglang_router_ip = args.sglang_router_ip or "127.0.0.1"
args.sglang_router_port = args.sglang_router_port or 0
return
if args.sglang_router_ip is not None:
return
args.sglang_router_ip = _wrap_ipv6(get_host_info()[1])
if args.sglang_router_port is None:
args.sglang_router_port = find_available_port(random.randint(3000, 4000))
if args.use_slime_router:
assert args.prefill_num_servers is None, "slime router does not support prefill_num_servers."
from slime.router.router import run_router
router_args = args
else:
from sglang_router.launch_router import RouterArgs
from slime.utils.http_utils import run_router
router_args = RouterArgs.from_cli_args(args, use_router_prefix=True)
router_args.host = args.sglang_router_ip
router_args.port = args.sglang_router_port
router_args.prometheus_port = find_available_port(random.randint(4000, 5000))
router_args.log_level = "warn"
if args.prefill_num_servers is not None:
router_args.pd_disaggregation = True
if hasattr(router_args, "request_timeout_secs"):
router_args.request_timeout_secs = args.sglang_router_request_timeout_secs
logger.info(f"Launch router with args: {router_args}")
process = multiprocessing.Process(
target=run_router,
args=(router_args,),
)
process.daemon = True # Set the process as a daemon
process.start()
# Wait 3 seconds
time.sleep(3)
assert process.is_alive()
logger.info(f"Router launched at {args.sglang_router_ip}:{args.sglang_router_port}")
def _log_eval_rollout_data(rollout_id, args, data, extra_metrics: dict[str, Any] | None = None):
if args.custom_eval_rollout_log_function_path is not None:
custom_log_func = load_function(args.custom_eval_rollout_log_function_path)
if custom_log_func(rollout_id, args, data, extra_metrics):
return
log_dict = extra_metrics or {}
for key in data.keys():
rewards = data[key]["rewards"]
log_dict[f"eval/{key}"] = sum(rewards) / len(rewards)
if (samples := data[key].get("samples")) is not None:
log_dict |= dict_add_prefix(compute_metrics_from_samples(args, samples), f"eval/{key}/")
if "truncated" in data[key]:
truncated = data[key]["truncated"]
log_dict[f"eval/{key}-truncated_ratio"] = sum(truncated) / len(truncated)
if args.log_passrate:
log_dict |= dict_add_prefix(
compute_pass_rate(
flat_rewards=rewards,
group_size=args.n_samples_per_eval_prompt,
),
f"eval/{key}-",
)
logger.info(f"eval {rollout_id}: {log_dict}")
step = compute_rollout_step(args, rollout_id)
log_dict["eval/step"] = step
tracking_utils.log(args, log_dict, step_key="eval/step")
return log_dict
def _log_rollout_data(rollout_id, args, samples, rollout_extra_metrics, rollout_time):
if args.custom_rollout_log_function_path is not None:
custom_log_func = load_function(args.custom_rollout_log_function_path)
if custom_log_func(rollout_id, args, samples, rollout_extra_metrics, rollout_time):
return
if args.load_debug_rollout_data:
return
log_dict = {**(rollout_extra_metrics or {})}
response_lengths = [sample.effective_response_length for sample in samples]
log_dict["perf/rollout_time"] = rollout_time
if args.rollout_num_gpus:
log_dict["perf/tokens_per_gpu_per_sec"] = sum(response_lengths) / rollout_time / args.rollout_num_gpus
log_dict["perf/longest_sample_tokens_per_sec"] = max(response_lengths) / rollout_time
log_dict |= dict_add_prefix(compute_metrics_from_samples(args, samples), "rollout/")
logger.info(f"perf {rollout_id}: {log_dict}")
step = compute_rollout_step(args, rollout_id)
log_dict["rollout/step"] = step
tracking_utils.log(args, log_dict, step_key="rollout/step")
def compute_metrics_from_samples(args, samples):
response_lengths = [sample.effective_response_length for sample in samples]
log_dict = {}
log_dict |= dict_add_prefix(compute_statistics(response_lengths), "response_len/")
log_dict |= _compute_zero_std_metrics(args, samples)
log_dict |= _compute_spec_metrics(args, samples)
log_dict |= _compute_reward_cat_metrics(args, samples)
log_dict["repetition_frac"] = np.mean([int(has_repetition(s.response)) for s in samples]).item()
log_dict["truncated_ratio"] = np.mean([int(s.status == Sample.Status.TRUNCATED) for s in samples]).item()
return log_dict
def _compute_zero_std_metrics(args, all_samples: list[Sample]):
# only compute in GRPO-like algorithms where one prompt has multiple responses
if args.advantage_estimator == "ppo":
return {}
def _is_zero_std(samples: list[Sample]):
rewards = [sample.get_reward_value(args) for sample in samples]
return len(rewards) == 0 or all(rewards[0] == r for r in rewards)
all_sample_groups = group_by(all_samples, lambda s: s.group_index)
interesting_sample_groups = [g for g in all_sample_groups.values() if _is_zero_std(g)]
def _format_reward(reward):
# Handle dict rewards (from RL samples with meta_info)
if isinstance(reward, dict):
return "dict"
try:
return str(round(reward, 1))
except (TypeError, ValueError):
return str(reward)
interesting_rewards = [_format_reward(g[0].get_reward_value(args)) for g in interesting_sample_groups]
return {f"zero_std/count_{reward}": len(items) for reward, items in group_by(interesting_rewards).items()}
def _compute_spec_metrics(args, all_samples: list[Sample]):
if args.sglang_speculative_algorithm is None:
return {}
num_samples = len(all_samples)
metrics = {}
metrics["rollout/spec_accept_rate"] = (
sum(sample.spec_info.spec_accept_rate for sample in all_samples) / num_samples
)
metrics["rollout/spec_accept_length"] = (
sum(sample.spec_info.spec_accept_length for sample in all_samples) / num_samples
)
return metrics
def _compute_reward_cat_metrics(args, all_samples: list[Sample]):
reward_cat_key = args.log_reward_category
if reward_cat_key is None:
return {}
samples_of_reward_cat = group_by(all_samples, lambda s: s.reward[reward_cat_key])
return {f"error_cat/{reward_cat}": len(s) / len(all_samples) for reward_cat, s in samples_of_reward_cat.items()}