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

26 lines
781 B
Python

# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import logging
from pathlib import Path
import torch
logger = logging.getLogger(__name__)
def save_debug_train_data(args, *, rollout_id, rollout_data):
if (path_template := args.save_debug_train_data) is not None:
rank = torch.distributed.get_rank()
path = Path(path_template.format(rollout_id=rollout_id, rank=rank))
logger.info(f"Save debug train data to {path}")
path.parent.mkdir(parents=True, exist_ok=True)
torch.save(
dict(
rollout_id=rollout_id,
rank=rank,
rollout_data=rollout_data,
),
path,
)