Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
23 commits
Select commit Hold shift + click to select a range
0fcf901
feat(data-plane): enable RDMA transport for TransferQueue
overloadedHenry Aug 11, 2026
75424b0
fix(data-plane): harden TQ RDMA lifecycle
overloadedHenry Aug 13, 2026
2b14c07
fix(data-plane): enforce TQ correctness contracts
overloadedHenry Aug 13, 2026
7e91be2
test(data-plane): real multimodal byte-exact tier
overloadedHenry Aug 13, 2026
0d8c26c
fix(data-plane): guard mooncake 0.3.10 TCP memcpy corruption
overloadedHenry Aug 14, 2026
4ab4578
docs: restructure transfer_queue_rdma.md as a usage guide
overloadedHenry Aug 14, 2026
0057df6
fix(data-plane): harden env and HCA probe guards
overloadedHenry Aug 14, 2026
e209e1b
fix(data-plane): close TQ owner, detach clients
overloadedHenry Aug 14, 2026
0f579fd
fix(data-plane): restore default path, bound attach
overloadedHenry Aug 14, 2026
18c758d
fix(data-plane): size segment from token budget
overloadedHenry Aug 14, 2026
682400a
refactor(data-plane): split runtime patches out of PR
overloadedHenry Aug 14, 2026
830cada
docs: sync rdma guide with review-round changes
overloadedHenry Aug 14, 2026
b15ee18
fix(tests): stub attach_tq_client now that SFT no longer imports tq
overloadedHenry Aug 14, 2026
66ec9d5
fix(data-plane): address RDMA review findings
overloadedHenry Aug 15, 2026
151b097
style(controller): sort TQ config imports
overloadedHenry Aug 15, 2026
dcbc85b
fix(data-plane): guard TQ client generations
overloadedHenry Aug 15, 2026
cfe6153
fix(data-plane): isolate TQ attach workers
overloadedHenry Aug 17, 2026
10ec8de
test(data-plane): harden TQ lifecycle tests
overloadedHenry Aug 17, 2026
6245ee9
feat(data-plane): version-gated mooncake loss guards
overloadedHenry Aug 14, 2026
fdee456
fix(data-plane): preserve Mooncake semantics
overloadedHenry Aug 15, 2026
6d984ce
test(data-plane): cover short get retry results
overloadedHenry Aug 17, 2026
c500eb4
fix(data-plane): gate guards by TQ revision
overloadedHenry Aug 18, 2026
2e2b99f
Merge branch 'main' into feat/tq-mooncake-loss-guards
overloadedHenry Aug 18, 2026
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
4 changes: 4 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -54,3 +54,7 @@ tensorboard_log
# .github
.github/copilot-instructions.md
env.sh

# Machine-local multimodal acceptance fixtures (generated by
# scripts/benchmarks/make_multimodal_fixture.py; hundreds of MB, never commit)
tests/fixtures/
222 changes: 222 additions & 0 deletions docs/draft/transfer_queue_rdma.md

Large diffs are not rendered by default.

24 changes: 21 additions & 3 deletions relax/backends/megatron/actor.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,6 @@
import requests
import torch
import torch.distributed as dist
import transfer_queue as tq
from megatron.core import mpu


