初始化项目,由ModelHub XC社区提供模型
Model: ayh015/myLightningOPD Source: Original Platform
This commit is contained in:
51
slime/utils/train_metric_utils.py
Normal file
51
slime/utils/train_metric_utils.py
Normal file
@@ -0,0 +1,51 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import logging
|
||||
from argparse import Namespace
|
||||
from collections.abc import Callable
|
||||
from copy import deepcopy
|
||||
|
||||
from slime.utils import tracking_utils
|
||||
from slime.utils.metric_utils import compute_rollout_step
|
||||
from slime.utils.timer import Timer
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def log_perf_data_raw(
|
||||
rollout_id: int, args: Namespace, is_primary_rank: bool, compute_total_fwd_flops: Callable
|
||||
) -> None:
|
||||
timer_instance = Timer()
|
||||
log_dict_raw = deepcopy(timer_instance.log_dict())
|
||||
timer_instance.reset()
|
||||
|
||||
if not is_primary_rank:
|
||||
return
|
||||
|
||||
log_dict = {f"perf/{key}_time": val for key, val in log_dict_raw.items()}
|
||||
|
||||
if ("perf/actor_train_time" in log_dict) and (compute_total_fwd_flops is not None):
|
||||
total_fwd_flops = compute_total_fwd_flops(seq_lens=timer_instance.seq_lens)
|
||||
|
||||
if "perf/log_probs_time" in log_dict:
|
||||
log_dict["perf/log_probs_tflops"] = total_fwd_flops / log_dict["perf/log_probs_time"]
|
||||
|
||||
if "perf/ref_log_probs_time" in log_dict:
|
||||
log_dict["perf/ref_log_probs_tflops"] = total_fwd_flops / log_dict["perf/ref_log_probs_time"]
|
||||
|
||||
if log_dict["perf/actor_train_time"] > 0:
|
||||
log_dict["perf/actor_train_tflops"] = 3 * total_fwd_flops / log_dict["perf/actor_train_time"]
|
||||
log_dict["perf/actor_train_tok_per_s"] = sum(timer_instance.seq_lens) / log_dict["perf/actor_train_time"]
|
||||
|
||||
if "perf/train_wait_time" in log_dict and "perf/train_time" in log_dict:
|
||||
total_time = log_dict["perf/train_wait_time"] + log_dict["perf/train_time"]
|
||||
if total_time > 0:
|
||||
log_dict["perf/step_time"] = total_time
|
||||
log_dict["perf/wait_time_ratio"] = log_dict["perf/train_wait_time"] / total_time
|
||||
|
||||
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")
|
||||
Reference in New Issue
Block a user