diff --git a/src/spindle/backends/miles_runtime/actor.py b/src/spindle/backends/miles_runtime/actor.py index 5192d0d..a084aa2 100644 --- a/src/spindle/backends/miles_runtime/actor.py +++ b/src/spindle/backends/miles_runtime/actor.py @@ -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 @@ -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 @@ -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.""" diff --git a/src/spindle/backends/miles_runtime/qwen3_vl_cp.py b/src/spindle/backends/miles_runtime/qwen3_vl_cp.py new file mode 100644 index 0000000..6c6d0f2 --- /dev/null +++ b/src/spindle/backends/miles_runtime/qwen3_vl_cp.py @@ -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 diff --git a/src/spindle/providers/modal/miles_image.py b/src/spindle/providers/modal/miles_image.py index 3eae0e8..f673794 100644 --- a/src/spindle/providers/modal/miles_image.py +++ b/src/spindle/providers/modal/miles_image.py @@ -1,3 +1,5 @@ +import shlex + import modal from .image_dependencies import ( @@ -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" +RELEASE_CHECK = "; ".join( + [ + "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) @@ -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, diff --git a/tests/backends/test_miles_actor.py b/tests/backends/test_miles_actor.py index af63665..d8ee2de 100644 --- a/tests/backends/test_miles_actor.py +++ b/tests/backends/test_miles_actor.py @@ -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" ) @@ -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( diff --git a/tests/backends/test_qwen3_vl_cp.py b/tests/backends/test_qwen3_vl_cp.py new file mode 100644 index 0000000..388270f --- /dev/null +++ b/tests/backends/test_qwen3_vl_cp.py @@ -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"}, + ]