Expand Down Expand Up @@ -61,6 +60,7 @@
from relax.utils.rotate_ckpt import rotate_ckpt
from relax.utils.s3_model_loader import prepare_model_maybe_update_args
from relax.utils.timer import Timer, inverse_timer, timer, with_defer
from relax.utils.tq_lifecycle import attach_tq_client, detach_tq_client
from relax.utils.tracking_utils import init_tracking
from relax.utils.training import train_dump_utils
from relax.utils.training.data_fields import build_data_fields
Expand Down Expand Up @@ -151,6 +151,20 @@ def _per_step_rollout(self) -> bool:
periodic predict steps; Megatron stays awake between."""
return not is_sft_mode(self.args)

def __del__(self) -> None:
# Best-effort detach on graceful teardown; ray.kill / fate-sharing
# kills skip destructors, in which case the Mooncake master TTL
# reclaims the segment.
generation = getattr(self, "_tq_client_generation", None)
if getattr(self, "data_system_client", None) is None or generation is None:
return
try:
detach_tq_client(generation)
self._tq_client_generation = None
self.data_system_client = None
except Exception: # destructor must never raise (interpreter shutdown)
return

def init(
self,
args: Namespace,
Expand Down Expand Up @@ -187,8 +201,12 @@ def _init(
init(args)
if repatch is not None:
repatch(args)
tq.init(args.tq_config)
self.data_system_client = tq.get_client()
self.data_system_client = attach_tq_client(
args.tq_config,
requested_gdr=getattr(args, "tq_use_gdr", False),
role=role,
lease_owner=self,
)
if is_megatron_main_rank():
init_tracking(args, primary=False)

Expand Down
10 changes: 7 additions & 3 deletions relax/components/actor.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,6 @@
from typing import Any, Dict, Optional

import ray
import transfer_queue as tq
from fastapi import FastAPI
from ray import serve

Expand All @@ -17,6 +16,7 @@
from relax.engine.sft.runtime import is_sft_mode, sft_partition_id, sft_task_name
from relax.utils.async_utils import run
from relax.utils.opd.opd_utils import set_managed_opd_teacher_on_train_group
from relax.utils.tq_lifecycle import attach_tq_client


app = FastAPI()
Expand Down Expand Up @@ -71,8 +71,12 @@ def __init__(

self.actor_model = allocate_train_group(args=config, num_gpus=num_gpus, pg=pgs, runtime_env=runtime_env)

tq.init(self.config.tq_config)
self.data_system_client = tq.get_client()
self.data_system_client = attach_tq_client(
self.config.tq_config,
requested_gdr=getattr(self.config, "tq_use_gdr", False),
role=self.role,
lease_owner=self,
)

self.steps = ray.get(
self.actor_model.async_init(
Expand Down
10 changes: 7 additions & 3 deletions relax/components/actor_fwd.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,12 +6,12 @@
from typing import Any, Optional

import ray
import transfer_queue as tq
from fastapi import FastAPI
from ray import serve

from relax.components.base import Base
from relax.distributed.ray.placement_group import allocate_train_group
from relax.utils.tq_lifecycle import attach_tq_client


app = FastAPI()
Expand All @@ -35,8 +35,12 @@ def __init__(
self._run_thread = None
self._done_event: Optional[asyncio.Event] = None
self._thread_error: Optional[Exception] = None
tq.init(self.config.tq_config)
self.data_system_client = tq.get_client()
self.data_system_client = attach_tq_client(
self.config.tq_config,
requested_gdr=getattr(self.config, "tq_use_gdr", False),
role=self.role,
lease_owner=self,
)
self.actor_model = allocate_train_group(args=config, num_gpus=num_gpus, pg=pgs, runtime_env=runtime_env)
ray.get(self.actor_model.async_init(config, role=self.role, with_ref=False))
self.step = 0
Expand Down
10 changes: 7 additions & 3 deletions relax/components/advantages.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,6 @@
from typing import Any, Dict

import torch
import transfer_queue as tq
from megatron.core import mpu
from ray import serve
from tensordict import TensorDict
Expand All @@ -17,6 +16,7 @@
apply_opd_to_advantages,
consume_opd_advantage_data,
)
from relax.utils.tq_lifecycle import attach_tq_client
from relax.utils.training.ppo_utils import (
compute_approx_kl,
get_advantages_and_returns_batch,
Expand All @@ -39,8 +39,12 @@ def __init__(
self._lock = threading.RLock()
self.healthy = healthy

tq.init(self.config.tq_config)
self.data_system_client = tq.get_client()
self.data_system_client = attach_tq_client(
self.config.tq_config,
requested_gdr=getattr(self.config, "tq_use_gdr", False),
role="advantages",
lease_owner=self,
)
self.step = 0

async def run(self) -> None:
Expand Down
18 changes: 18 additions & 0 deletions relax/components/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -89,6 +89,24 @@ def __init__(self) -> None:
self.step = 0
self._logger_instance = None
self._lock = threading.Lock()
self._tq_client_generation: int | None = None

def __del__(self) -> None:
# Ray Serve calls the destructor on replica shutdown (normal stop,
# global restart, in-place restart). Components that attached a
# TransferQueue client must detach so a MooncakeStore segment
# deregisters before client_ttl instead of leaving a stale endpoint.
generation = getattr(self, "_tq_client_generation", None)
if getattr(self, "data_system_client", None) is None or generation is None:
return
try:
from relax.utils.tq_lifecycle import detach_tq_client

detach_tq_client(generation)
self._tq_client_generation = None
self.data_system_client = None
except Exception: # destructor must never raise (interpreter shutdown)
return

@property
def _logger(self):
Expand Down
10 changes: 7 additions & 3 deletions relax/components/critic.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,6 @@
from typing import Any, Optional

import ray
import transfer_queue as tq
from ray import serve
from ray.serve.schema import LoggingConfig

Expand All @@ -16,6 +15,7 @@
from relax.distributed.ray.placement_group import allocate_train_group
from relax.engine.sft.runtime import sft_partition_id
from relax.utils.async_utils import run
from relax.utils.tq_lifecycle import attach_tq_client


@serve.deployment(
Expand All @@ -40,8 +40,12 @@ def __init__(
self.healthy = healthy
self.role = role

tq.init(self.config.tq_config)
self.data_system_client = tq.get_client()
self.data_system_client = attach_tq_client(
self.config.tq_config,
requested_gdr=getattr(self.config, "tq_use_gdr", False),
role=self.role,
lease_owner=self,
)

self.critic_model = allocate_train_group(
args=config, num_gpus=num_gpus, pg=pgs, role=self.role, runtime_env=runtime_env
Expand Down
10 changes: 7 additions & 3 deletions relax/components/rollout.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,6 @@

import httpx
import ray
import transfer_queue as tq
from fastapi import FastAPI, HTTPException, Request
from fastapi.responses import StreamingResponse
from pydantic import BaseModel, Field
Expand All @@ -20,6 +19,7 @@
from relax.distributed.ray.placement_group import create_rollout_manager
from relax.utils.env import Envs
from relax.utils.http_utils import _wrap_ipv6
from relax.utils.tq_lifecycle import attach_tq_client


app = FastAPI()
Expand Down Expand Up @@ -333,8 +333,12 @@ def __init__(
self.config = config
self.healthy = healthy

tq.init(self.config.tq_config)
self.data_system_client = tq.get_client()
self.data_system_client = attach_tq_client(
self.config.tq_config,
requested_gdr=getattr(self.config, "tq_use_gdr", False),
role="rollout",
lease_owner=self,
)
self.rollout_manager, self.num_rollout_per_epoch = create_rollout_manager(
config, pg, data_source=data_source, runtime_env=runtime_env
)
Expand Down
10 changes: 7 additions & 3 deletions relax/components/sft.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,6 @@
import random
from typing import Any

import transfer_queue as tq
from ray import serve
from transformers import AutoConfig, AutoTokenizer

Expand All @@ -39,6 +38,7 @@
from relax.utils.data.processor_pool import ProcessorPool
from relax.utils.misc import load_function
from relax.utils.s3_model_loader import prepare_model_maybe_update_args
from relax.utils.tq_lifecycle import attach_tq_client
from relax.utils.training.eval_config import build_named_prompt_data_configs
from relax.utils.utils import dict_to_tensordict

Expand Down Expand Up @@ -81,8 +81,12 @@ def __init__(self, healthy, pgs, num_gpus, config, role, runtime_env=None): # n
self.healthy = healthy
self.step = getattr(config, "start_rollout_id", 0)

tq.init(self.config.tq_config)
self.data_system_client = tq.get_client()
self.data_system_client = attach_tq_client(
self.config.tq_config,
requested_gdr=getattr(self.config, "tq_use_gdr", False),
role=self.role,
lease_owner=self,
)

self._dataset: Any | None = None
self._eval_dataset: Any | None = None
Expand Down
Loading
Loading