mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
[4/N] Tiny enable UP ruleset in Ruff (#287)
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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 = [
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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,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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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())
|
||||
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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`
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 = []
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
]
|
||||
)
|
||||
|
||||
@@ -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,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,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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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,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]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,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=}"
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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.
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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()}
|
||||
|
||||
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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 = ""
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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("#")]
|
||||
|
||||
|
||||
|
||||
@@ -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() != ""]
|
||||
|
||||
|
||||
|
||||
@@ -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"))
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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"]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user