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: 7 additions & 0 deletions src/spindle/backends/miles_runtime/actor.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
import torch
import torch.distributed as dist
from megatron.bridge import AutoBridge
from megatron.bridge.models.qwen_vl.modelling_qwen3_vl import model as qwen3_vl_model
from megatron.core import dist_checkpointing
from miles.backends.fsdp_utils import actor as fsdp_actor
from miles.backends.megatron_utils import actor as megatron_actor
Expand All @@ -27,8 +28,10 @@
from miles.backends.training_utils.replay_data import fill_replay_data
from miles.backends.training_utils.weight_update import snapshot_publisher
from miles.utils.replay_base import routing_replay_manager
from miles_plugins.models import qwen3_vl as miles_qwen3_vl

from .profiling import RankProfiler, TorchProfileConfig
from .qwen3_vl_cp import install_qwen3_vl_cp_position_ids
from .replay import install_replay_hooks


Expand Down Expand Up @@ -182,6 +185,10 @@ def slice_log_prob_with_cp(
manager=routing_replay_manager,
)

install_qwen3_vl_cp_position_ids(
bridge_model=qwen3_vl_model, miles_qwen3_vl=miles_qwen3_vl
)


def _checkpoint_volume_path(path: Path) -> str | None:
"""Locate ``path`` inside the checkpoint volume, or ``None`` if outside."""
Expand Down
33 changes: 33 additions & 0 deletions src/spindle/backends/miles_runtime/qwen3_vl_cp.py

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This is fine for now but can we make an upstream PR to Miles?

Original file line number Diff line number Diff line change
@@ -0,0 +1,33 @@
"""Hand Miles' rank-local Qwen3-VL MRoPE positions to Bridge under CP.

Miles pre-shards packed THD rows across context-parallel ranks and supplies
per-segment MRoPE positions through a patched ``get_rope_index``. The pinned
Megatron-Bridge accepts pre-sharded rows only with explicit 3D
``position_ids`` and raises before it calls ``get_rope_index``.
"""

from __future__ import annotations

from functools import wraps

_INSTALLED = "_spindle_cp_position_ids"


def install_qwen3_vl_cp_position_ids(*, bridge_model, miles_qwen3_vl) -> None:
model_cls = bridge_model.Qwen3VLModel
forward = model_cls.forward
if getattr(forward, _INSTALLED, False):
return

@wraps(forward)
def forward_with_positions(self, *args, **kwargs):
if kwargs.get("position_ids") is None:
parsed = miles_qwen3_vl._parse_packed_thd(args, kwargs)
if miles_qwen3_vl._prepare_cp_local_context(parsed) is not None:
kwargs["position_ids"] = miles_qwen3_vl._build_packed_positions(
self, parsed, kwargs, bridge_model.get_rope_index
)
return forward(self, *args, **kwargs)

setattr(forward_with_positions, _INSTALLED, True)
model_cls.forward = forward_with_positions
63 changes: 36 additions & 27 deletions src/spindle/providers/modal/miles_image.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
import shlex

import modal

from .image_dependencies import (
Expand All @@ -8,17 +10,41 @@
ignore_config_source,
)

BASE_IMAGE = "radixark/miles:v0.1.0"
MILES_REPOSITORY = "https://github.com/radixark/miles.git"
# Update this when Miles main should be picked up, then refresh the trainer app.
MILES_COMMIT = "5510af675238be8271c0a24740f70d5116f6d32b"
# The release tag pins Miles together with the Megatron-LM and Megatron-Bridge
# revisions its image was built and tested with. Update all three by moving to a
# newer release tag, digest, and commit together, then refresh the trainer app.
MILES_RELEASE = "v0.1.1"
BASE_IMAGE = (
f"radixark/miles:{MILES_RELEASE}"
"@sha256:6355834f16bacd35d5d40c43f142e3758376f7b2e8d678bccfe870c092bd96bf"
)
MILES_COMMIT = "2806267d060d51b1d3b62f85a1f9b145047aeef9"
MILES_PATH = "/root/miles"
MEGATRON_REPOSITORY = "https://github.com/radixark/Megatron-LM.git"
MEGATRON_REVISION = "8c1e05747eb612b382df2632783df5c83a853646"
MEGATRON_PATH = "/root/Megatron-LM"
BRIDGE_REPOSITORY = "https://github.com/radixark/Megatron-Bridge.git"
BRIDGE_REVISION = "582783a05442245647239e4c5e7d733d7f0e00ea"
BRIDGE_PATH = "/root/Megatron-Bridge"
Comment on lines +13 to -21

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

solid change btw, good find

