Skip to content
Open
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
7 changes: 6 additions & 1 deletion examples/vime_qwen3_8b_tp4_cp2_200/run_arm.py
Original file line number Diff line number Diff line change
Expand Up @@ -352,6 +352,7 @@ def build_parser() -> argparse.ArgumentParser:
parser.add_argument("--rollout-top-k", type=int, default=-1)
parser.add_argument("--lr", type=float, default=5e-7)
parser.add_argument("--weight-decay", type=float, default=0.1)
parser.add_argument("--entropy-coef", type=float, default=0.0)
parser.add_argument("--require-updates", action="store_true")
parser.add_argument(
"--use-kl-loss",
Expand Down Expand Up @@ -406,6 +407,8 @@ def main(argv: list[str] | None = None) -> int:
raise ValueError("--lr must be finite and positive")
if not math.isfinite(args.weight_decay) or args.weight_decay < 0:
raise ValueError("--weight-decay must be finite and nonnegative")
if not math.isfinite(args.entropy_coef) or args.entropy_coef < 0:
raise ValueError("--entropy-coef must be finite and nonnegative")
if not math.isfinite(args.kl_loss_coef) or args.kl_loss_coef < 0:
raise ValueError("--kl-loss-coef must be finite and nonnegative")
# VIME passes this value to vLLM ParallelConfig.prefill_context_parallel_size.
Expand Down Expand Up @@ -592,7 +595,7 @@ def main(argv: list[str] | None = None) -> int:
"--adam-beta2",
"0.98",
"--entropy-coef",
"0",
str(args.entropy_coef),
"--global-batch-size",
str(args.global_batch_size),
"--balance-data",
Expand All @@ -602,6 +605,7 @@ def main(argv: list[str] | None = None) -> int:
str(128 * canonical_tp // int(topology["tp"])),
"--context-parallel-size",
str(topology["cp"]),
"--sequence-parallel",
"--cp-comm-type",
"p2p",
"--pipeline-model-parallel-size",
Expand Down Expand Up @@ -721,6 +725,7 @@ def main(argv: list[str] | None = None) -> int:
},
"algorithm": {
"optimizer": {"name": "adam", "lr": args.lr, "weight_decay": args.weight_decay},
"entropy_coefficient": args.entropy_coef,
"require_updates": args.require_updates,
"advantage_estimator": "grpo",
"reward_model": "deepscaler",
Expand Down
90 changes: 78 additions & 12 deletions rl_engine/integrations/megatron_runtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -415,6 +415,40 @@ def backward(ctx: Any, grad_output: torch.Tensor) -> tuple[torch.Tensor, None]:
return collective.all_reduce(grad_output.contiguous()), None


class _DeterministicGatherFromSequenceParallelRegion(torch.autograd.Function):
"""Gather sequence shards for a TP column projection; scatter its dgrad."""

@staticmethod
def forward(ctx: Any, input_value: torch.Tensor, collective: Any | None) -> torch.Tensor:
ctx.collective = collective
if collective is None:
return input_value
return collective.all_gather(input_value.contiguous())

@staticmethod
def backward(ctx: Any, grad_output: torch.Tensor) -> tuple[torch.Tensor, None]:
if ctx.collective is None:
return grad_output, None
return ctx.collective.reduce_scatter(grad_output.contiguous()), None


class _DeterministicReduceScatterToSequenceParallelRegion(torch.autograd.Function):
"""Reduce a TP row projection into sequence shards; gather its dgrad."""

@staticmethod
def forward(ctx: Any, input_value: torch.Tensor, collective: Any | None) -> torch.Tensor:
ctx.collective = collective
if collective is None:
return input_value
return collective.reduce_scatter(input_value.contiguous())

@staticmethod
def backward(ctx: Any, grad_output: torch.Tensor) -> tuple[torch.Tensor, None]:
if ctx.collective is None:
return grad_output, None
return ctx.collective.all_gather(grad_output.contiguous()), None


def _deterministic_reduce_from_tensor_model_parallel_region(
input_value: torch.Tensor,
collective: Any | None,
Expand Down Expand Up @@ -567,9 +601,24 @@ def forward(
weight: torch.Tensor,
bias: torch.Tensor | None,
tp_group: Any,
sequence_parallel: bool,
) -> torch.Tensor:
from rl_engine.kernels.ops.matmul.det_gemm import det_gemm_linear

ctx.collective = None
ctx.input_shape = input_value.shape
if sequence_parallel and _tp_world_size(tp_group) > 1:
from rl_engine.distributed.collectives import collective_for_group

ctx.collective = collective_for_group(
tp_group,
min_size_bytes=input_value.numel()
* input_value.element_size()
* _tp_world_size(tp_group),
)
if ctx.collective is None:
raise RuntimeError("strict SP LM head requires a deterministic TP collective")
input_value = ctx.collective.all_gather(input_value.contiguous())
if input_value.ndim == 3:
input_2d = input_value.transpose(0, 1).contiguous().reshape(-1, input_value.shape[-1])
ctx.batch_major = True
Expand All @@ -581,7 +630,7 @@ def forward(
if bias is not None:
output_2d = (output_2d.float() + bias.float().reshape(1, -1)).to(torch.bfloat16)
ctx.save_for_backward(input_2d, weight_2d)
ctx.input_shape = input_value.shape
ctx.full_input_shape = input_value.shape
ctx.input_dtype = input_value.dtype
ctx.weight_dtype = weight.dtype
ctx.bias_dtype = None if bias is None else bias.dtype
Expand Down Expand Up @@ -612,15 +661,22 @@ def backward(ctx: Any, grad_output: torch.Tensor):
grad_input = _canonical_column_input_gradient(
dlogits, weight, _tp_world_size(ctx.tp_group)
)
_deterministic_tp_all_reduce_(grad_input, ctx.tp_group)
if ctx.collective is None:
_deterministic_tp_all_reduce_(grad_input, ctx.tp_group)
if ctx.batch_major:
grad_input = (
grad_input.reshape(ctx.input_shape[1], ctx.input_shape[0], ctx.input_shape[2])
grad_input.reshape(
ctx.full_input_shape[1], ctx.full_input_shape[0], ctx.full_input_shape[2]
)
.transpose(0, 1)
.contiguous()
)
else:
grad_input = grad_input.reshape(ctx.input_shape)
grad_input = grad_input.reshape(ctx.full_input_shape)
if ctx.collective is not None:
grad_input = ctx.collective.reduce_scatter(grad_input.contiguous())
if grad_input.shape != ctx.input_shape:
raise RuntimeError("strict TP LM-head dgrad does not match the input layout")
grad_input = grad_input.to(ctx.input_dtype)
if ctx.needs_input_grad[1]:
tp_world = _tp_world_size(ctx.tp_group)
Expand All @@ -632,7 +688,7 @@ def backward(ctx: Any, grad_output: torch.Tensor):
).to(ctx.weight_dtype)
if ctx.has_bias and ctx.needs_input_grad[2]:
grad_bias = dlogits.float().sum(dim=0).to(ctx.bias_dtype)
return grad_input, grad_weight, grad_bias, None
return grad_input, grad_weight, grad_bias, None, None


def _optional_class(path: str) -> type[Any] | None:
Expand Down Expand Up @@ -793,6 +849,8 @@ def strict_tp_copy(
input_value: torch.Tensor,
) -> torch.Tensor:
if copy_to_tp is not None:
if sequence_parallel_enabled(module):
raise RuntimeError("strict SP Attention requires a deterministic TP collective")
record_collective_backend(
core_attention,
_MEGATRON_TP_QKV_DGRAD_COLLECTIVE_ATTR,
Expand All @@ -805,6 +863,8 @@ def strict_tp_copy(
_MEGATRON_TP_QKV_DGRAD_COLLECTIVE_ATTR,
_collective_backend_id(collective),
)
if sequence_parallel_enabled(module):
return _DeterministicGatherFromSequenceParallelRegion.apply(input_value, collective)
return _deterministic_copy_to_tensor_model_parallel_region(input_value, collective)

def strict_tp_reduce(
Expand All @@ -813,6 +873,8 @@ def strict_tp_reduce(
input_value: torch.Tensor,
) -> torch.Tensor:
if reduce_from_tp is not None:
if sequence_parallel_enabled(module):
raise RuntimeError("strict SP Attention requires a deterministic TP collective")
record_collective_backend(
core_attention,
_MEGATRON_TP_OUTPUT_PROJECTION_COLLECTIVE_ATTR,
Expand All @@ -825,6 +887,10 @@ def strict_tp_reduce(
_MEGATRON_TP_OUTPUT_PROJECTION_COLLECTIVE_ATTR,
_collective_backend_id(collective),
)
if sequence_parallel_enabled(module):
return _DeterministicReduceScatterToSequenceParallelRegion.apply(
input_value, collective
)
return _deterministic_reduce_from_tensor_model_parallel_region(input_value, collective)

def bind_collective_identity(module: Any, core_attention: Any, attribute: str) -> None:
Expand Down Expand Up @@ -894,10 +960,8 @@ def attention_init_wrapped(instance: Any, *args: Any, **kwargs: Any) -> None:
qkv = instance.linear_qkv
projection = instance.linear_proj
core_attention = getattr(instance, "core_attention", None)
if sequence_parallel_enabled(qkv) or sequence_parallel_enabled(projection):
raise RuntimeError(
"strict Attention projection collectives do not support sequence parallelism"
)
if sequence_parallel_enabled(qkv) != sequence_parallel_enabled(projection):
raise RuntimeError("strict Attention QKV and output projection SP settings differ")
setattr(qkv, _STRICT_ATTENTION_PROJECTION_MARKER, "qkv")
setattr(projection, _STRICT_ATTENTION_PROJECTION_MARKER, "o_proj")
object.__setattr__(qkv, _STRICT_ATTENTION_CORE_MARKER, core_attention)
Expand Down Expand Up @@ -1062,15 +1126,17 @@ def wrapped(
)
if gather_output:
raise RuntimeError("strict reusable TP LM head does not support gathered logits")
if bool(getattr(instance, "sequence_parallel", False)):
raise RuntimeError("strict reusable TP LM head does not support sequence parallelism")
if bool(getattr(instance, "explicit_expert_comm", False)) or bool(
getattr(instance, "disable_grad_reduce", False)
):
raise RuntimeError("strict reusable TP LM head requires ordinary TP dgrad reduction")
bias = instance.bias if not instance.skip_bias_add else None
output = _DeterministicTPOutputProjection.apply(
input_, output_weight, bias, instance.tp_group
input_,
output_weight,
bias,
instance.tp_group,
bool(getattr(instance, "sequence_parallel", False)),
)
instance._rl_kernel_local_logits = output
output_bias = instance.bias if instance.skip_bias_add else None
Expand Down
2 changes: 1 addition & 1 deletion tests/test_canonical_lm_head_backward.py
Original file line number Diff line number Diff line change
Expand Up @@ -65,7 +65,7 @@ def test_lm_head_autograd_uses_canonical_subtree(monkeypatch, shape):
x = torch.randn(shape, generator=gen).bfloat16().requires_grad_()
w = torch.randn(32, 8, generator=gen).bfloat16().requires_grad_()
b = torch.randn(32, generator=gen).bfloat16().requires_grad_()
output = runtime._DeterministicTPOutputProjection.apply(x, w, b, object())
output = runtime._DeterministicTPOutputProjection.apply(x, w, b, object(), False)
dy = torch.randn(output.shape, generator=gen).bfloat16()
output.backward(dy)

Expand Down
154 changes: 150 additions & 4 deletions tests/test_framework_runtime_adapters.py
Original file line number Diff line number Diff line change
Expand Up @@ -785,10 +785,156 @@ def all_reduce(self, value):
output.sum().backward()

assert torch.equal(output, value.detach() * 4)
assert torch.equal(value.grad, torch.ones_like(value))


def test_vllm_qwen3_strict_model_installs_without_debug_environment(monkeypatch):
assert torch.equal(value.grad, torch.ones_like(value))


def test_megatron_sequence_parallel_projection_collectives_preserve_autograd_contract():
from rl_engine.integrations.megatron_runtime import (
_DeterministicGatherFromSequenceParallelRegion,
_DeterministicReduceScatterToSequenceParallelRegion,
)

class Collective:
def __init__(self):
self.operations = []

def all_gather(self, value):
self.operations.append(("all_gather", tuple(value.shape)))
return torch.cat((value, value), dim=0)

def reduce_scatter(self, value):
self.operations.append(("reduce_scatter", tuple(value.shape)))
return value.chunk(2, dim=0)[0] + value.chunk(2, dim=0)[1]

collective = Collective()
local = torch.tensor([[[1.0, 2.0]]], requires_grad=True)
gathered = _DeterministicGatherFromSequenceParallelRegion.apply(local, collective)
assert gathered.shape == (2, 1, 2)
gathered.sum().backward()
assert torch.equal(local.grad, torch.full_like(local, 2.0))
assert collective.operations == [("all_gather", (1, 1, 2)), ("reduce_scatter", (2, 1, 2))]

collective.operations.clear()
full = torch.ones((2, 1, 2), requires_grad=True)
scattered = _DeterministicReduceScatterToSequenceParallelRegion.apply(full, collective)
assert scattered.shape == (1, 1, 2)
scattered.sum().backward()
assert torch.equal(full.grad, torch.ones_like(full))
assert collective.operations == [("reduce_scatter", (2, 1, 2)), ("all_gather", (1, 1, 2))]


def test_megatron_strict_attention_sequence_parallel_uses_tp_collective(monkeypatch):
from rl_engine.kernels.ops.matmul import det_gemm

class Collective:
backend_id = "test.fixed_tree"

def __init__(self):
self.operations = []

def all_gather(self, value):
self.operations.append("gather")
return torch.cat((value, value), dim=0)

def reduce_scatter(self, value):
self.operations.append("scatter")
return value.chunk(2, dim=0)[0] + value.chunk(2, dim=0)[1]

class ColumnLinear:
def __init__(self):
self.weight = torch.eye(2, requires_grad=True)
self.sequence_parallel = True
self.gather_output = False
self.skip_bias_add = False
self.bias = None

def forward(self, input):
return self._forward_impl(input, self.weight), None

def _forward_impl(self, input, weight, **kwargs):
return input @ weight.t()

class RowLinear:
def __init__(self):
self.weight = torch.eye(2, requires_grad=True)
self.sequence_parallel = True
self.input_is_parallel = True
self.skip_bias_add = False
self.bias = None

def forward(self, input):
return self._forward_impl(input, self.weight), None

def _forward_impl(self, input, weight, **kwargs):
return input @ weight.t()

class SelfAttention:
def __init__(self):
self.linear_qkv = ColumnLinear()
self.linear_proj = RowLinear()

collective = Collective()
monkeypatch.setattr(
"rl_engine.integrations.megatron_runtime._fixed_tree_collective",
lambda module, input_value=None: collective,
)
monkeypatch.setattr(det_gemm, "det_gemm_linear_input_gradient", lambda lhs, rhs: lhs @ rhs)
monkeypatch.setattr(det_gemm, "det_gemm_linear_weight_gradient", lambda lhs, rhs: rhs.t() @ lhs)
_patch_strict_attention_projections(
self_attention_cls=SelfAttention,
column_linear_cls=ColumnLinear,
row_linear_cls=RowLinear,
det_gemm=lambda lhs, rhs: lhs @ rhs,
)
attention = SelfAttention()
local = torch.tensor([[[1.0, 2.0]]], requires_grad=True)
qkv, _ = attention.linear_qkv.forward(local)
output, _ = attention.linear_proj.forward(qkv)
assert qkv.shape == (2, 1, 2)
assert output.shape == local.shape
output.sum().backward()
assert torch.equal(local.grad, torch.full_like(local, 2.0))
assert collective.operations == ["gather", "scatter", "gather", "scatter"]


def test_megatron_strict_lm_head_sequence_parallel_dgrad_is_sharded(monkeypatch):
from rl_engine.integrations.megatron_runtime import _DeterministicTPOutputProjection
from rl_engine.kernels.ops.matmul import det_gemm

class Collective:
def __init__(self):
self.operations = []

def all_gather(self, value):
self.operations.append(("gather", tuple(value.shape)))
return torch.cat((value, value * 3), dim=0)

def reduce_scatter(self, value):
self.operations.append(("scatter", tuple(value.shape)))
return value.chunk(2, dim=0)[0] + value.chunk(2, dim=0)[1]

collective = Collective()
monkeypatch.setattr("rl_engine.integrations.megatron_runtime._tp_world_size", lambda group: 2)
monkeypatch.setattr(
"rl_engine.distributed.collectives.collective_for_group",
lambda group, min_size_bytes: collective,
)
monkeypatch.setattr(det_gemm, "det_gemm_linear", lambda lhs, rhs: lhs @ rhs.t())
monkeypatch.setattr(det_gemm, "det_gemm_linear_input_gradient", lambda lhs, rhs: lhs @ rhs)
monkeypatch.setattr(det_gemm, "det_gemm_linear_weight_gradient", lambda lhs, rhs: rhs.t() @ lhs)
local = torch.tensor([[[1.0, 2.0]]], dtype=torch.bfloat16, requires_grad=True)
weight = torch.eye(2, dtype=torch.bfloat16, requires_grad=True)
logits = _DeterministicTPOutputProjection.apply(local, weight, None, object(), True)
assert logits.shape == (2, 1, 2)
assert torch.equal(logits[:, 0], torch.tensor([[1, 2], [3, 6]], dtype=torch.bfloat16))
logits.sum().backward()
assert local.grad.shape == local.shape
assert torch.equal(local.grad, torch.full_like(local, 2))
assert torch.equal(weight.grad, torch.tensor([[4, 8], [4, 8]], dtype=torch.bfloat16))
assert collective.operations == [("gather", (1, 1, 2)), ("scatter", (2, 1, 2))]


def test_vllm_qwen3_strict_model_installs_without_debug_environment(monkeypatch):
monkeypatch.delenv("RL_KERNEL_MODEL_DEBUG_DIR", raising=False)

class RMSNorm:
Expand Down
Loading
Loading