[4/N] Tiny enable UP ruleset in Ruff (#287)

This commit is contained in:
fzyzcjy
2025-12-01 20:03:31 +08:00
committed by GitHub
parent 16b8e569df
commit 0e2733d41e
70 changed files with 390 additions and 419 deletions
+25 -26
View File
@@ -1,6 +1,7 @@
import logging
from collections.abc import Iterable, Mapping, Sequence
from dataclasses import dataclass, field, fields
from typing import Any, Dict, Iterable, List, Mapping, Optional, Sequence
from typing import Any, Optional
from omegaconf import OmegaConf
@@ -14,7 +15,7 @@ def _first_not_none(*values: Any) -> Any:
return None
def _pick_from_mapping(data: Optional[Mapping[str, Any]], keys: Iterable[str]) -> Any:
def _pick_from_mapping(data: Mapping[str, Any] | None, keys: Iterable[str]) -> Any:
if not data:
return None
for key in keys:
@@ -28,11 +29,11 @@ class EvalEnvDatasetConfig:
"""Dataset-level generation parameters shared across delegate clients."""
name: str = ""
n_samples_per_eval_prompt: Optional[int] = None
temperature: Optional[float] = None
top_p: Optional[float] = None
top_k: Optional[int] = None
max_response_len: Optional[int] = None
n_samples_per_eval_prompt: int | None = None
temperature: float | None = None
top_p: float | None = None
top_k: int | None = None
max_response_len: int | None = None
# TODO: This is ugly, temporarily leave this. We should unify all the config name for dataset, default, and args. (advice from Tom.)
FIELD_SPECS = {
@@ -75,7 +76,7 @@ class EvalEnvDatasetConfig:
"Colon in dataset name is not allowed; use `n_samples_per_eval_prompt` to configure samples per prompt."
)
values: Dict[str, Any] = {"name": name}
values: dict[str, Any] = {"name": name}
for field_name, spec in cls.FIELD_SPECS.items():
dataset_value = _pick_from_mapping(dataset_cfg, spec["dataset_keys"])
default_value = _pick_from_mapping(defaults, spec["default_keys"])
@@ -88,9 +89,9 @@ class EvalEnvDatasetConfig:
obj = cls(**obj)
return obj
def to_payload(self) -> Dict[str, Any]:
def to_payload(self) -> dict[str, Any]:
"""Return a JSON-serializable payload for this dataset configuration."""
payload: Dict[str, Any] = {}
payload: dict[str, Any] = {}
for field_info in fields(self):
value = getattr(self, field_info.name)
if value is None:
@@ -104,11 +105,11 @@ class EvalEnvConfig:
"""Environment definition shared across delegate implementations."""
name: str = ""
url: Optional[str] = None
url: str | None = None
timeout_secs: int = 3600
max_retries: int = 1
headers: Dict[str, Any] = field(default_factory=dict)
defaults: Dict[str, Any] = field(default_factory=dict)
headers: dict[str, Any] = field(default_factory=dict)
defaults: dict[str, Any] = field(default_factory=dict)
@classmethod
def parse(cls, raw: Mapping[str, Any], defaults: Mapping[str, Any]) -> "EvalEnvConfig":
@@ -121,9 +122,9 @@ class EvalEnvConfig:
def _rebuild_delegate_config(
args, raw_delegate_config: Optional[Sequence[Mapping[str, Any]]], defaults: Optional[Mapping[str, Any]]
) -> List[EvalEnvConfig]:
envs: List[EvalEnvConfig] = []
args, raw_delegate_config: Sequence[Mapping[str, Any]] | None, defaults: Mapping[str, Any] | None
) -> list[EvalEnvConfig]:
envs: list[EvalEnvConfig] = []
defaults = defaults or {}
for env in raw_delegate_config or []:
env_name = str(env.get("name", "")).strip().lower()
@@ -151,13 +152,13 @@ class EvalClient:
def __init__(self, name: str):
self.name = name
def evaluate(self, args, rollout_id: int) -> tuple[Dict[str, Any], Dict[str, Any]]:
def evaluate(self, args, rollout_id: int) -> tuple[dict[str, Any], dict[str, Any]]:
raise NotImplementedError("Subclasses must implement this method")
def _flatten(result: Dict[str, Any], prefix: Optional[str] = None) -> Dict[str, Any]:
def _flatten(result: dict[str, Any], prefix: str | None = None) -> dict[str, Any]:
"""Flatten nested metric dicts into slash separated keys."""
flattened: Dict[str, Any] = {}
flattened: dict[str, Any] = {}
for key, value in (result or {}).items():
full_key = f"{prefix}/{key}" if prefix else key
if isinstance(value, dict):
@@ -174,15 +175,13 @@ class EvalDelegateClient:
self._delegates = list(delegates)
@classmethod
def maybe_create(
cls, args, env_configs: Optional[Sequence[EvalEnvConfig]] = None
) -> Optional["EvalDelegateClient"]:
def maybe_create(cls, args, env_configs: Sequence[EvalEnvConfig] | None = None) -> Optional["EvalDelegateClient"]:
env_configs = list(env_configs) if env_configs is not None else getattr(args, "eval_delegate_config", None)
if not env_configs:
return None
router_addr = f"http://{args.sglang_router_ip}:{args.sglang_router_port}"
delegates: List[EvalClient] = []
delegates: list[EvalClient] = []
for env_cfg in env_configs:
delegate = cls._create_delegate(env_cfg, router_addr)
if delegate is not None:
@@ -201,9 +200,9 @@ class EvalDelegateClient:
logger.warning("No delegate client registered for environment: %s", env_name)
return None
def evaluate(self, args, rollout_id: int) -> tuple[Dict[str, Any], Dict[str, Any]]:
aggregated_metrics: Dict[str, Any] = {}
raw_responses: Dict[str, Any] = {}
def evaluate(self, args, rollout_id: int) -> tuple[dict[str, Any], dict[str, Any]]:
aggregated_metrics: dict[str, Any] = {}
raw_responses: dict[str, Any] = {}
for delegate in self._delegates:
metrics, response = delegate.evaluate(args, rollout_id)
if metrics:
+7 -7
View File
@@ -3,7 +3,7 @@ from __future__ import annotations
import logging
import os
from pathlib import Path
from typing import Any, Optional
from typing import Any
from examples.eval.eval_delegate import EvalDelegateClient, _rebuild_delegate_config
from omegaconf import OmegaConf
@@ -13,7 +13,7 @@ from miles.rollout.sglang_rollout import generate_rollout as base_generate_rollo
logger = logging.getLogger(__name__)
_DELEGATE_CACHE: dict[str, tuple[Optional[float], Optional[EvalDelegateClient]]] = {}
_DELEGATE_CACHE: dict[str, tuple[float | None, EvalDelegateClient | None]] = {}
def generate_rollout(
@@ -32,7 +32,7 @@ def generate_rollout(
return result
def _get_delegate_client(args) -> Optional[EvalDelegateClient]:
def _get_delegate_client(args) -> EvalDelegateClient | None:
config_path = getattr(args, "eval_config", None)
if not config_path:
return None
@@ -48,7 +48,7 @@ def _get_delegate_client(args) -> Optional[EvalDelegateClient]:
return client
def _build_delegate_client(args, config_path: str) -> Optional[EvalDelegateClient]:
def _build_delegate_client(args, config_path: str) -> EvalDelegateClient | None:
cfg = OmegaConf.load(config_path)
cfg_dict = OmegaConf.to_container(cfg, resolve=True)
if not isinstance(cfg_dict, dict):
@@ -70,14 +70,14 @@ def _build_delegate_client(args, config_path: str) -> Optional[EvalDelegateClien
return EvalDelegateClient.maybe_create(args, env_configs=env_configs)
def _safe_mtime(path: str) -> Optional[float]:
def _safe_mtime(path: str) -> float | None:
try:
return os.path.getmtime(path)
except OSError:
return None
def _log_delegate_metrics(args, rollout_id: int, metrics: dict | None, raw_response: Optional[dict]) -> dict:
def _log_delegate_metrics(args, rollout_id: int, metrics: dict | None, raw_response: dict | None) -> dict:
flattened = _flatten_metrics(metrics)
if raw_response is not None:
logger.info("External eval raw response for rollout %s: %s", rollout_id, raw_response)
@@ -85,7 +85,7 @@ def _log_delegate_metrics(args, rollout_id: int, metrics: dict | None, raw_respo
return flattened
def _flatten_metrics(metric_source: Optional[dict]) -> dict:
def _flatten_metrics(metric_source: dict | None) -> dict:
flattened_metrics: dict[str, float] = {}
if not isinstance(metric_source, dict):
return flattened_metrics
+5 -5
View File
@@ -1,6 +1,6 @@
import logging
import time
from typing import Any, Dict, Optional
from typing import Any
import requests
from examples.eval.eval_delegate import EvalClient, EvalDelegateError
@@ -28,7 +28,7 @@ class SkillsEvalClient(EvalClient):
return None
return cls(config, router_url)
def evaluate(self, args, rollout_id: int) -> tuple[Dict[str, Any], Dict[str, Any]]:
def evaluate(self, args, rollout_id: int) -> tuple[dict[str, Any], dict[str, Any]]:
if not self._config.datasets:
logger.warning("No Skills datasets configured; skipping delegate evaluation.")
return {}, {}
@@ -38,7 +38,7 @@ class SkillsEvalClient(EvalClient):
metrics = response["raw_metrics"]
return metrics, response
def _build_payload(self, args, rollout_id: int) -> Dict[str, Any]:
def _build_payload(self, args, rollout_id: int) -> dict[str, Any]:
benchmarks = [cfg.to_payload() for cfg in self._config.datasets]
benchmarks = [cfg for cfg in benchmarks if cfg]
return {
@@ -47,8 +47,8 @@ class SkillsEvalClient(EvalClient):
"benchmarks": benchmarks,
}
def _request(self, payload: Dict[str, Any]) -> Dict[str, Any]:
last_error: Optional[Exception] = None
def _request(self, payload: dict[str, Any]) -> dict[str, Any]:
last_error: Exception | None = None
for attempt in range(1, self._max_retries + 1):
try:
response = self._session.post(
+4 -3
View File
@@ -1,7 +1,8 @@
from __future__ import annotations
from collections.abc import Mapping
from dataclasses import dataclass, field
from typing import Any, List, Mapping
from typing import Any
from examples.eval.eval_delegate import EvalEnvConfig, EvalEnvDatasetConfig
@@ -35,10 +36,10 @@ class SkillsEvalEnvDatasetConfig(EvalEnvDatasetConfig):
class SkillsEvalEnvConfig(EvalEnvConfig):
"""Environment configuration shared by the Skills client/server."""
datasets: List[SkillsEvalEnvDatasetConfig] = field(default_factory=list)
datasets: list[SkillsEvalEnvDatasetConfig] = field(default_factory=list)
@classmethod
def parse(cls, args, raw_env_config: Mapping[str, Any], defaults: Mapping[str, Any]) -> "SkillsEvalEnvConfig":
def parse(cls, args, raw_env_config: Mapping[str, Any], defaults: Mapping[str, Any]) -> SkillsEvalEnvConfig:
base_cfg: SkillsEvalEnvConfig = super().parse(raw_env_config, defaults)
datasets = raw_env_config.get("datasets") or []
base_cfg.datasets = [
+17 -16
View File
@@ -30,9 +30,10 @@ import sys
import threading
import time
import uuid
from collections.abc import Mapping
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Dict, List, Mapping
from typing import Any
REPO_ROOT = Path(__file__).resolve().parents[3]
if str(REPO_ROOT) not in sys.path:
@@ -56,8 +57,8 @@ logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(me
class EvalRequestPayload:
rollout_id: int
router_url: str
defaults: Dict[str, Any] = field(default_factory=dict)
benchmarks: List[SkillsEvalEnvDatasetConfig] = field(default_factory=list)
defaults: dict[str, Any] = field(default_factory=dict)
benchmarks: list[SkillsEvalEnvDatasetConfig] = field(default_factory=list)
# ---------------------------------------------------------------------------
@@ -83,8 +84,8 @@ def _hydra_overrides_from_benchmark(
router_url: str,
openai_model_name: str,
max_concurrent_requests: int,
) -> List[str]:
overrides: List[str] = []
) -> list[str]:
overrides: list[str] = []
for key, hydra_key in HYDRA_OVERRIDE_MAP.items():
value = getattr(benchmark_cfg, key, None)
if value is None:
@@ -114,7 +115,7 @@ class ServerConfig:
max_concurrent_requests: int = 512
@classmethod
def from_args(cls, args: argparse.Namespace) -> "ServerConfig":
def from_args(cls, args: argparse.Namespace) -> ServerConfig:
return cls(
output_root=Path(args.output_root).expanduser().resolve(),
cluster=args.cluster,
@@ -130,7 +131,7 @@ class SkillsEvaluator:
self._lock = threading.Lock()
self._config.output_root.mkdir(parents=True, exist_ok=True)
def evaluate(self, payload: EvalRequestPayload) -> Dict[str, Any]:
def evaluate(self, payload: EvalRequestPayload) -> dict[str, Any]:
if not payload.benchmarks:
warning_msg = "No benchmarks specified in delegate config; skipping NeMo Skills evaluation."
logger.warning(warning_msg)
@@ -149,8 +150,8 @@ class SkillsEvaluator:
run_dir = self._config.output_root / f"{int(time.time())}-{exp_name}"
run_dir.mkdir(parents=True, exist_ok=True)
runs: List[Dict[str, Any]] = []
raw_metrics: Dict[str, Any] = {}
runs: list[dict[str, Any]] = []
raw_metrics: dict[str, Any] = {}
with self._lock:
for benchmark in payload.benchmarks:
result = self._run_single_benchmark(
@@ -182,7 +183,7 @@ class SkillsEvaluator:
exp_name: str,
router_url: str,
run_dir: Path,
) -> Dict[str, Any]:
) -> dict[str, Any]:
name = benchmark.name
benchmark_run_dir = run_dir / name
benchmark_run_dir.mkdir(parents=True, exist_ok=True)
@@ -220,7 +221,7 @@ class SkillsEvaluator:
run_dir: Path,
defaults: Mapping[str, Any],
benchmark_cfg: SkillsEvalEnvDatasetConfig,
) -> List[str]:
) -> list[str]:
base_cmd = [
"ns",
"eval",
@@ -250,29 +251,29 @@ class SkillsEvaluator:
)
return base_cmd + hydra_overrides
def _build_env(self) -> Dict[str, str]:
def _build_env(self) -> dict[str, str]:
env = os.environ.copy()
return env
@staticmethod
def _run_command(cmd: List[str], *, env: Dict[str, str], log_path: Path):
def _run_command(cmd: list[str], *, env: dict[str, str], log_path: Path):
with open(log_path, "w", encoding="utf-8") as log_file:
process = subprocess.Popen(cmd, stdout=log_file, stderr=subprocess.STDOUT, env=env)
retcode = process.wait()
if retcode != 0:
with open(log_path, "r", encoding="utf-8", errors="ignore") as log_file:
with open(log_path, encoding="utf-8", errors="ignore") as log_file:
tail = "".join(log_file.readlines()[-200:])
raise RuntimeError(f"`ns eval` failed with exit code {retcode}. See {log_path}\n{tail}")
@staticmethod
def _collect_metrics(run_dir: Path, benchmark: str) -> Dict[str, Any]:
def _collect_metrics(run_dir: Path, benchmark: str) -> dict[str, Any]:
benchmark_name = benchmark.split(":")[0]
metrics_path = run_dir / "eval-results" / benchmark_name / "metrics.json"
if not metrics_path.exists():
logger.warning("Metrics file missing for %s at %s", benchmark_name, metrics_path)
return {}
try:
with open(metrics_path, "r", encoding="utf-8") as fp:
with open(metrics_path, encoding="utf-8") as fp:
metrics_data = json.load(fp)
except json.JSONDecodeError as exc:
logger.warning("Failed to parse %s: %s", metrics_path, exc)
@@ -2,7 +2,6 @@ import datetime
import os
import random
import time
from typing import List
import ray
import requests
@@ -25,7 +24,7 @@ class KiminaServerAndClientCluster:
class _KiminaClientCluster:
def __init__(self, servers: List["_KiminaServerActor"]):
def __init__(self, servers: list["_KiminaServerActor"]):
self._clients = [AsyncKiminaClient(api_url=ray.get(server.get_api_url.remote())) for server in servers]
self._next_client_index = 0
@@ -35,7 +34,7 @@ class _KiminaClientCluster:
return await client.check(*args, **kwargs)
def _create_actor_per_node(actor_cls) -> List:
def _create_actor_per_node(actor_cls) -> list:
# for simplicity, we use all available nodes
nodes = [n for n in ray.nodes() if n.get("Alive")]
assert len(nodes) > 0
@@ -4,7 +4,7 @@ import random
import re
from concurrent.futures import ProcessPoolExecutor
from pathlib import Path
from typing import Annotated, Optional
from typing import Annotated
import polars as pl
import torch
@@ -33,7 +33,7 @@ def process_flc(
dir_output: Path,
train_flc_select_num_rows: int,
val_flc_select_num_rows: int,
filter_difficulty: Optional[int],
filter_difficulty: int | None,
filter_solvable_by_rollout_dumps: str,
):
ds = load_dataset("m-a-p/FineLeanCorpus", split="train")
@@ -235,8 +235,8 @@ def main(
output_name: Annotated[str, typer.Option()] = None,
train_flc_select_num_rows: Annotated[int, typer.Option()] = 20000,
val_flc_select_num_rows: Annotated[int, typer.Option()] = 100,
filter_difficulty: Annotated[Optional[int], typer.Option()] = None,
filter_solvable_by_rollout_dumps: Annotated[Optional[str], typer.Option()] = None,
filter_difficulty: Annotated[int | None, typer.Option()] = None,
filter_solvable_by_rollout_dumps: Annotated[str | None, typer.Option()] = None,
):
dir_output = Path(dir_output_base) / (
output_name or f"{datetime.datetime.now().strftime('%Y%m%d%H%M%S')}-{random.randint(0, 1000000)}"
@@ -2,7 +2,6 @@ import asyncio
import logging
import re
from types import SimpleNamespace
from typing import Optional, Tuple
from kimina_client import SnippetStatus
@@ -47,7 +46,7 @@ def _single(arr):
return arr[0]
def _assemble_code(prompt: str, response: str) -> Tuple[Optional[str], Optional[str]]:
def _assemble_code(prompt: str, response: str) -> tuple[str | None, str | None]:
prompt_code_block = _extract_last_full_code_block(prompt)
assert prompt_code_block is not None
@@ -76,7 +75,7 @@ def _extract_last_full_code_block(text):
return matches[-1] if matches else None
_REWARD_FN: Optional[RewardFn] = None
_REWARD_FN: RewardFn | None = None
async def reward_fn(*args, **kwargs):
+2 -3
View File
@@ -3,7 +3,6 @@ import atexit
import queue
import threading
import time
from typing import List
# Import core functions from sglang_rollout directly to avoid code duplication
from miles.rollout.sglang_rollout import GenerateState, generate_and_rm_group
@@ -131,7 +130,7 @@ class AsyncRolloutWorker:
self.worker_thread.join(timeout=5)
print("Stopped async worker thread")
def get_completed_groups(self) -> List[tuple]:
def get_completed_groups(self) -> list[tuple]:
"""Get completed sample groups"""
completed = []
while True:
@@ -147,7 +146,7 @@ class AsyncRolloutWorker:
return self.output_queue.qsize()
async def generate_rollout_async(args, rollout_id: int, data_buffer) -> List[List[Sample]]:
async def generate_rollout_async(args, rollout_id: int, data_buffer) -> list[list[Sample]]:
"""
Simplified asynchronous rollout generation - using global continuous worker
"""
+3 -4
View File
@@ -3,7 +3,6 @@ import re
import time
import traceback
from copy import deepcopy
from typing import List
from miles.rollout.rm_hub import batched_async_rm
from miles.utils.http_utils import post
@@ -114,7 +113,7 @@ class RewriterAgent(Agent):
def __init__(self):
super().__init__()
async def rewrite(self, args, problem_statement, previous_solutions: List[str]) -> str:
async def rewrite(self, args, problem_statement, previous_solutions: list[str]) -> str:
"""Generates the rewrited solution."""
# 动态生成模板
@@ -135,7 +134,7 @@ class SelectorAgent(Agent):
def __init__(self):
super().__init__()
async def select(self, args, problem_statement, candidate_solutions: List[str]) -> str:
async def select(self, args, problem_statement, candidate_solutions: list[str]) -> str:
"""Generates the rewrited solution."""
# 动态生成模板
@@ -149,7 +148,7 @@ class SelectorAgent(Agent):
prompt = template.format(**format_params)
return await self.run(args, prompt, max_retries=10, key="selector")
def extract_selected_solution_idx(self, response: str, candidate_solutions: List[str]) -> int:
def extract_selected_solution_idx(self, response: str, candidate_solutions: list[str]) -> int:
"""Extracts the selected solution ID from the response."""
PATTERN = re.compile("Judgment:\s*(\d+)")
matched = PATTERN.findall(response)
+2 -2
View File
@@ -1,6 +1,6 @@
# Adapted from https://github.com/volcengine/verl/blob/cb809d66e46dfd3342d008628891a14a054fa424/recipe/retool/retool.py
import re
from typing import Any, Dict, List
from typing import Any
try:
from jinja2 import Template
@@ -59,7 +59,7 @@ For each function call, return a json object with function name and arguments wi
def format_conversation_with_tools(
prompt: str, tools: List[Dict[str, Any]] = None, system_prompt: str = None, messages: List[Dict[str, Any]] = None
prompt: str, tools: list[dict[str, Any]] = None, system_prompt: str = None, messages: list[dict[str, Any]] = None
) -> str:
"""Format conversation using Jinja2 template with tool support"""
template = Template(TOOL_TEMPLATE)
+5 -5
View File
@@ -14,7 +14,7 @@ import re
import subprocess
import tempfile
from contextlib import contextmanager
from typing import Any, Dict, List
from typing import Any
import psutil
@@ -328,15 +328,15 @@ class ToolRegistry:
},
)
def register_tool(self, name: str, tool_spec: Dict[str, Any]):
def register_tool(self, name: str, tool_spec: dict[str, Any]):
"""Register a new tool in the registry"""
self.tools[name] = tool_spec
def get_tool_specs(self) -> List[Dict[str, Any]]:
def get_tool_specs(self) -> list[dict[str, Any]]:
"""Get all tool specifications as a list"""
return list(self.tools.values())
async def execute_tool(self, tool_name: str, arguments: Dict[str, Any]) -> str:
async def execute_tool(self, tool_name: str, arguments: dict[str, Any]) -> str:
"""Execute a tool call with the given arguments"""
if tool_name not in self.tools:
return f"Error: Tool '{tool_name}' not found"
@@ -347,7 +347,7 @@ class ToolRegistry:
else:
return f"Error: Tool '{tool_name}' not implemented"
async def _execute_python(self, arguments: Dict[str, Any]) -> str:
async def _execute_python(self, arguments: dict[str, Any]) -> str:
"""Execute Python code using the sandbox"""
code = arguments.get("code", "")
if not code.strip():
+4 -5
View File
@@ -2,14 +2,13 @@ import asyncio
import os
import random
import re
from typing import Dict, List
import aiohttp
import chardet
# --- Utilities ---
def parse_snippet(snippet: str) -> List[str]:
def parse_snippet(snippet: str) -> list[str]:
segments = snippet.split("...")
return [s.strip() for s in segments if len(s.strip().split()) > 5]
@@ -28,7 +27,7 @@ def sanitize_search_query(query: str) -> str:
return sanitized_query
def filter_links(search_results: List[Dict]) -> List[str]:
def filter_links(search_results: list[dict]) -> list[str]:
links = []
for result in search_results:
for item in result.get("items", []):
@@ -61,7 +60,7 @@ async def fetch(session: aiohttp.ClientSession, url: str, semaphore: asyncio.Sem
return ""
async def fetch_all(urls: List[str], limit: int = 8) -> List[str]:
async def fetch_all(urls: list[str], limit: int = 8) -> list[str]:
semaphore = asyncio.Semaphore(limit)
timeout = aiohttp.ClientTimeout(total=5)
connector = aiohttp.TCPConnector(limit_per_host=limit, force_close=True)
@@ -92,7 +91,7 @@ def collect_context(snippet: str, doc: str) -> str:
return "\n".join(ctx_paras)
async def google_search(api_key, query, top_k=5, timeout: int = 60, proxy=None, snippet_only=False) -> List[Dict]:
async def google_search(api_key, query, top_k=5, timeout: int = 60, proxy=None, snippet_only=False) -> list[dict]:
timeout_obj = aiohttp.ClientTimeout(total=timeout)
session_kwargs = {}
if proxy:
@@ -18,7 +18,6 @@
import argparse
import json
import warnings
from typing import Optional
import datasets
import faiss
@@ -319,7 +318,7 @@ class Config:
class QueryRequest(BaseModel):
queries: list[str]
topk: Optional[int] = None
topk: int | None = None
return_scores: bool = False
+2 -4
View File
@@ -19,8 +19,6 @@ Usage:
}
"""
from typing import Dict, List, Optional
import aiohttp
@@ -29,8 +27,8 @@ async def local_search(
query: str,
top_k: int = 5,
timeout: int = 60,
proxy: Optional[str] = None,
) -> List[Dict]:
proxy: str | None = None,
) -> list[dict]:
"""
Call local search engine server and format results to match google_search_server.py output.
+2 -2
View File
@@ -8,7 +8,7 @@ results to the format expected by miles's training pipeline.
import logging
import os
from typing import Any, Dict
from typing import Any
from tau_bench.envs import get_env
from tau_bench.types import RunConfig
@@ -97,7 +97,7 @@ def res_to_sample(res: InteractionResult, task_index: int) -> Sample:
return sample
async def generate(args: Dict[str, Any], sample: Sample, sampling_params: dict) -> Sample:
async def generate(args: dict[str, Any], sample: Sample, sampling_params: dict) -> Sample:
"""
Generate a complete agent-environment interaction trajectory for tau-bench.
+10 -10
View File
@@ -1,7 +1,7 @@
import json
import logging
from dataclasses import dataclass
from typing import Any, Dict, List, Optional
from typing import Any
from sglang_tool_parser import parse_tools
from tau_bench.agents.tool_calling_agent import RESPOND_ACTION_NAME
@@ -17,7 +17,7 @@ class OpenAIToolCall:
id: str
type: str = "function"
function: Dict[str, Any] = None
function: dict[str, Any] = None
@dataclass
@@ -25,8 +25,8 @@ class OpenAIAssistantMessage:
"""OpenAI format assistant message structure"""
role: str = "assistant"
content: Optional[str] = None
tool_calls: Optional[List[OpenAIToolCall]] = None
content: str | None = None
tool_calls: list[OpenAIToolCall] | None = None
class OpenAICompatibleToolCallAdapter:
@@ -37,7 +37,7 @@ class OpenAICompatibleToolCallAdapter:
and provides OpenAI format output interface.
"""
def __init__(self, tools_info: List[Dict[str, Any]], parser_type: str = "qwen25"):
def __init__(self, tools_info: list[dict[str, Any]], parser_type: str = "qwen25"):
"""
Initialize adapter
@@ -48,7 +48,7 @@ class OpenAICompatibleToolCallAdapter:
self.tools_info = tools_info
self.parser_type = parser_type
def parse_response_to_openai_format(self, response: str) -> Dict[str, Any]:
def parse_response_to_openai_format(self, response: str) -> dict[str, Any]:
"""
Parse sglang response to OpenAI compatible format
@@ -78,7 +78,7 @@ class OpenAICompatibleToolCallAdapter:
logger.warning(f"Parsing failed with error: {str(e)}")
return {"openai_message": None, "parsed_result": None, "success": False, "error": str(e)}
def _convert_to_openai_message(self, normal_text: str, calls: List[Dict[str, Any]]) -> OpenAIAssistantMessage:
def _convert_to_openai_message(self, normal_text: str, calls: list[dict[str, Any]]) -> OpenAIAssistantMessage:
"""
Convert parsing results to OpenAI format assistant message
@@ -108,7 +108,7 @@ class OpenAICompatibleToolCallAdapter:
)
return result
def _call_to_action_sglang(self, calls: List[Any], text_response: str) -> Action:
def _call_to_action_sglang(self, calls: list[Any], text_response: str) -> Action:
"""
Convert sglang tool calls to Action object
@@ -136,7 +136,7 @@ class OpenAICompatibleToolCallAdapter:
return action
def get_openai_tools_format(self) -> List[Dict[str, Any]]:
def get_openai_tools_format(self) -> list[dict[str, Any]]:
"""
Get OpenAI format tool definitions
@@ -160,7 +160,7 @@ class OpenAICompatibleToolCallAdapter:
# Usage examples and factory functions
def create_openai_adapter(
tools_info: List[Dict[str, Any]], parser_type: str = "qwen25"
tools_info: list[dict[str, Any]], parser_type: str = "qwen25"
) -> OpenAICompatibleToolCallAdapter:
"""
Factory function to create OpenAI compatible tool call adapter
+2 -2
View File
@@ -1,10 +1,10 @@
from typing import Any, Dict, List
from typing import Any
from sglang.srt.function_call.function_call_parser import FunctionCallParser
from sglang.srt.managers.io_struct import Function, Tool
def parse_tools(response: str, tools: List[Dict[str, Any]], parser: str = "qwen25"):
def parse_tools(response: str, tools: list[dict[str, Any]], parser: str = "qwen25"):
"""
This function mimics the function call parser API from
https://github.com/sgl-project/sglang/blob/main/python/sglang/srt/entrypoints/http_server.py#L952
+26 -26
View File
@@ -2,7 +2,7 @@ import json
import logging
from dataclasses import dataclass
from enum import Enum
from typing import Any, Dict, List, Optional, Tuple
from typing import Any
from openai_tool_adapter import create_openai_adapter
from tau_bench.agents.base import Agent
@@ -27,15 +27,15 @@ class Status(Enum):
class InteractionResult:
prompt: str
reward: float
messages: List[Dict[str, Any]]
info: Dict[str, Any]
messages: list[dict[str, Any]]
info: dict[str, Any]
response: str = ""
loss_mask: Optional[List[int]] = None
tokens: Optional[int] = None
loss_mask: list[int] | None = None
tokens: int | None = None
status: Status = Status.COMPLETED
def call_to_action_sglang(calls: List[Any], text_response: str) -> Action:
def call_to_action_sglang(calls: list[Any], text_response: str) -> Action:
"""
Convert sglang response message to Action, similar to original message_to_action
but adapted for sglang response format.
@@ -87,7 +87,7 @@ class TrainableAgentMixin:
"""
return text.replace("You may call one or more functions to assist with the user query.", TOOL_INSTRUCTION)
async def _call_llm(self, url: str, payload: Dict[str, Any]) -> Dict[str, Any]:
async def _call_llm(self, url: str, payload: dict[str, Any]) -> dict[str, Any]:
"""
Make an LLM call tracking.
@@ -100,7 +100,7 @@ class TrainableAgentMixin:
"""
return await post(url, payload)
def _parse_tool(self, response: str) -> Dict[str, Any]:
def _parse_tool(self, response: str) -> dict[str, Any]:
"""
Parse tool calls from LLM response string.
@@ -125,7 +125,7 @@ class TrainableAgentMixin:
"""
return env.step(action)
def _initialize_environment(self, env, task_index: Optional[int]) -> Tuple[str, Dict[str, Any]]:
def _initialize_environment(self, env, task_index: int | None) -> tuple[str, dict[str, Any]]:
"""
Initialize the environment and get initial observation.
@@ -142,7 +142,7 @@ class TrainableAgentMixin:
env_reset_res = env.reset()
return env_reset_res.observation, env_reset_res.info.model_dump()
def _build_initial_messages(self, obs: str) -> List[Dict[str, Any]]:
def _build_initial_messages(self, obs: str) -> list[dict[str, Any]]:
"""
Build initial conversation messages.
@@ -154,7 +154,7 @@ class TrainableAgentMixin:
"""
return [{"role": "system", "content": self.wiki}, {"role": "user", "content": obs}]
def _prepare_prompt_tokens(self, state: GenerateState, messages: List[Dict[str, Any]]) -> Tuple[str, List[int]]:
def _prepare_prompt_tokens(self, state: GenerateState, messages: list[dict[str, Any]]) -> tuple[str, list[int]]:
"""
Prepare prompt text and tokenize it.
@@ -176,9 +176,9 @@ class TrainableAgentMixin:
async def asolve(
self,
env,
rollout_args: Dict[str, Any],
sampling_params: Dict[str, Any],
task_index: Optional[int] = None,
rollout_args: dict[str, Any],
sampling_params: dict[str, Any],
task_index: int | None = None,
max_num_steps: int = 30,
) -> InteractionResult:
"""
@@ -331,7 +331,7 @@ class TrainableAgentMixin:
res, total_reward, info, messages, loss_masks, prompt_token_ids, response_token_ids
)
def _get_token_delta(self, tokenizer: AutoTokenizer, messages: List[Dict]) -> Tuple[List[int], List[int]]:
def _get_token_delta(self, tokenizer: AutoTokenizer, messages: list[dict]) -> tuple[list[int], list[int]]:
"""
Calculate token delta for multi-turn conversations.
@@ -370,11 +370,11 @@ class TrainableAgentMixin:
self,
res: InteractionResult,
total_reward: float,
info: Dict[str, Any],
messages: List[Dict[str, Any]],
loss_masks: List[int],
prompt_token_ids: List[int],
response_token_ids: List[int],
info: dict[str, Any],
messages: list[dict[str, Any]],
loss_masks: list[int],
prompt_token_ids: list[int],
response_token_ids: list[int],
) -> InteractionResult:
"""
Build the final interaction result with all collected data.
@@ -420,13 +420,13 @@ class TrainableToolCallingAgent(ToolCallingAgent, TrainableAgentMixin):
def __init__(
self,
tools_info: List[Dict[str, Any]],
tools_info: list[dict[str, Any]],
wiki: str,
model: str,
provider: str,
temperature: float = 0.0,
rollout_args: Optional[Dict[str, Any]] = None,
sampling_params: Optional[Dict[str, Any]] = None,
rollout_args: dict[str, Any] | None = None,
sampling_params: dict[str, Any] | None = None,
):
# Initialize the parent ToolCallingAgent
super().__init__(
@@ -454,11 +454,11 @@ class TrainableToolCallingAgent(ToolCallingAgent, TrainableAgentMixin):
def agent_factory(
tools_info: List[Dict[str, Any]],
tools_info: list[dict[str, Any]],
wiki,
config: RunConfig,
rollout_args: Optional[Dict[str, Any]] = None,
sampling_params: Optional[Dict[str, Any]] = None,
rollout_args: dict[str, Any] | None = None,
sampling_params: dict[str, Any] | None = None,
) -> Agent:
if config.agent_strategy == "tool-calling":
return TrainableToolCallingAgent(
+13 -13
View File
@@ -1,4 +1,4 @@
from typing import Any, Dict, Optional, Tuple
from typing import Any
import torch
@@ -15,7 +15,7 @@ def masked_mean(x: torch.Tensor, loss_mask: torch.Tensor, expand: bool = False)
return result.expand_as(x) if expand else result
def metrics_append(metrics: Dict[str, list[torch.Tensor]], key: str, value: torch.Tensor) -> None:
def metrics_append(metrics: dict[str, list[torch.Tensor]], key: str, value: torch.Tensor) -> None:
"""
Every metrics-dict value is a list of 1D tensor, i.e., [torch.Tensor] with shapes exactly the same as log_probs.
@@ -58,14 +58,14 @@ def metrics_append(metrics: Dict[str, list[torch.Tensor]], key: str, value: torc
def calculate_veto_mask(
log_ratio: torch.Tensor,
loss_mask: torch.Tensor,
veto_threshold: Optional[float],
metrics: Dict[str, list[torch.Tensor]],
veto_threshold: float | None,
metrics: dict[str, list[torch.Tensor]],
) -> torch.Tensor:
if veto_threshold is None:
return torch.ones_like(log_ratio)
log_veto_threshold = torch.log(torch.tensor(veto_threshold, device=log_ratio.device))
# For each sequence, if it has any catastrophic tokens, return 0 for the sequence
catastrophic_tokens = ((log_ratio < log_veto_threshold)) & loss_mask.bool()
catastrophic_tokens = (log_ratio < log_veto_threshold) & loss_mask.bool()
has_catastrophic = catastrophic_tokens.any()
veto_mask = (~has_catastrophic).float().expand_as(log_ratio)
@@ -75,7 +75,7 @@ def calculate_veto_mask(
def truncate(
weights: torch.Tensor, loss_mask: torch.Tensor, metrics: Dict[str, list[torch.Tensor]], upper_bound: float
weights: torch.Tensor, loss_mask: torch.Tensor, metrics: dict[str, list[torch.Tensor]], upper_bound: float
) -> torch.Tensor:
assert upper_bound is not None
metrics_append(metrics, "truncate_fraction", (weights > upper_bound).int())
@@ -85,7 +85,7 @@ def truncate(
def clip(
weights: torch.Tensor,
loss_mask: torch.Tensor,
metrics: Dict[str, list[torch.Tensor]],
metrics: dict[str, list[torch.Tensor]],
lower_bound: float,
upper_bound: float,
) -> torch.Tensor:
@@ -98,10 +98,10 @@ def clip(
def mask(
weights: torch.Tensor,
loss_mask: torch.Tensor,
metrics: Dict[str, list[torch.Tensor]],
metrics: dict[str, list[torch.Tensor]],
lower_bound: float,
upper_bound: float,
) -> Tuple[torch.Tensor, torch.Tensor]:
) -> tuple[torch.Tensor, torch.Tensor]:
assert lower_bound is not None and upper_bound is not None and lower_bound < upper_bound
metrics_append(metrics, "mask_fraction_low", (weights < lower_bound).int())
metrics_append(metrics, "mask_fraction_high", (weights > upper_bound).int())
@@ -118,7 +118,7 @@ def compute_mis_weights(
train_log_probs: list[torch.Tensor],
rollout_log_probs: list[torch.Tensor],
loss_masks: list[torch.Tensor],
) -> Tuple[list[torch.Tensor], list[torch.Tensor], Dict[str, list[torch.Tensor]]]:
) -> tuple[list[torch.Tensor], list[torch.Tensor], dict[str, list[torch.Tensor]]]:
"""
Compute the importance sampling (IS) weights and metrics between the inference and training engine.
Args:
@@ -134,7 +134,7 @@ def compute_mis_weights(
metrics: The metrics for the importance sampling weights, a dict of list[torch.Tensor]. 1D tensor each.
"""
metrics: Dict[str, list[torch.Tensor]] = {}
metrics: dict[str, list[torch.Tensor]] = {}
tis_lower_bound = args.tis_lower_bound if args.tis_lower_bound is not None else 1.0 / args.tis_upper_bound
rs_lower_bound = args.rs_lower_bound if args.rs_lower_bound is not None else tis_lower_bound
@@ -266,7 +266,7 @@ def compute_mis_weights_with_cp(
total_lengths: list[int],
response_lengths: list[int],
**kwargs: Any,
) -> Tuple[torch.Tensor, list[torch.Tensor], Dict[str, torch.Tensor]]:
) -> tuple[torch.Tensor, list[torch.Tensor], dict[str, torch.Tensor]]:
"""
Compute the importance sampling (IS) weights and metrics with context parallel.
Args:
@@ -330,7 +330,7 @@ def add_ppl_metrics(
train_log_prob: torch.Tensor,
rollout_log_prob: torch.Tensor,
loss_mask: torch.Tensor,
metrics: Dict[str, list[torch.Tensor]],
metrics: dict[str, list[torch.Tensor]],
):
loss_mask = loss_mask.float()
+1 -1
View File
@@ -75,7 +75,7 @@ def parse_fsdp_cli(extra_args_provider=None):
def load_fsdp_args(extra_args_provider=None):
args = parse_fsdp_cli(extra_args_provider)
if args.config:
with open(args.config, "r") as f:
with open(args.config) as f:
data = yaml.safe_load(f) or {}
for k, v in data.items():
if not hasattr(args, k):
@@ -1,4 +1,4 @@
from typing import Any, Dict, List
from typing import Any
import torch
import torch.nn as nn
@@ -18,13 +18,13 @@ class FSDPCPUAdamWrapper:
Following the parameter copy pattern from update_weight_utils.py
"""
def __init__(self, optimizer_config: Dict[str, Any], model: nn.Module) -> None:
def __init__(self, optimizer_config: dict[str, Any], model: nn.Module) -> None:
from deepspeed.ops.adam import DeepSpeedCPUAdam
self.model: nn.Module = model
self.gpu_params: List[nn.Parameter] = list(model.parameters())
self.optimizer_config: Dict[str, Any] = optimizer_config
self.cpu_params: List[torch.Tensor] = []
self.gpu_params: list[nn.Parameter] = list(model.parameters())
self.optimizer_config: dict[str, Any] = optimizer_config
self.cpu_params: list[torch.Tensor] = []
self.cpu_optimizer: DeepSpeedCPUAdam
# Create CPU shadow copies of parameters using the pattern from update_weight_utils.py
+12 -14
View File
@@ -1,5 +1,3 @@
from typing import List, Optional
import torch
import torch.nn as nn
import torch.nn.functional as F
@@ -19,8 +17,8 @@ def fused_experts_impl(
w2: torch.Tensor,
topk_weights: torch.Tensor,
topk_ids: torch.Tensor,
b1: Optional[torch.Tensor] = None,
b2: Optional[torch.Tensor] = None,
b1: torch.Tensor | None = None,
b2: torch.Tensor | None = None,
inplace: bool = True,
activation: str = "silu",
apply_router_weight_on_input: bool = False,
@@ -29,17 +27,17 @@ def fused_experts_impl(
use_int8_w8a16: bool = False,
use_int4_w4a16: bool = False,
per_channel_quant: bool = False,
w1_scale: Optional[torch.Tensor] = None,
w2_scale: Optional[torch.Tensor] = None,
w1_zp: Optional[torch.Tensor] = None,
w2_zp: Optional[torch.Tensor] = None,
a1_scale: Optional[torch.Tensor] = None,
a2_scale: Optional[torch.Tensor] = None,
block_shape: Optional[List[int]] = None,
w1_scale: torch.Tensor | None = None,
w2_scale: torch.Tensor | None = None,
w1_zp: torch.Tensor | None = None,
w2_zp: torch.Tensor | None = None,
a1_scale: torch.Tensor | None = None,
a2_scale: torch.Tensor | None = None,
block_shape: list[int] | None = None,
no_combine: bool = False,
routed_scaling_factor: Optional[float] = None,
gemm1_alpha: Optional[float] = None,
gemm1_limit: Optional[float] = None,
routed_scaling_factor: float | None = None,
gemm1_alpha: float | None = None,
gemm1_limit: float | None = None,
filter_expert: bool = True,
):
padded_size = 0
+5 -6
View File
@@ -3,7 +3,6 @@ import os
import socket
from argparse import Namespace
from contextlib import nullcontext
from typing import Dict, Optional
import ray
import torch
@@ -50,7 +49,7 @@ class MegatronTrainRayActor(TrainRayActor):
args: Namespace,
role: str,
with_ref: bool = False,
) -> Optional[int]:
) -> int | None:
monkey_patch_torch_dist()
super().init(args, role, with_ref)
@@ -283,7 +282,7 @@ class MegatronTrainRayActor(TrainRayActor):
data_iterator: list[DataIterator],
num_microbatches: list[int],
store_prefix: str = "",
) -> Dict[str, list[torch.Tensor]]:
) -> dict[str, list[torch.Tensor]]:
self.weights_backuper.restore(model_tag)
with timer(f"{store_prefix}log_probs"):
@@ -501,9 +500,9 @@ class MegatronTrainRayActor(TrainRayActor):
def connect_actor_critic(
self,
actor_handle: Optional[ActorHandle] = None,
master_address: Optional[str] = None,
master_port: Optional[int] = None,
actor_handle: ActorHandle | None = None,
master_address: str | None = None,
master_port: int | None = None,
) -> None:
if self.role == "actor":
master_address = ray.util.get_node_ip_address()
@@ -1,5 +1,5 @@
import logging
from typing import Callable, Dict, List
from collections.abc import Callable
logger = logging.getLogger(__name__)
@@ -10,9 +10,9 @@ class MapperRegistry:
"""
def __init__(self):
self._mappers: Dict[str, Callable] = {}
self._mappers: dict[str, Callable] = {}
def register(self, model_types: List[str], mapper_func: Callable):
def register(self, model_types: list[str], mapper_func: Callable):
if not callable(mapper_func):
raise TypeError(f"Mapper for {model_types} must be callable")
@@ -30,7 +30,7 @@ class MapperRegistry:
raise ValueError(f"Mapper for {name} is not registered.")
return self._mappers[name]
def list_registered_mappers(self) -> List[str]:
def list_registered_mappers(self) -> list[str]:
return list(self._mappers.keys())
+3 -3
View File
@@ -1,4 +1,4 @@
from typing import Callable, Union
from collections.abc import Callable
import torch
import torch.distributed as dist
@@ -184,10 +184,10 @@ def slice_with_cp(tokens: torch.Tensor, pad_value: tuple[int, float, Callable])
def slice_log_prob_with_cp(
log_prob: Union[list[float], torch.Tensor],
log_prob: list[float] | torch.Tensor,
total_length: int,
response_length: int,
) -> Union[list[float], torch.Tensor]:
) -> list[float] | torch.Tensor:
assert len(log_prob) == response_length
cp_size = mpu.get_context_parallel_world_size()
+9 -9
View File
@@ -1,6 +1,6 @@
import logging
from argparse import Namespace
from typing import Optional, Sequence, Union
from collections.abc import Sequence
import numpy as np
import torch
@@ -26,7 +26,7 @@ def get_batch(
data_iterator: "DataIterator",
keys: Sequence[str],
pad_multiplier: int = 128,
) -> dict[str, Union[torch.Tensor, PackedSeqParams, list[torch.Tensor], None]]:
) -> dict[str, torch.Tensor | PackedSeqParams | list[torch.Tensor] | None]:
"""
Generate a CP-ready micro-batch with packed sequence parameters.
@@ -98,7 +98,7 @@ def gather_log_data(
args: Namespace,
rollout_id: int,
log_dict: dict[str, float],
) -> Optional[dict[str, float]]:
) -> dict[str, float] | None:
"""
Gather per-rank metrics, reduce by mean on the DP source rank, and log.
@@ -150,8 +150,8 @@ class DataIterator:
def __init__(
self,
rollout_data: RolloutBatch,
micro_batch_size: Optional[int] = None,
micro_batch_indices: Optional[list[list[int]]] = None,
micro_batch_size: int | None = None,
micro_batch_indices: list[list[int]] | None = None,
) -> None:
"""Initialize an iterator over `rollout_data`.
@@ -167,7 +167,7 @@ class DataIterator:
assert micro_batch_size is None or micro_batch_indices is None
self.offset = 0
def get_next(self, keys: Sequence[str]) -> dict[str, Optional[list[object]]]:
def get_next(self, keys: Sequence[str]) -> dict[str, list[object] | None]:
"""Return the next micro-batch for the requested keys.
- If `micro_batch_indices` is provided, selects rows according to the current
@@ -206,7 +206,7 @@ class DataIterator:
def get_data_iterator(
args: Namespace,
model: Union[torch.nn.Module, Sequence[torch.nn.Module]],
model: torch.nn.Module | Sequence[torch.nn.Module],
rollout_data: RolloutBatch,
) -> tuple[list[DataIterator], list[int]]:
"""
@@ -438,8 +438,8 @@ def log_perf_data(rollout_id: int, args: Namespace) -> None:
def sync_actor_critic_data(
args: Namespace,
rollout_data: Optional[RolloutBatch] = None,
group: Optional[dist.ProcessGroup] = None,
rollout_data: RolloutBatch | None = None,
group: dist.ProcessGroup | None = None,
) -> None:
"""
Broadcast `values` (from critic) and optionally `log_probs`/`ref_log_probs`
+1 -1
View File
@@ -63,7 +63,7 @@ def init(args):
# Random seeds for reproducibility.
if args.rank == 0:
logger.info("> setting random seeds to {} ...".format(args.seed))
logger.info(f"> setting random seeds to {args.seed} ...")
_set_random_seed(
args.seed,
args.data_parallel_random_init,
+3 -3
View File
@@ -1,6 +1,6 @@
from argparse import Namespace
from collections.abc import Callable, Iterator
from typing import Any, Dict, Tuple, Union
from typing import Any
import torch
from megatron.core import mpu
@@ -214,7 +214,7 @@ def compute_advantages_and_returns(args: Namespace, rollout_data: RolloutBatch)
log_probs: list[torch.Tensor] = rollout_data.get("rollout_log_probs" if args.use_rollout_logprobs else "log_probs")
ref_log_probs: list[torch.Tensor] = rollout_data.get("ref_log_probs")
rewards: list[float] = rollout_data.get("rewards")
values: Union[None, list[torch.Tensor]] = rollout_data.get("values")
values: None | list[torch.Tensor] = rollout_data.get("values")
response_lengths: list[int] = rollout_data.get("response_lengths")
loss_masks: list[torch.Tensor] = rollout_data.get("loss_masks")
total_lengths: list[int] = rollout_data.get("total_lengths")
@@ -444,7 +444,7 @@ def policy_loss_function(
rollout_log_probs: list[torch.Tensor],
loss_masks: list[torch.Tensor],
**kwargs: Any,
) -> Tuple[torch.Tensor, list[torch.Tensor], Dict[str, torch.Tensor]]:
) -> tuple[torch.Tensor, list[torch.Tensor], dict[str, torch.Tensor]]:
rollout_log_probs = torch.cat(rollout_log_probs, dim=0)
old_log_probs = torch.cat(train_log_probs, dim=0)
tis = torch.exp(old_log_probs - rollout_log_probs)
@@ -2,7 +2,7 @@
import argparse
import inspect
from contextlib import nullcontext
from typing import Literal, Optional
from typing import Literal
import torch
from megatron.core import tensor_parallel
@@ -39,8 +39,8 @@ class LinearForLastLayer(torch.nn.Linear):
def forward(
self,
input_: torch.Tensor,
weight: Optional[torch.Tensor] = None,
runtime_gather_output: Optional[bool] = None,
weight: torch.Tensor | None = None,
runtime_gather_output: bool | None = None,
) -> tuple[torch.Tensor, None]:
logits = super().forward(input_)
logits = logits.float()
@@ -53,9 +53,7 @@ def get_model_provider_func(
args: argparse.Namespace,
role: Literal["actor", "critic"] = "actor",
):
def model_provider(
pre_process: bool = True, post_process: bool = True, vp_stage: Optional[int] = None
) -> GPTModel:
def model_provider(pre_process: bool = True, post_process: bool = True, vp_stage: int | None = None) -> GPTModel:
"""Builds the model.
If you set the use_legacy_models to True, it will return the legacy GPT model and if not the mcore GPT model.
@@ -1,8 +1,7 @@
import socket
import time
from argparse import Namespace
from collections.abc import Mapping, Sequence
from typing import Callable
from collections.abc import Callable, Mapping, Sequence
import ray
import torch
@@ -1,6 +1,6 @@
from argparse import Namespace
from collections.abc import Mapping, Sequence
from typing import Any, Callable, Tuple
from collections.abc import Callable, Mapping, Sequence
from typing import Any
import ray
import torch
@@ -133,7 +133,7 @@ class UpdateWeightFromTensor:
dist.barrier(group=get_gloo_group())
def _send_hf_params(self, hf_named_tensors) -> Tuple[list[ObjectRef], Any]:
def _send_hf_params(self, hf_named_tensors) -> tuple[list[ObjectRef], Any]:
all_refs = []
refs_colocated, long_lived_tensors = _send_to_colocated_engine(
@@ -166,7 +166,7 @@ def _send_to_colocated_engine(
ipc_gather_src,
ipc_gather_group,
weight_version,
) -> Tuple[list[ObjectRef], Any]:
) -> tuple[list[ObjectRef], Any]:
# TODO improve
long_live_tensors = []
+12 -13
View File
@@ -2,7 +2,6 @@ import dataclasses
import logging
import multiprocessing
import time
from typing import List, Optional
import requests
import sglang_router
@@ -157,7 +156,7 @@ class SGLangEngine(RayActor):
)
response.raise_for_status()
def _make_request(self, endpoint: str, payload: Optional[dict] = None):
def _make_request(self, endpoint: str, payload: dict | None = None):
"""Make a POST request to the specified endpoint with the given payload.
Args:
@@ -203,10 +202,10 @@ class SGLangEngine(RayActor):
def update_weights_from_tensor(
self,
serialized_named_tensors: List[str],
load_format: Optional[str] = None,
serialized_named_tensors: list[str],
load_format: str | None = None,
flush_cache: bool = False,
weight_version: Optional[str] = None,
weight_version: str | None = None,
):
"""
Update model weights from tensor data. The HTTP server will only post meta data, and the real weights will be copied directly from GPUs.
@@ -273,7 +272,7 @@ class SGLangEngine(RayActor):
self.flush_cache()
return self._make_request("release_memory_occupation")
def resume_memory_occupation(self, tags: List[str] = None):
def resume_memory_occupation(self, tags: list[str] = None):
"""
Available tags for multi-stage resume: weights, kv_cache
"""
@@ -311,7 +310,7 @@ class SGLangEngine(RayActor):
pass
def update_weights_from_distributed(
self, names, dtypes, shapes, group_name, flush_cache=False, weight_version: Optional[str] = None
self, names, dtypes, shapes, group_name, flush_cache=False, weight_version: str | None = None
):
payload = {
"names": names,
@@ -340,16 +339,16 @@ class SGLangEngine(RayActor):
def start_profile(
self,
# The output directory
output_dir: Optional[str] = None,
output_dir: str | None = None,
# If set, it profile as many as this number of steps.
# If it is set, profiling is automatically stopped after this step, and
# the caller doesn't need to run stop_profile.
start_step: Optional[int] = None,
num_steps: Optional[int] = None,
activities: Optional[List[str]] = None,
start_step: int | None = None,
num_steps: int | None = None,
activities: list[str] | None = None,
profile_by_stage: bool = False,
with_stack: Optional[bool] = None,
record_shapes: Optional[bool] = None,
with_stack: bool | None = None,
record_shapes: bool | None = None,
):
response = requests.post(
f"http://{self.server_host}:{self.server_port}/start_profile",
+1 -1
View File
@@ -135,7 +135,7 @@ class RayTrainGroup:
def connect(self, critic_group):
return ray.get(
[
actor.connect_actor_critic.remote((critic))
actor.connect_actor_critic.remote(critic)
for actor, critic in zip(self._actor_handlers, critic_group._actor_handlers, strict=False)
]
)
+9 -9
View File
@@ -5,7 +5,7 @@ import random
import time
from glob import glob
from pathlib import Path
from typing import Any, Dict, List, Optional, Union
from typing import Any
import numpy as np
import ray
@@ -129,7 +129,7 @@ class RolloutManager:
def offload(self):
return ray.get([engine.release_memory_occupation.remote() for engine in self.rollout_engines])
def onload(self, tags: List[str] = None):
def onload(self, tags: list[str] = None):
return ray.get([engine.resume_memory_occupation.remote(tags=tags) for engine in self.rollout_engines])
def check_weights(self, action: str):
@@ -184,7 +184,7 @@ class RolloutManager:
torch.save(dict(rollout_id=rollout_id, **dump_data), path)
def _post_process_rewards(self, samples: Union[list[Sample], list[list[Sample]]]):
def _post_process_rewards(self, samples: list[Sample] | list[list[Sample]]):
if self.custom_reward_post_process_func is not None:
return self.custom_reward_post_process_func(self.args, samples)
@@ -211,7 +211,7 @@ class RolloutManager:
return raw_rewards, raw_rewards
def _convert_samples_to_train_data(self, samples: Union[list[Sample], list[list[Sample]]]):
def _convert_samples_to_train_data(self, samples: list[Sample] | list[list[Sample]]):
"""
Convert inference generated samples to training data.
"""
@@ -484,7 +484,7 @@ def _start_router(args):
logger.info(f"Router launched at {args.sglang_router_ip}:{args.sglang_router_port}")
def _log_eval_rollout_data(rollout_id, args, data, extra_metrics: Optional[Dict[str, Any]] = None):
def _log_eval_rollout_data(rollout_id, args, data, extra_metrics: dict[str, Any] | None = None):
log_dict = extra_metrics or {}
for key in data.keys():
rewards = data[key]["rewards"]
@@ -542,12 +542,12 @@ def _compute_metrics_from_samples(args, samples):
return log_dict
def _compute_zero_std_metrics(args, all_samples: List[Sample]):
def _compute_zero_std_metrics(args, all_samples: list[Sample]):
# only compute in GRPO-like algorithms where one prompt has multiple responses
if args.advantage_estimator == "ppo":
return {}
def _is_zero_std(samples: List[Sample]):
def _is_zero_std(samples: list[Sample]):
rewards = [sample.get_reward_value(args) for sample in samples]
return len(rewards) == 0 or all(rewards[0] == r for r in rewards)
@@ -559,7 +559,7 @@ def _compute_zero_std_metrics(args, all_samples: List[Sample]):
return {f"zero_std/count_{reward}": len(items) for reward, items in group_by(interesting_rewards).items()}
def _compute_spec_metrics(args, all_samples: List[Sample]):
def _compute_spec_metrics(args, all_samples: list[Sample]):
if args.sglang_speculative_algorithm is None:
return {}
num_samples = len(all_samples)
@@ -573,7 +573,7 @@ def _compute_spec_metrics(args, all_samples: List[Sample]):
return metrics
def _compute_reward_cat_metrics(args, all_samples: List[Sample]):
def _compute_reward_cat_metrics(args, all_samples: list[Sample]):
reward_cat_key = args.log_reward_category
if reward_cat_key is None:
return {}
+1 -2
View File
@@ -1,8 +1,7 @@
from dataclasses import dataclass
from typing import Optional
@dataclass
class DynamicFilterOutput:
keep: bool
reason: Optional[str] = None
reason: str | None = None
+1 -2
View File
@@ -1,5 +1,4 @@
import asyncio
from typing import Union
import aiohttp
@@ -68,7 +67,7 @@ async def batched_async_rm(
args,
samples: list[Sample],
**kwargs,
) -> list[Union[int, float]]:
) -> list[int | float]:
if args.custom_rm_path is not None:
# Ensure the custom reward function is implemented in batch mode
rm_function = load_function(args.custom_rm_path)
+3 -3
View File
@@ -1,6 +1,6 @@
import re
import string
from typing import Iterable, Optional
from collections.abc import Iterable
DEFAULT_VALID_LETTERS = list(string.ascii_uppercase[:8])
@@ -19,7 +19,7 @@ def _normalize_text(text: str) -> str:
return re.sub(r"[^a-z0-9]+", " ", text.lower()).strip()
def _extract_letter_from_response(response: str, valid_letters: Iterable[str]) -> Optional[str]:
def _extract_letter_from_response(response: str, valid_letters: Iterable[str]) -> str | None:
"""
Best-effort extraction of the selected option letter from the model response.
"""
@@ -51,7 +51,7 @@ def _extract_letter_from_response(response: str, valid_letters: Iterable[str]) -
return None
def compute_gpqa_reward(response: str, label, metadata: Optional[dict] = None) -> float:
def compute_gpqa_reward(response: str, label, metadata: dict | None = None) -> float:
"""Rule-based scorer for GPQA-style multiple-choice evaluation."""
if response is None:
return 0.0
+12 -11
View File
@@ -5,8 +5,9 @@ import logging
import os
import subprocess
import sys
from collections.abc import Sequence
from pathlib import Path
from typing import Any, Dict, List, Optional, Sequence, Union
from typing import Any
logger = logging.getLogger(__name__)
@@ -23,7 +24,7 @@ def _ensure_ifbench_repo() -> Path:
if not repo_path.exists():
clone_cmd = ["git", "clone", "https://github.com/allenai/IFBench.git", str(repo_path)]
try:
subprocess.run(clone_cmd, check=True, stdout=subprocess.PIPE, stderr=subprocess.PIPE)
subprocess.run(clone_cmd, check=True, capture_output=True)
except Exception as exc:
raise ImportError(
"Unable to automatically clone IFBench. Please clone "
@@ -78,14 +79,14 @@ evaluation_lib = _load_evaluation_lib()
InputExample = evaluation_lib.InputExample
JsonDict = Dict[str, Any]
KwargsDict = Dict[str, Optional[Union[str, int, float]]]
JsonDict = dict[str, Any]
KwargsDict = dict[str, str | int | float | None]
def _normalize_instruction_ids(raw_ids: Sequence[Any]) -> List[str]:
def _normalize_instruction_ids(raw_ids: Sequence[Any]) -> list[str]:
"""Ensure instruction identifiers are clean strings."""
normalized: List[str] = []
normalized: list[str] = []
for entry in raw_ids or []:
if entry is None:
continue
@@ -99,11 +100,11 @@ def _normalize_instruction_ids(raw_ids: Sequence[Any]) -> List[str]:
def _coerce_kwargs_list(
raw_kwargs: Any,
num_instructions: int,
) -> List[KwargsDict]:
) -> list[KwargsDict]:
"""Convert stored kwargs into the list structure expected by IFBench."""
if isinstance(raw_kwargs, list):
processed: List[KwargsDict] = []
processed: list[KwargsDict] = []
for entry in raw_kwargs:
if isinstance(entry, dict):
processed.append(dict(entry))
@@ -121,13 +122,13 @@ def _coerce_kwargs_list(
processed = processed[:num_instructions]
# Remove explicit None values to match official preprocessing.
sanitized: List[KwargsDict] = []
sanitized: list[KwargsDict] = []
for entry in processed:
sanitized.append({k: v for k, v in entry.items() if v is not None})
return sanitized
def _build_input_example(metadata: JsonDict) -> Optional[InputExample]:
def _build_input_example(metadata: JsonDict) -> InputExample | None:
instruction_ids = _normalize_instruction_ids(metadata.get("instruction_id_list") or [])
if not instruction_ids:
logger.debug("Missing instruction identifiers in metadata: %s", metadata)
@@ -150,7 +151,7 @@ def _build_input_example(metadata: JsonDict) -> Optional[InputExample]:
)
def compute_ifbench_reward(response: str, label: Any, metadata: Optional[JsonDict] = None) -> float:
def compute_ifbench_reward(response: str, label: Any, metadata: JsonDict | None = None) -> float:
"""Score a model response using the official IFBench rules."""
if metadata is None:
+4 -7
View File
@@ -15,10 +15,9 @@
import re
import signal
from typing import Optional
def last_boxed_only_string(string: str) -> Optional[str]:
def last_boxed_only_string(string: str) -> str | None:
"""Extract the last LaTeX boxed expression from a string.
Args:
@@ -213,9 +212,7 @@ def is_correct_minerva(
return (pred == gt), pred
def is_correct_strict_box(
pred: str, gt: str, pause_tokens_index: Optional[list[int]] = None
) -> tuple[int, Optional[str]]:
def is_correct_strict_box(pred: str, gt: str, pause_tokens_index: list[int] | None = None) -> tuple[int, str | None]:
"""Check if the prediction is correct using strict boxed answer criteria.
Args:
@@ -241,7 +238,7 @@ def is_correct_strict_box(
def verify(
solution_str: str, answer: str, strict_box_verify: bool = False, pause_tokens_index: Optional[list[int]] = None
solution_str: str, answer: str, strict_box_verify: bool = False, pause_tokens_index: list[int] | None = None
) -> bool:
"""Verify if the solution is correct.
@@ -266,7 +263,7 @@ def compute_score(
solution_str: str,
ground_truth: str,
strict_box_verify: bool = False,
pause_tokens_index: Optional[list[int]] = None,
pause_tokens_index: list[int] | None = None,
) -> float:
"""Compute the reward score for a solution.
+2 -3
View File
@@ -5,7 +5,6 @@ Answer checker API that uses sympy to simplify expressions and check for equalit
Call grade_answer(given_answer: str, ground_truth: str).
"""
import re
from typing import Optional
import sympy
from pylatexenc import latex2text
@@ -13,7 +12,7 @@ from sympy.parsing import sympy_parser
# Dan Hendrycks' code
def mathd_normalize_answer(answer: Optional[str]) -> Optional[str]:
def mathd_normalize_answer(answer: str | None) -> str | None:
if answer is None:
return None
answer = answer.strip()
@@ -67,7 +66,7 @@ def _strip_string(string):
try:
a = int(a)
b = int(b)
assert string == "{}/{}".format(a, b)
assert string == f"{a}/{b}"
new_string = "\\frac{" + str(a) + "}{" + str(b) + "}"
return new_string
except Exception:
+6 -5
View File
@@ -5,7 +5,8 @@ import io
import logging
from argparse import Namespace
from collections import defaultdict
from typing import Any, Callable, Optional, Union
from collections.abc import Callable
from typing import Any
import numpy as np
import sglang_router
@@ -214,10 +215,10 @@ async def generate(args: Namespace, sample: Sample, sampling_params: dict[str, A
async def generate_and_rm(
args: Namespace,
sample: Union[Sample, list[Sample]],
sample: Sample | list[Sample],
sampling_params: dict[str, Any],
evaluation: bool = False,
) -> Union[Sample, list[Sample]]:
) -> Sample | list[Sample]:
# For samples with existing response, check if they're complete
if sample.status == Sample.Status.COMPLETED or sample.status == Sample.Status.TRUNCATED:
assert sample.response is not None
@@ -434,7 +435,7 @@ class _MetricGatherer:
def __init__(self):
self._dynamic_filter_drop_reason_count = defaultdict(lambda: 0)
def on_dynamic_filter_drop(self, reason: Optional[str]):
def on_dynamic_filter_drop(self, reason: str | None):
if not reason:
return
self._dynamic_filter_drop_reason_count[reason] += 1
@@ -615,7 +616,7 @@ async def eval_rollout_single_dataset(
# TODO remove this temp function
def generate_rollout(
args: Namespace, rollout_id: int, data_buffer: Any, evaluation: bool = False
) -> Union[RolloutFnTrainOutput, RolloutFnEvalOutput]:
) -> RolloutFnTrainOutput | RolloutFnEvalOutput:
"""An example to implement the generate_rollout function for an rule based rm rollout generation.
Args:
+24 -24
View File
@@ -8,7 +8,7 @@ Optimized for string prefixes with corresponding token IDs.
import threading
import time
from dataclasses import dataclass
from typing import Any, Dict, List, Optional
from typing import Any
@dataclass
@@ -16,11 +16,11 @@ class MatchResult:
"""Result of prefix matching operation."""
matched_prefix: str
token_ids: List[int]
logp: List[float]
loss_mask: List[int] # Added loss mask for model generation parts
token_ids: list[int]
logp: list[float]
loss_mask: list[int] # Added loss mask for model generation parts
remaining_string: str
last_node: "StringTreeNode"
last_node: StringTreeNode
class StringTreeNode:
@@ -28,16 +28,16 @@ class StringTreeNode:
counter = 0
def __init__(self, node_id: Optional[int] = None):
def __init__(self, node_id: int | None = None):
# Core tree structure
self.children: List[StringTreeNode] = [] # Use list to store children
self.parent: Optional[StringTreeNode] = None
self.children: list[StringTreeNode] = [] # Use list to store children
self.parent: StringTreeNode | None = None
# Node data
self.string_key: str = "" # The string fragment this node represents
self.token_ids: Optional[List[int]] = None # Token IDs for this node only (not cumulative)
self.logp: Optional[List[float]] = None # Log probabilities for this node's tokens
self.loss_mask: Optional[List[int]] = None # Loss mask for model generation parts
self.token_ids: list[int] | None = None # Token IDs for this node only (not cumulative)
self.logp: list[float] | None = None # Log probabilities for this node's tokens
self.loss_mask: list[int] | None = None # Loss mask for model generation parts
# Access tracking
self.last_access_time = time.monotonic()
@@ -47,7 +47,7 @@ class StringTreeNode:
self.ref_count = 0
# Weight version tracking
self.weight_version: Optional[int] = None # Weight version for this node
self.weight_version: int | None = None # Weight version for this node
# Node identification
self.id = StringTreeNode.counter if node_id is None else node_id
@@ -201,10 +201,10 @@ class StringRadixTrie:
def insert(
self,
text: str,
token_ids: List[int],
logp: Optional[List[float]] = None,
loss_mask: Optional[List[int]] = None,
weight_version: Optional[int] = None,
token_ids: list[int],
logp: list[float] | None = None,
loss_mask: list[int] | None = None,
weight_version: int | None = None,
) -> bool:
"""
Insert a string and its corresponding token IDs, log probabilities, and loss mask into the trie.
@@ -276,10 +276,10 @@ class StringRadixTrie:
def _insert(
self,
text: str,
token_ids: List[int],
logp: List[float],
loss_mask: List[int],
weight_version: Optional[int] = None,
token_ids: list[int],
logp: list[float],
loss_mask: list[int],
weight_version: int | None = None,
) -> bool:
"""Insert tokens - skip tokens for existing nodes just like we skip text."""
@@ -371,7 +371,7 @@ class StringRadixTrie:
return removed_count > 0
return False
def _find_node_by_text(self, text: str) -> Optional[StringTreeNode]:
def _find_node_by_text(self, text: str) -> StringTreeNode | None:
"""
Find node by exact text match.
Args:
@@ -436,7 +436,7 @@ class StringRadixTrie:
return True
return False
def gc_by_weight_version(self, current_weight_version: Optional[int] = None) -> int:
def gc_by_weight_version(self, current_weight_version: int | None = None) -> int:
"""
Perform garbage collection based on weight version.
Remove nodes with weight_version < (current_weight_version - gc_threshold_k).
@@ -470,7 +470,7 @@ class StringRadixTrie:
return removed_count
def _find_outdated_nodes(self, gc_threshold: int) -> List[StringTreeNode]:
def _find_outdated_nodes(self, gc_threshold: int) -> list[StringTreeNode]:
"""
Find nodes that should be removed based on weight version threshold.
Uses layer-by-layer traversal - if parent is outdated, children are not checked.
@@ -521,7 +521,7 @@ class StringRadixTrie:
# Start validation from the node itself
validate_recursive(node, node.weight_version)
def get_stats(self) -> Dict[str, Any]:
def get_stats(self) -> dict[str, Any]:
"""Get cache statistics."""
with self._lock:
total_requests = self.cache_hits + self.cache_misses
+4 -4
View File
@@ -2,7 +2,7 @@ import argparse
import json
import logging
import os
from typing import Any, Dict
from typing import Any
import yaml
from sglang_router.launch_router import RouterArgs
@@ -1282,7 +1282,7 @@ def _resolve_eval_datasets(args) -> list[EvalDatasetConfig]:
Build evaluation dataset configurations from either --eval-config or --eval-prompt-data.
"""
datasets_config = []
defaults: Dict[str, Any] = {}
defaults: dict[str, Any] = {}
if args.eval_config:
from omegaconf import OmegaConf
@@ -1498,7 +1498,7 @@ def miles_validate_args(args):
args.use_routing_replay = True
if args.custom_config_path:
with open(args.custom_config_path, "r") as f:
with open(args.custom_config_path) as f:
data = yaml.safe_load(f) or {}
for k, v in data.items():
if hasattr(args, k):
@@ -1554,7 +1554,7 @@ def hf_validate_args(args, hf_config):
raise AssertionError("hf_validate_args failed: " + "; ".join(errors))
def _validate_and_update_megatron_args_from_hf(args, args_from_hf_config: Dict[str, Any]):
def _validate_and_update_megatron_args_from_hf(args, args_from_hf_config: dict[str, Any]):
for key, value in args_from_hf_config.items():
if hasattr(args, key) and getattr(args, key) != value:
raise ValueError(
@@ -1,7 +1,7 @@
import json
from pathlib import Path
from types import SimpleNamespace
from typing import Annotated, List
from typing import Annotated
import torch
import typer
@@ -27,7 +27,7 @@ def main(
load_debug_rollout_data: Annotated[str, typer.Option()],
show_metrics: bool = True,
show_samples: bool = True,
category: List[str] = None,
category: list[str] = None,
):
if category is None:
category = ["train", "eval"]
@@ -56,7 +56,7 @@ def main(
print(json.dumps({k: v for k, v in sample.items() if k in _WHITELIST_KEYS}))
def _get_rollout_dump_paths(load_debug_rollout_data: str, categories: List[str]):
def _get_rollout_dump_paths(load_debug_rollout_data: str, categories: list[str]):
# may improve later
for rollout_id in range(1000):
for category in categories:
+7 -7
View File
@@ -1,5 +1,5 @@
from datetime import timedelta
from typing import Any, Optional, Union
from typing import Any
import torch
import torch.distributed as dist
@@ -36,14 +36,14 @@ def get_gloo_group():
# Copy from pytorch to allow creating multiple main groups.
# https://github.com/pytorch/pytorch/blob/main/torch/distributed/distributed_c10d.py
def init_process_group(
backend: Union[str, Backend] = None,
init_method: Optional[str] = None,
timeout: Optional[timedelta] = None,
backend: str | Backend = None,
init_method: str | None = None,
timeout: timedelta | None = None,
world_size: int = -1,
rank: int = -1,
store: Optional[Store] = None,
store: Store | None = None,
group_name: str = None,
pg_options: Optional[Any] = None,
pg_options: Any | None = None,
):
assert (store is None) or (init_method is None), "Cannot specify both init_method and store."
@@ -94,7 +94,7 @@ def init_process_group(
def distributed_masked_whiten(
values: torch.Tensor,
mask: torch.Tensor,
process_group: Optional[dist.ProcessGroup] = None,
process_group: dist.ProcessGroup | None = None,
shift_mean: bool = True,
epsilon: float = 1e-8,
):
+25 -24
View File
@@ -1,14 +1,15 @@
from __future__ import annotations
import copy
from typing import Any, Dict, Iterable, List, Optional, Sequence, Tuple
from collections.abc import Iterable, Sequence
from typing import Any
from pydantic import BaseModel, ConfigDict, Field, field_validator
_EMPTY_VALUES = (None, [], {})
def _ensure_metadata_overrides(value: Any) -> Dict[str, Any]:
def _ensure_metadata_overrides(value: Any) -> dict[str, Any]:
if value is None:
return {}
if not isinstance(value, dict):
@@ -21,34 +22,34 @@ class EvalDatasetConfig(BaseModel):
name: str
path: str
rm_type: Optional[str] = None
rm_type: str | None = None
# Dataset-specific overrides
prompt_key: Optional[str] = None
label_key: Optional[str] = None
tool_key: Optional[str] = None
metadata_key: Optional[str] = None
prompt_key: str | None = None
label_key: str | None = None
tool_key: str | None = None
metadata_key: str | None = None
n_samples_per_eval_prompt: Optional[int] = None
n_samples_per_eval_prompt: int | None = None
temperature: Optional[float] = None
top_p: Optional[float] = None
top_k: Optional[int] = None
max_response_len: Optional[int] = None
min_new_tokens: Optional[int] = None
temperature: float | None = None
top_p: float | None = None
top_k: int | None = None
max_response_len: int | None = None
min_new_tokens: int | None = None
stop: Optional[Sequence[str]] = None
stop_token_ids: Optional[Sequence[int]] = None
stop: Sequence[str] | None = None
stop_token_ids: Sequence[int] | None = None
metadata_overrides: Dict[str, Any] = Field(default_factory=dict)
metadata_overrides: dict[str, Any] = Field(default_factory=dict)
model_config = ConfigDict(validate_assignment=True, extra="forbid")
@field_validator("metadata_overrides", mode="before")
def _validate_metadata_overrides(cls, value: Any) -> Dict[str, Any]:
def _validate_metadata_overrides(cls, value: Any) -> dict[str, Any]:
return _ensure_metadata_overrides(value)
def apply_defaults(self, defaults: Dict[str, Any]) -> None:
def apply_defaults(self, defaults: dict[str, Any]) -> None:
for key, value in defaults.items():
if not hasattr(self, key):
continue
@@ -60,7 +61,7 @@ class EvalDatasetConfig(BaseModel):
setattr(self, key, value)
@property
def cache_key(self) -> Tuple[Any, ...]:
def cache_key(self) -> tuple[Any, ...]:
"""Return a tuple uniquely identifying dataset config for caching."""
return (
self.name,
@@ -71,7 +72,7 @@ class EvalDatasetConfig(BaseModel):
self.metadata_key,
)
def inject_metadata(self, sample_metadata: Any) -> Dict[str, Any]:
def inject_metadata(self, sample_metadata: Any) -> dict[str, Any]:
"""Return updated metadata merging overrides."""
if not isinstance(sample_metadata, dict):
metadata = {}
@@ -87,7 +88,7 @@ class EvalDatasetConfig(BaseModel):
return metadata
def ensure_dataset_list(config: Any) -> List[Dict[str, Any]]:
def ensure_dataset_list(config: Any) -> list[dict[str, Any]]:
"""
Normalize OmegaConf containers into a list of dicts.
Accepts either a list or dictionary keyed by dataset name.
@@ -116,9 +117,9 @@ def ensure_dataset_list(config: Any) -> List[Dict[str, Any]]:
def build_eval_dataset_configs(
raw_config: Iterable[Dict[str, Any]], defaults: Dict[str, Any]
) -> List[EvalDatasetConfig]:
datasets: List[EvalDatasetConfig] = []
raw_config: Iterable[dict[str, Any]], defaults: dict[str, Any]
) -> list[EvalDatasetConfig]:
datasets: list[EvalDatasetConfig] = []
for cfg in raw_config:
dataset = EvalDatasetConfig(**cfg)
dataset.apply_defaults(defaults)
+4 -5
View File
@@ -9,7 +9,6 @@ import random
import time
from dataclasses import dataclass
from pathlib import Path
from typing import Optional
from miles.utils.misc import exec_command
from miles.utils.typer_utils import dataclass_cli
@@ -26,7 +25,7 @@ def convert_checkpoint(
multinode: bool = False,
extra_args: str = "",
dir_dst: str = "/root",
hf_checkpoint: Optional[str] = None,
hf_checkpoint: str | None = None,
):
hf_checkpoint = hf_checkpoint or f"/root/models/{model_name}"
@@ -94,11 +93,11 @@ class ExecuteTrainConfig:
def execute_train(
train_args: str,
num_gpus_per_node: int,
megatron_model_type: Optional[str],
megatron_model_type: str | None,
train_script: str = "train.py",
before_ray_job_submit=None,
extra_env_vars=None,
config: Optional[ExecuteTrainConfig] = None,
config: ExecuteTrainConfig | None = None,
):
if extra_env_vars is None:
extra_env_vars = {}
@@ -198,7 +197,7 @@ def check_has_nvlink():
return int(output) > 0
def get_default_wandb_args(test_file: str, run_name_prefix: Optional[str] = None, run_id: Optional[str] = None):
def get_default_wandb_args(test_file: str, run_name_prefix: str | None = None, run_id: str | None = None):
if not os.environ.get("WANDB_API_KEY"):
print("Skip wandb configuration since WANDB_API_KEY is not found")
return ""
+1 -3
View File
@@ -1,5 +1,3 @@
from typing import Tuple
import torch
import triton
import triton.language as tl
@@ -60,7 +58,7 @@ def _blockwise_cast_to_fp8_triton(
tl.store(S + pid_m * stride_sm + pid_n * stride_sn, x_s)
def blockwise_cast_to_fp8_triton(x: torch.Tensor, block_size=None) -> Tuple[torch.Tensor, torch.Tensor]:
def blockwise_cast_to_fp8_triton(x: torch.Tensor, block_size=None) -> tuple[torch.Tensor, torch.Tensor]:
BLOCK_M, BLOCK_N = 128, 128
if block_size:
BLOCK_M, BLOCK_N = block_size[0], block_size[1]
+2 -3
View File
@@ -6,7 +6,6 @@ import multiprocessing
import os
import random
import socket
from typing import Optional
import httpx
@@ -34,7 +33,7 @@ def is_port_available(port):
s.bind(("", port))
s.listen(1)
return True
except socket.error:
except OSError:
return False
except OverflowError:
return False
@@ -115,7 +114,7 @@ def terminate_process(process: multiprocessing.Process, timeout: float = 1.0) ->
process.join()
_http_client: Optional[httpx.AsyncClient] = None
_http_client: httpx.AsyncClient | None = None
_client_concurrency: int = 0
# Optional Ray-based distributed POST dispatch
+4 -3
View File
@@ -1,5 +1,6 @@
from collections import defaultdict
from typing import Any, Callable, Iterable, List, Tuple
from collections.abc import Callable, Iterable
from typing import Any
import torch
@@ -14,7 +15,7 @@ def group_by(iterable, key=None):
# TODO fsdp can also use this
def chunk_named_params_by_size(named_params: Iterable[Tuple[str, torch.Tensor]], chunk_size: int):
def chunk_named_params_by_size(named_params: Iterable[tuple[str, torch.Tensor]], chunk_size: int):
return _chunk_by_size(
named_params,
compute_size=lambda named_weight: named_weight[1].nbytes,
@@ -23,7 +24,7 @@ def chunk_named_params_by_size(named_params: Iterable[Tuple[str, torch.Tensor]],
def _chunk_by_size(objects: Iterable[Any], compute_size: Callable[[Any], int], chunk_size: int):
bucket: List[Any] = []
bucket: list[Any] = []
bucket_size = 0
for obj in objects:
+8 -10
View File
@@ -1,9 +1,7 @@
from typing import Dict, List, Tuple
from transformers import AutoTokenizer
def get_response_lengths(loss_masks: List[List[int]]) -> List[int]:
def get_response_lengths(loss_masks: list[list[int]]) -> list[int]:
return [len(mask[mask.index(1) :]) if 1 in mask else 0 for mask in loss_masks]
@@ -13,7 +11,7 @@ class MultiTurnLossMaskGenerator:
self.system_message_length, self.gen_token_length = self.get_system_message_length()
self.tokenizer_type = tokenizer_type
def get_response_lengths(self, loss_masks: List[List[int]]) -> List[int]:
def get_response_lengths(self, loss_masks: list[list[int]]) -> list[int]:
return get_response_lengths(loss_masks)
def find_all_sublist_indices(self, main_list, sublist):
@@ -24,7 +22,7 @@ class MultiTurnLossMaskGenerator:
indices.append(i)
return indices
def get_system_message_length(self) -> Tuple[int, int]:
def get_system_message_length(self) -> tuple[int, int]:
test_string = "FOR TESTING ONLY"
test_messages = [
{"role": "user", "content": test_string},
@@ -46,7 +44,7 @@ class MultiTurnLossMaskGenerator:
system_message_length = idx_1 - ((idx_2 - idx_1) - end_interval - len(raw_token_ids))
return system_message_length, gen_token_length
def gen_multi_turn_loss_mask_qwen(self, messages: List[Dict]) -> Tuple[List[int], List[int]]:
def gen_multi_turn_loss_mask_qwen(self, messages: list[dict]) -> tuple[list[int], list[int]]:
all_loss_masks = []
all_token_ids = []
@@ -69,7 +67,7 @@ class MultiTurnLossMaskGenerator:
return all_token_ids, all_loss_masks
def gen_multi_turn_loss_mask_qwen3(self, messages: List[Dict]) -> Tuple[List[int], List[int]]:
def gen_multi_turn_loss_mask_qwen3(self, messages: list[dict]) -> tuple[list[int], list[int]]:
all_loss_masks = []
all_token_ids = []
@@ -96,7 +94,7 @@ class MultiTurnLossMaskGenerator:
return all_token_ids, all_loss_masks
def gen_multi_turn_loss_mask_distill_qwen(self, messages: List[Dict]) -> Tuple[List[int], List[int]]:
def gen_multi_turn_loss_mask_distill_qwen(self, messages: list[dict]) -> tuple[list[int], list[int]]:
prompt = self.tokenizer.apply_chat_template(messages[:1], tokenize=False, add_generation_prompt=True)
response = messages[-1]["content"]
prompt_tokens = self.tokenizer(prompt, add_special_tokens=False)["input_ids"]
@@ -110,7 +108,7 @@ class MultiTurnLossMaskGenerator:
loss_mask = [0] * len(token_ids)
return token_ids, loss_mask
def get_loss_mask(self, messages: List[Dict]) -> List[int]:
def get_loss_mask(self, messages: list[dict]) -> list[int]:
if self.tokenizer_type == "qwen":
if "<|Assistant|>" in self.tokenizer.get_added_vocab():
return self.gen_multi_turn_loss_mask_distill_qwen(messages)
@@ -123,7 +121,7 @@ class MultiTurnLossMaskGenerator:
else:
raise ValueError(f"Unsupported tokenizer type: {self.tokenizer_type}")
def get_text_from_loss_mask(self, token_ids: List[int], loss_masks: List[int]) -> List[str]:
def get_text_from_loss_mask(self, token_ids: list[int], loss_masks: list[int]) -> list[str]:
selected_texts = []
current_tokens = []
+1 -2
View File
@@ -1,5 +1,4 @@
import logging
from typing import Dict
logger = logging.getLogger(__name__)
@@ -15,7 +14,7 @@ class MetricChecker:
self.args = args
self._exists_check_success = False
def on_eval(self, metrics: Dict[str, float]):
def on_eval(self, metrics: dict[str, float]):
actual_value = metrics.get(self.args.ci_metric_checker_key)
assert actual_value is not None, f"{metrics=} {self.args.ci_metric_checker_key=}"
+7 -7
View File
@@ -1,17 +1,17 @@
import math
from typing import Any, Dict, List, Literal, Optional, Tuple, Union
from typing import Any, Literal
import numpy as np
def dict_add_prefix(d: Dict[str, Any], prefix: str) -> Dict[str, Any]:
def dict_add_prefix(d: dict[str, Any], prefix: str) -> dict[str, Any]:
return {f"{prefix}{k}": v for k, v in d.items()}
def compute_pass_rate(
flat_rewards: List[float],
flat_rewards: list[float],
group_size: int,
num_groups: Optional[int] = None,
num_groups: int | None = None,
):
if group_size == 1:
return {}
@@ -53,7 +53,7 @@ def _estimate_pass_at_k(num_samples, num_correct, k):
return np.array([estimator(int(n), int(c), k) for n, c in zip(num_samples, num_correct, strict=False)])
def compute_statistics(values: List[float]) -> Dict[str, float]:
def compute_statistics(values: list[float]) -> dict[str, float]:
values = np.array(values)
return {
"mean": np.mean(values).item(),
@@ -62,12 +62,12 @@ def compute_statistics(values: List[float]) -> Dict[str, float]:
def compression_ratio(
data: Union[str, bytes],
data: str | bytes,
*,
encoding: str = "utf-8",
algorithm: Literal["zlib", "gzip", "bz2", "lzma"] = "zlib",
level: int = 9,
) -> Tuple[float, float]:
) -> tuple[float, float]:
if isinstance(data, str):
raw = data.encode(encoding)
else:
+1 -2
View File
@@ -1,6 +1,5 @@
import importlib
import subprocess
from typing import Optional
import ray
@@ -32,7 +31,7 @@ class SingletonMeta(type):
return cls._instances[cls]
def exec_command(cmd: str, capture_output: bool = False) -> Optional[str]:
def exec_command(cmd: str, capture_output: bool = False) -> str | None:
print(f"EXEC: {cmd}", flush=True)
try:
+11 -12
View File
@@ -1,6 +1,5 @@
# Adapt from https://github.com/OpenRLHF/OpenRLHF/blob/10c733694ed9fbb78a0a2ff6a05efc7401584d46/openrlhf/models/utils.py
# and https://github.com/OpenRLHF/OpenRLHF/blob/10c733694ed9fbb78a0a2ff6a05efc7401584d46/openrlhf/trainer/ppo_utils/experience_maker.py
from typing import List, Optional, Tuple
import torch
import torch.distributed as dist
@@ -52,7 +51,7 @@ def compute_policy_loss(
advantages: torch.Tensor,
eps_clip: float,
eps_clip_high: float,
eps_clip_c: Optional[float] = None,
eps_clip_c: float | None = None,
):
ratio = (-ppo_kl).exp()
pg_losses1 = -ratio * advantages
@@ -73,7 +72,7 @@ def compute_policy_loss(
return pg_losses, clipfrac
def compute_log_probs(logits: torch.Tensor, tokens: torch.Tensor, process_group: Optional[dist.ProcessGroup]):
def compute_log_probs(logits: torch.Tensor, tokens: torch.Tensor, process_group: dist.ProcessGroup | None):
from megatron.core.fusions.fused_cross_entropy import fused_vocab_parallel_cross_entropy
# convert to [seq_len, batch_size, vocab_size] as expected by fused_vocab_parallel_cross_entropy
@@ -134,13 +133,13 @@ def get_grpo_returns(
def get_reinforce_plus_plus_returns(
rewards: torch.Tensor,
kl: List[torch.Tensor],
loss_masks: List[torch.Tensor],
response_lengths: List[int],
total_lengths: List[int],
kl: list[torch.Tensor],
loss_masks: list[torch.Tensor],
response_lengths: list[int],
total_lengths: list[int],
kl_coef: float,
gamma: float,
) -> List[torch.Tensor]:
) -> list[torch.Tensor]:
"""
Calculates discounted returns for REINFORCE++ (https://arxiv.org/pdf/2501.03262)
@@ -203,10 +202,10 @@ def get_reinforce_plus_plus_returns(
def get_reinforce_plus_plus_baseline_advantages(
rewards: torch.Tensor,
kl: List[torch.Tensor],
loss_masks: List[torch.Tensor],
kl: list[torch.Tensor],
loss_masks: list[torch.Tensor],
kl_coef: float,
) -> List[torch.Tensor]:
) -> list[torch.Tensor]:
"""
Calculates the unwhitened advantages for the REINFORCE++-baseline algorithm.
Broadcasting the scalar (reward - group_baseline) to each token.
@@ -238,7 +237,7 @@ def get_advantages_and_returns(
rewards: torch.Tensor,
gamma: float,
lambd: float,
) -> Tuple[torch.Tensor, torch.Tensor]:
) -> tuple[torch.Tensor, torch.Tensor]:
"""Function that computes advantages and returns from rewards and values.
Calculated as in the original PPO paper: https://arxiv.org/abs/1707.06347
Note that rewards may include a KL divergence loss term.
+4 -5
View File
@@ -15,10 +15,9 @@
import copy
import heapq
from typing import List, Tuple
def karmarkar_karp(seqlen_list: List[int], k_partitions: int, equal_size: bool):
def karmarkar_karp(seqlen_list: list[int], k_partitions: int, equal_size: bool):
# see: https://en.wikipedia.org/wiki/Largest_differencing_method
class Set:
@@ -44,7 +43,7 @@ def karmarkar_karp(seqlen_list: List[int], k_partitions: int, equal_size: bool):
class State:
def __init__(self, items: List[Tuple[int, int]], k: int) -> None:
def __init__(self, items: list[tuple[int, int]], k: int) -> None:
self.k = k
# sets should always be decreasing order
self.sets = [Set() for _ in range(k)]
@@ -124,7 +123,7 @@ def karmarkar_karp(seqlen_list: List[int], k_partitions: int, equal_size: bool):
return partitions
def greedy_partition(seqlen_list: List[int], k_partitions: int, equal_size: bool):
def greedy_partition(seqlen_list: list[int], k_partitions: int, equal_size: bool):
bias = sum(seqlen_list) + 1 if equal_size else 0
sorted_seqlen = [(seqlen + bias, i) for i, seqlen in enumerate(seqlen_list)]
partitions = [[] for _ in range(k_partitions)]
@@ -144,7 +143,7 @@ def greedy_partition(seqlen_list: List[int], k_partitions: int, equal_size: bool
return partitions
def get_seqlen_balanced_partitions(seqlen_list: List[int], k_partitions: int, equal_size: bool):
def get_seqlen_balanced_partitions(seqlen_list: list[int], k_partitions: int, equal_size: bool):
"""get order of seq lengths to make partitions balanced, this is
used in balacing sum of seqlength across dp ranks and microbatches
Parameters:
+4 -4
View File
@@ -1,10 +1,10 @@
from abc import ABC, abstractmethod
from collections import defaultdict
from typing import Callable, Dict, Iterable, Tuple
from collections.abc import Callable, Iterable
import torch
_SourceGetter = Callable[[], Iterable[Tuple[str, torch.Tensor]]]
_SourceGetter = Callable[[], Iterable[tuple[str, torch.Tensor]]]
class TensorBackuper(ABC):
@@ -43,7 +43,7 @@ class TensorBackuper(ABC):
class _TensorBackuperNormal(TensorBackuper):
def __init__(self, source_getter):
super().__init__(source_getter=source_getter)
self._backups: Dict[str, Dict[str, torch.Tensor]] = defaultdict(dict)
self._backups: dict[str, dict[str, torch.Tensor]] = defaultdict(dict)
@property
def backup_tags(self):
@@ -103,7 +103,7 @@ class _TensorBackuperNoop(TensorBackuper):
torch.cuda.synchronize()
def _compute_hash_dict(tensors: Dict[str, torch.Tensor]):
def _compute_hash_dict(tensors: dict[str, torch.Tensor]):
return {k: _compute_hash_tensor(v) for k, v in tensors.items()}
+2 -1
View File
@@ -1,7 +1,8 @@
import logging
from argparse import Namespace
from collections.abc import Callable
from copy import deepcopy
from typing import Callable
from miles.utils import tracking_utils
from miles.utils.metric_utils import compute_rollout_step
from miles.utils.timer import Timer
+10 -10
View File
@@ -1,6 +1,6 @@
from dataclasses import dataclass, field
from enum import Enum
from typing import Any, Optional, Union
from typing import Any
import torch
@@ -9,20 +9,20 @@ import torch
class Sample:
"""The sample generated"""
group_index: Optional[int] = None
index: Optional[int] = None
group_index: int | None = None
index: int | None = None
# prompt
prompt: Union[str, list[dict[str, str]]] = ""
prompt: str | list[dict[str, str]] = ""
tokens: list[int] = field(default_factory=list)
# response
response: str = ""
response_length: int = 0
label: Optional[str] = None
reward: Optional[Union[float, dict[str, Any]]] = None
loss_mask: Optional[list[int]] = None
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: Optional[list[float]] = None # Log probabilities from rollout engine
rollout_routed_experts: Optional[list[list[int]]] = None # Routed experts from rollout engine
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):
@@ -35,7 +35,7 @@ class Sample:
metadata: dict = field(default_factory=dict)
# metadata used during training, e.g., what loss to use for this sample.
train_metadata: Optional[dict] = None
train_metadata: dict | None = None
class SpecInfo:
spec_accept_token_num: int = 0
+10 -11
View File
@@ -1,5 +1,4 @@
from abc import ABC, abstractmethod
from typing import Optional, Tuple, Union
import torch
import torch.distributed as dist
@@ -39,17 +38,17 @@ class HuggingfaceAttention(MegatronModule, ABC):
self,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor,
key_value_states: Optional[torch.Tensor] = None,
inference_context: Optional[BaseInferenceContext] = None,
rotary_pos_emb: Optional[Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]] = None,
rotary_pos_cos: Optional[torch.Tensor] = None,
rotary_pos_sin: Optional[torch.Tensor] = None,
attention_bias: Optional[torch.Tensor] = None,
packed_seq_params: Optional[PackedSeqParams] = None,
sequence_len_offset: Optional[int] = None,
key_value_states: torch.Tensor | None = None,
inference_context: BaseInferenceContext | None = None,
rotary_pos_emb: torch.Tensor | tuple[torch.Tensor, torch.Tensor] | None = None,
rotary_pos_cos: torch.Tensor | None = None,
rotary_pos_sin: torch.Tensor | None = None,
attention_bias: torch.Tensor | None = None,
packed_seq_params: PackedSeqParams | None = None,
sequence_len_offset: int | None = None,
*,
inference_params: Optional[BaseInferenceContext] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
inference_params: BaseInferenceContext | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
assert packed_seq_params is not None
cu_seqlens = packed_seq_params.cu_seqlens_q
+1 -1
View File
@@ -28,7 +28,7 @@ select = [
"E", # Pycodestyle Errors (Structural/Fundamental Errors like bad indentation)
"F", # Pyflakes (Core Errors: Unused imports, undefined names)
"B", # Flake8-Bugbear (Logic Bugs: Variable shadowing, dangerous default arguments)
# "UP", # pyupgrade (Modernization and compatibility issues) # TODO
"UP", # pyupgrade (Modernization and compatibility issues)
]
ignore = [
"E402", # module-import-not-at-top-of-file
+3 -3
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass
from typing import Literal, Optional
from typing import Literal
import typer
@@ -12,7 +12,7 @@ class ScriptArgs(U.ExecuteTrainConfig):
use_ref: bool = False
colocate: bool = True
model_name: str = "Qwen3-4B-Instruct-2507"
num_gpus_per_node: Optional[int] = None
num_gpus_per_node: int | None = None
hardware: Literal["H100", "GB300"] = "H100"
mode: Literal["normal", "debug_minimal"] = "normal"
run_id: str = U.create_run_id()
@@ -20,7 +20,7 @@ class ScriptArgs(U.ExecuteTrainConfig):
true_on_policy: bool = False
dynamic_sampling: bool = False
enable_eval: bool = True
megatron_model_type: Optional[str] = None
megatron_model_type: str | None = None
extra_args: str = ""
def __post_init__(self):
+2 -2
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass
from typing import Literal, Optional
from typing import Literal
import typer
@@ -12,7 +12,7 @@ class ScriptArgs(U.ExecuteTrainConfig):
run_id: str = U.create_run_id()
model_name: str = "Qwen3-30B-A3B"
megatron_model_type: str = "qwen3-30B-A3B"
num_gpus_per_node: Optional[int] = None
num_gpus_per_node: int | None = None
hardware: Literal["H100", "GB200", "GB300"] = "H100"
enable_eval: bool = True
extra_args: str = ""
+3 -3
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass
from typing import Literal, Optional
from typing import Literal
import typer
@@ -11,8 +11,8 @@ class ScriptArgs(U.ExecuteTrainConfig):
mode: Literal["normal", "debug_minimal"] = "normal"
run_id: str = U.create_run_id()
model_name: str = "Qwen3-4B"
megatron_model_type: Optional[str] = None
num_gpus_per_node: Optional[int] = None
megatron_model_type: str | None = None
num_gpus_per_node: int | None = None
hardware: Literal["H100", "GB200", "GB300"] = "H100"
extra_args: str = ""
multi_eval: bool = False
+1 -1
View File
@@ -6,7 +6,7 @@ from wheel.bdist_wheel import bdist_wheel as _bdist_wheel
def _fetch_requirements(path):
with open(path, "r") as fd:
with open(path) as fd:
return [r.strip() for r in fd.readlines() if r.strip() and not r.startswith("#")]
+3 -4
View File
@@ -5,7 +5,6 @@ import os
import random
import sys
import time
from typing import List
SLEEP_BACKOFF = 5.0
@@ -92,7 +91,7 @@ def _try_acquire(args):
return _try_acquire_count(args.count, args.total_gpus, args.lock_path_pattern, args.timeout)
def _try_acquire_specific(devs: List[int], path_pattern: str, timeout: int):
def _try_acquire_specific(devs: list[int], path_pattern: str, timeout: int):
fd_locks = []
start = time.time()
try:
@@ -121,7 +120,7 @@ def _try_acquire_count(count: int, total_gpus: int, path_pattern: str, timeout:
start = time.time()
_ensure_lock_files(path_pattern, total_gpus)
while True:
fd_locks: List = []
fd_locks: list = []
for gpu_id in range(total_gpus):
fd_lock = FdLock(path_pattern, gpu_id=gpu_id)
fd_lock.open()
@@ -188,7 +187,7 @@ def _get_lock_path(path_pattern: str, gpu_id: int) -> str:
return path_pattern.format(gpu_id=gpu_id)
def _parse_devices(s: str) -> List[int]:
def _parse_devices(s: str) -> list[int]:
return [int(x) for x in s.split(",") if x.strip() != ""]
+2 -3
View File
@@ -28,7 +28,6 @@ import json
import os
import shutil
from collections import defaultdict
from typing import Optional, Tuple
import torch
from compressed_tensors.compressors import unpack_from_int32
@@ -36,10 +35,10 @@ from safetensors.torch import safe_open, save_file
from tqdm import tqdm
def _load_config(model_dir: str, config_path: Optional[str]) -> Tuple[int, int, int]:
def _load_config(model_dir: str, config_path: str | None) -> tuple[int, int, int]:
"""Read config.json and return hidden_size, inter_size, and group_size."""
cfg_path = config_path or os.path.join(model_dir, "config.json")
with open(cfg_path, "r") as f:
with open(cfg_path) as f:
cfg = json.load(f)
hidden_size = int(cfg.get("hidden_size"))
inter_size = int(cfg.get("moe_intermediate_size"))
+1 -2
View File
@@ -6,7 +6,6 @@ import re
import shutil
import time
from typing import Optional
import safetensors.torch
import torch
@@ -51,7 +50,7 @@ class EmptyStateDictLoadPlanner(dist_cp.default_planner.DefaultLoadPlanner):
def set_up_planner(
self,
state_dict: dist_cp.metadata.STATE_DICT_TYPE,
metadata: Optional[dist_cp.metadata.Metadata] = None,
metadata: dist_cp.metadata.Metadata | None = None,
is_coordinator: bool = False,
) -> None:
for k, v in metadata.state_dict_metadata.items():
+1 -1
View File
@@ -47,7 +47,7 @@ def main(fp8_path, bf16_path):
os.system("cp -rf " + fp8_path + "/tokenizer* " + bf16_path)
os.system("cp -rf " + fp8_path + "/chat_template* " + bf16_path)
model_index_file = os.path.join(fp8_path, "model.safetensors.index.json")
with open(model_index_file, "r") as f:
with open(model_index_file) as f:
model_index = json.load(f)
weight_map = model_index["weight_map"]