RELEASE_CHECK = "; ".join(

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we move this into a normal Python runction then run it with Modal's Image.run_function?

Something like

def _check_release(miles_commit: str) -> None:
    lock = json.loads(Path(MILES_PATH, "release-lock.json").read_text())
    assert _head(MILES_PATH) == miles_commit, "Miles is not at MILES_COMMIT"
    assert _head(MEGATRON_PATH) == lock["megatron_commit"], "Megatron-LM differs from release-lock.json"
    ...

image = (
    modal.Image.from_registry(BASE_IMAGE)
    ...
    .run_function(_check_release, kwargs={"miles_commit": MILES_COMMIT})

[
"import json, re, subprocess",
"from importlib.metadata import distribution",
(
"head = lambda path: subprocess.check_output("
"['git', '-C', path, 'rev-parse', 'HEAD'], text=True).strip()"
),
f"assert head('{MILES_PATH}') == '{MILES_COMMIT}', 'Miles is not at MILES_COMMIT'",
f"lock = json.load(open('{MILES_PATH}/release-lock.json'))",
(
f"assert head('{MEGATRON_PATH}') == lock['megatron_commit'], "
"'Megatron-LM differs from release-lock.json'"
),
f"dockerfile = open('{MILES_PATH}/docker/Dockerfile').read()",
"bridge = re.search(r'Megatron-Bridge[.]git@([0-9a-f]{40})', dockerfile)[1]",
"url = json.loads(distribution('megatron-bridge').read_text('direct_url.json'))",
(
"assert url['vcs_info']['commit_id'] == bridge, "
"'Megatron-Bridge differs from the Miles Dockerfile'"
),
"print('megatron-lm', lock['megatron_commit'], 'megatron-bridge', bridge)",
]
)

image = (
modal.Image.from_registry(BASE_IMAGE)
Expand All @@ -30,24 +56,7 @@
"SPINDLE_MILES_COMMIT": MILES_COMMIT,
}
)
.apt_install("git")
.run_commands(
f"rm -rf {MILES_PATH}"
f" && git clone --filter=blob:none {MILES_REPOSITORY} {MILES_PATH}"
f" && git -C {MILES_PATH} fetch --depth 1 origin {MILES_COMMIT}"
f" && git -C {MILES_PATH} checkout --detach FETCH_HEAD",
f"pip install --no-build-isolation --no-deps -e {MILES_PATH}",
f"git -C {MEGATRON_PATH} fetch --depth 1"
f" {MEGATRON_REPOSITORY} {MEGATRON_REVISION}"
f" && git -C {MEGATRON_PATH} checkout --detach FETCH_HEAD",
f"pip install --no-build-isolation --no-deps -e {MEGATRON_PATH}",
f"rm -rf {BRIDGE_PATH}"
f" && git clone --filter=blob:none {BRIDGE_REPOSITORY} {BRIDGE_PATH}"
f" && git -C {BRIDGE_PATH} fetch --depth 1 origin {BRIDGE_REVISION}"
f" && git -C {BRIDGE_PATH} checkout --detach FETCH_HEAD",
"pip uninstall -y megatron-bridge megatron_bridge || true",
f"pip install --no-build-isolation --no-deps -e {BRIDGE_PATH}",
)
.run_commands(f"python -c {shlex.quote(RELEASE_CHECK)}")
.pip_install(*CORE_PACKAGES, STITCH_PACKAGE)
.pip_install(
*MEGATRON_RUNTIME_PACKAGES,
Expand Down
5 changes: 5 additions & 0 deletions tests/backends/test_miles_actor.py
Original file line number Diff line number Diff line change
Expand Up @@ -104,6 +104,8 @@ def _load_actor(monkeypatch):
_module(monkeypatch, "torch.distributed")
megatron_bridge = _module(monkeypatch, "megatron.bridge")
megatron_bridge.AutoBridge = type("AutoBridge", (), {})
_module(monkeypatch, "megatron.bridge.models.qwen_vl.modelling_qwen3_vl.model")
_module(monkeypatch, "miles_plugins.models.qwen3_vl")
_module(monkeypatch, "megatron.core").dist_checkpointing = types.ModuleType(
"megatron.core.dist_checkpointing"
)
Expand Down Expand Up @@ -174,6 +176,9 @@ def _load_actor(monkeypatch):
_module(
monkeypatch, "spindle.backends.miles_runtime.replay"
).install_replay_hooks = lambda **kwargs: None
_module(
monkeypatch, "spindle.backends.miles_runtime.qwen3_vl_cp"
).install_qwen3_vl_cp_position_ids = lambda **kwargs: None

path = Path(__file__).parents[2] / "src/spindle/backends/miles_runtime/actor.py"
spec = importlib.util.spec_from_file_location(
Expand Down
53 changes: 53 additions & 0 deletions tests/backends/test_qwen3_vl_cp.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,53 @@
from types import SimpleNamespace

from spindle.backends.miles_runtime.qwen3_vl_cp import install_qwen3_vl_cp_position_ids


def _install(pre_sharded: bool):
calls = []

class Qwen3VLModel:
def forward(self, *args, **kwargs):
calls.append(kwargs)
return "out"

def get_rope_index(*args, **kwargs):
raise AssertionError("positions come from Miles")

def build_positions(model, parsed, kwargs, rope_index):
assert rope_index is get_rope_index
return ("positions", parsed)

bridge_model = SimpleNamespace(
Qwen3VLModel=Qwen3VLModel, get_rope_index=get_rope_index
)
miles_qwen3_vl = SimpleNamespace(
_parse_packed_thd=lambda args, kwargs: kwargs["input_ids"],
_prepare_cp_local_context=lambda parsed: (
{"psp": parsed} if pre_sharded else None
),
_build_packed_positions=build_positions,
)
install_qwen3_vl_cp_position_ids(
bridge_model=bridge_model, miles_qwen3_vl=miles_qwen3_vl
)
install_qwen3_vl_cp_position_ids(
bridge_model=bridge_model, miles_qwen3_vl=miles_qwen3_vl
)
return Qwen3VLModel(), calls


def test_pre_sharded_cp_input_gets_explicit_position_ids():
model, calls = _install(pre_sharded=True)
assert model.forward(input_ids="row", position_ids=None) == "out"
assert calls == [{"input_ids": "row", "position_ids": ("positions", "row")}]


def test_other_inputs_pass_through():
model, calls = _install(pre_sharded=False)
model.forward(input_ids="row", position_ids=None)
model.forward(input_ids="row", position_ids="explicit")
assert calls == [
{"input_ids": "row", "position_ids": None},
{"input_ids": "row", "position_ids": "explicit"},
]
Loading