iamrahulreddy's picture
add: sparse_moe architecture source
d8f717b verified
Raw
History Blame Contribute Delete
1.5 kB
from __future__ import annotations
import random
from typing import Any, Dict, List
import torch
import numpy as np
# Label building
def build_labels(batch: Dict[str, torch.Tensor]) -> torch.Tensor:
labels = batch["input_ids"].clone() # (batch_size, seq_len)
if "attention_mask" in batch:
labels[batch["attention_mask"] == 0] = -100
return labels # (batch_size, seq_len)
# Reproducibility
def set_seed(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
torch.backends.cudnn.benchmark = True
# Serialisation
def to_serializable(value: Any) -> Any:
if isinstance(value, torch.Tensor):
return value.detach().cpu().tolist()
if isinstance(value, np.ndarray):
return value.tolist()
if isinstance(value, np.integer):
return int(value)
if isinstance(value, np.floating):
return float(value)
if isinstance(value, dict):
return {k: to_serializable(v) for k, v in value.items()}
if isinstance(value, (list, tuple)):
return [to_serializable(v) for v in value]
if hasattr(value, "__dataclass_fields__"):
return {k: to_serializable(v) for k, v in value.__dict__.items()}
return value
# Multi-seed helpers
def parse_seed_list(raw: str) -> List[int]:
return [int(s.strip()) for s in raw.split(",") if s.strip()]