myLightningOPD / slime /utils /types.py
ayh015's picture
Upload folder using huggingface_hub
6011e08 verified
Raw
History Blame Contribute Delete
5.5 kB
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from enum import Enum
from typing import Any
import torch
@dataclass
class Sample:
"""The sample generated"""
group_index: int | None = None
index: int | None = None
# prompt - can be:
# - str: raw text prompt
# - list[dict[str, str]]: chat messages format
prompt: str | list[dict[str, str]] = ""
tokens: list[int] = field(default_factory=list)
multimodal_inputs: dict[str, Any] = None # raw multimodal data, e.g. images, videos, etc.
multimodal_train_inputs: dict[str, Any] = None # processed multimodal data, e.g. pixel_values, etc.
# response
response: str = ""
response_length: int = 0
label: str | None = None
reward: float | dict[str, Any] | None = None
loss_mask: list[int] | None = None
weight_versions: list[str] = field(default_factory=list)
rollout_log_probs: list[float] | None = None # Log probabilities from rollout engine
rollout_routed_experts: list[list[int]] | None = None # Routed experts from rollout engine
remove_sample: bool = False
class Status(Enum):
PENDING = "pending"
COMPLETED = "completed"
TRUNCATED = "truncated"
ABORTED = "aborted"
# Indicates a recoverable or non-critical failure during generation (e.g., tool call failure,
# external API error, parsing error). Unlike ABORTED, FAILED samples may still contain partial
# valid output and can be retried or handled gracefully.
FAILED = "failed"
status: Status = Status.PENDING
metadata: dict = field(default_factory=dict)
# metadata used during training, e.g., what loss to use for this sample.
train_metadata: dict | None = None
class SpecInfo:
spec_accept_token_num: int = 0
spec_draft_token_num: int = 0
spec_verify_ct: int = 0
spec_accept_rate: float = 0.0
spec_accept_length: float = 0.0
def add(self, meta_info: dict, response_length: int):
self.spec_accept_token_num += meta_info["spec_accept_token_num"]
self.spec_draft_token_num += meta_info["spec_draft_token_num"]
self.spec_verify_ct += meta_info["spec_verify_ct"]
if self.spec_draft_token_num > 0:
# Notice: this does not iclude the bonus token generated by verify step.
self.spec_accept_rate = self.spec_accept_token_num / self.spec_draft_token_num
# self.spec_accept_rate = meta_info["spec_accept_rate"] #
if self.spec_verify_ct > 0:
self.spec_accept_length = response_length / self.spec_verify_ct
def to_dict(self):
return {
"spec_accept_token_num": self.spec_accept_token_num,
"spec_draft_token_num": self.spec_draft_token_num,
"spec_verify_ct": self.spec_verify_ct,
"spec_accept_rate": self.spec_accept_rate,
"spec_accept_length": self.spec_accept_length,
}
@staticmethod
def from_dict(data: dict):
info = Sample.SpecInfo()
info.spec_accept_token_num = data.get("spec_accept_token_num", 0)
info.spec_draft_token_num = data.get("spec_draft_token_num", 0)
info.spec_verify_ct = data.get("spec_verify_ct", 0)
info.spec_accept_rate = data.get("spec_accept_rate", 0.0)
info.spec_accept_length = data.get("spec_accept_length", 0.0)
return info
spec_info: SpecInfo = field(default_factory=SpecInfo)
def to_dict(self):
value = self.__dict__.copy()
value["status"] = self.status.value
value["spec_info"] = self.spec_info.to_dict()
return value
@staticmethod
def from_dict(data: dict):
data["status"] = Sample.Status(data["status"])
data["spec_info"] = Sample.SpecInfo.from_dict(data.get("spec_info", {}))
return Sample(**data)
def get_reward_value(self, args) -> float:
return self.reward if not args.reward_key else self.reward[args.reward_key]
@property
def effective_response_length(self):
return sum(self.loss_mask) if self.loss_mask is not None else self.response_length
@dataclass(frozen=True)
class ParamInfo:
name: str
dtype: torch.dtype
shape: torch.Size
attrs: dict
size: int
src_rank: int
# A dict-based batch produced along the rollout -> training path
# In Megatron backend, several fields are converted to torch.Tensor lists on GPU
# before being consumed by data iterators (see megatron_utils.actor._get_rollout_data).
RolloutBatch = dict[str, list[torch.Tensor] | list[int] | list[float] | list[str]]
@dataclass
class MultimodalType:
name: str # Type identifier used in message content (e.g., "image")
placeholder: str # Placeholder token in conversation messages (e.g., "<image>")
class MultimodalTypes:
IMAGE = MultimodalType(name="image", placeholder="<image>")
VIDEO = MultimodalType(name="video", placeholder="<video>")
AUDIO = MultimodalType(name="audio", placeholder="<audio>")
@classmethod
def all(cls) -> list[MultimodalType]:
return [cls.IMAGE, cls.VIDEO, cls.AUDIO]
@classmethod
def get(cls, name: str) -> MultimodalType | None:
return next((m for m in cls.all() if m.name == name), None)