Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion .github/workflows/format-check.yml
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@ jobs:
- name: Set up Python
uses: actions/setup-python@v4
with:
python-version: '3.10'
python-version: '3.12'

- name: Install Python dependencies
run: |
Expand Down
8 changes: 4 additions & 4 deletions areal/api/cli_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@
from dataclasses import asdict, dataclass, field, fields
from enum import Enum
from pathlib import Path
from typing import TYPE_CHECKING, Any, ClassVar
from typing import TYPE_CHECKING, Any, ClassVar, TypeVar

import uvloop
import yaml
Expand All @@ -28,6 +28,8 @@

logger = logging.getLogger("CLIArgs")

ConfigT = TypeVar("ConfigT")


@dataclass
class NormConfig:
Expand Down Expand Up @@ -2246,9 +2248,7 @@ def to_structured_cfg(cfg, config_cls):
return cfg


def load_expr_config[ConfigT](
argv: list[str], config_cls: type[ConfigT]
) -> tuple[ConfigT, str]:
def load_expr_config(argv: list[str], config_cls: type[ConfigT]) -> tuple[ConfigT, str]:
cfg, config_file = parse_cli_args(argv)
cfg = to_structured_cfg(cfg, config_cls=config_cls)
cfg = OmegaConf.to_object(cfg)
Expand Down
13 changes: 8 additions & 5 deletions areal/infra/async_task_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,10 +14,13 @@
import time
from collections.abc import Awaitable, Callable, Coroutine
from dataclasses import dataclass
from typing import Any, cast
from typing import Any, Generic, TypeVar, cast

import uvloop

# Type variable for generic result types
T = TypeVar("T")

# Polling configuration
DEFAULT_POLL_WAIT_TIME = 0.05 # 50ms
DEFAULT_POLL_SLEEP_TIME = 0.5 # 500ms
Expand All @@ -28,7 +31,7 @@ class TaskQueueFullError(RuntimeError):


@dataclass
class TimedResult[T]:
class TimedResult(Generic[T]):
"""Wrapper for task results with creation timestamp.

Attributes
Expand All @@ -47,7 +50,7 @@ class TimedResult[T]:


@dataclass
class _TaskInput[T]:
class _TaskInput(Generic[T]):
"""Internal wrapper for task input with async function and arguments."""

async_fn: Callable[..., Awaitable[T]]
Expand All @@ -57,15 +60,15 @@ class _TaskInput[T]:


@dataclass
class _Task[T]:
class _Task(Generic[T]):
"""Internal wrapper for running task with metadata."""

create_time: int # nanoseconds from time.monotonic_ns()
task: asyncio.Task
task_input: _TaskInput[T]


class AsyncTaskRunner[T]:
class AsyncTaskRunner(Generic[T]):
"""Generic asynchronous task runner with queue management and pause/resume control.

This class provides a reusable async task executor that runs a background thread
Expand Down
14 changes: 10 additions & 4 deletions areal/infra/workflow_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@
import time
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Protocol
from typing import TYPE_CHECKING, Any, Generic, TypeVar, Protocol
from collections.abc import Generator
from collections import deque
import torch
Expand Down Expand Up @@ -250,7 +250,11 @@ class WithTaskID(Protocol):
task_id: int


class BatchTaskDispatcher[TInput: WithTaskID, TResult]:
TInput = TypeVar("TInput", bound=WithTaskID)
TResult = TypeVar("TResult")


class BatchTaskDispatcher(Generic[TInput, TResult]):
"""Generic dispatcher for asynchronous task execution with staleness control.

Manages background threads for task submission and result collection.
Expand Down Expand Up @@ -368,8 +372,10 @@ def _commit_loop(self) -> None:
with self._input_cv:
self._pending_inputs.appendleft(task_input)
self._input_cv.wait_for(
lambda: self._shutdown_event.is_set()
or self._has_runner_capacity()
lambda: (
self._shutdown_event.is_set()
or self._has_runner_capacity()
)
)
# Allow other threads to make progress before retrying
continue
Expand Down
9 changes: 6 additions & 3 deletions areal/models/mcore/common.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
from typing import TypeVar

import torch
import torch.nn.functional as F
from megatron.core.transformer import TransformerConfig
Expand All @@ -8,6 +10,9 @@
logger = logging.getLogger("MCoreCommon")


T = TypeVar("T", bound=TransformerConfig)


# Modified from verl:
# https://github.com/volcengine/verl/blob/ea885f32f04d86c3a81de18083db7eef0d781421/verl/models/mcore/config_converter.py
def hf_to_mcore_base_args(
Expand Down Expand Up @@ -68,9 +73,7 @@ def hf_to_mcore_base_args(

# Modified from verl:
# https://github.com/volcengine/verl/blob/ea885f32f04d86c3a81de18083db7eef0d781421/verl/models/mcore/config_converter.py
def check_and_construct_configs[T: TransformerConfig](
original_config: dict, cls: type[T]
) -> T:
def check_and_construct_configs(original_config: dict, cls: type[T]) -> T:
"""
Check and disable incompatible configurations for older Megatron version.

Expand Down
5 changes: 4 additions & 1 deletion areal/utils/functional/vocab_parallel.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,14 @@
import functools
from collections.abc import Callable
from typing import TypeVar

import torch
from torch import distributed as dist

from areal.infra.platforms import is_npu_available

T = TypeVar("T", torch.Tensor, tuple[torch.Tensor, torch.Tensor])


def _gather_logprobs(
logits: torch.Tensor, labels: torch.Tensor, temperature: float = 1.0
Expand Down Expand Up @@ -33,7 +36,7 @@ def _should_use_torch_compile() -> bool:
_gather_logprobs_entropy = torch.compile(_gather_logprobs_entropy)


def _chunked_apply[T: (torch.Tensor, tuple[torch.Tensor, torch.Tensor])](
def _chunked_apply(
fn: Callable[[torch.Tensor, torch.Tensor], T],
logits: torch.Tensor,
labels: torch.Tensor,
Expand Down
5 changes: 3 additions & 2 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@ name = "areal"
description = "AReaL: A Large-Scale Asynchronous Reinforcement Learning System"
readme = "README.md"
license = {text = "Apache-2.0"}
requires-python = ">=3.12,<3.13"
requires-python = ">=3.11,<3.13"
version = "1.0.1"
authors = [
{name = "AReaL Team"},
Expand All @@ -34,6 +34,7 @@ classifiers = [
"License :: OSI Approved :: Apache Software License",
"Operating System :: POSIX :: Linux",
"Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.11",
"Programming Language :: Python :: 3.12",
"Topic :: Scientific/Engineering :: Artificial Intelligence",
"Topic :: System :: Distributed Computing",
Expand Down Expand Up @@ -245,7 +246,7 @@ markers = [

[tool.ruff]
line-length = 88
target-version = "py312"
target-version = "py311"

[tool.ruff.lint]
select = [
Expand Down
Loading