初始化项目,由ModelHub XC社区提供模型
Model: ayh015/myLightningOPD Source: Original Platform
This commit is contained in:
143
slime/utils/types.py
Normal file
143
slime/utils/types.py
Normal file
@@ -0,0 +1,143 @@
|
||||
# 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)
|
||||
Reference in New Issue
Block a user