diff --git a/agentlightning/config/server.yaml b/agentlightning/config/server.yaml index 28a9b1970..21ac156f3 100644 --- a/agentlightning/config/server.yaml +++ b/agentlightning/config/server.yaml @@ -4,6 +4,7 @@ key: "" default_proxy: model_name: "Qwen/Qwen2.5-7B-Instruct" include_log_probs: True + include_routed_experts: False train: temperature: 1 val: diff --git a/agentlightning/server/proxy.py b/agentlightning/server/proxy.py index fb63884cb..01676fe55 100644 --- a/agentlightning/server/proxy.py +++ b/agentlightning/server/proxy.py @@ -44,6 +44,7 @@ def __init__(self, default_proxy: Mapping[str, Any]) -> None: self._train_temperature = float(default_proxy["train"]["temperature"]) self._val_temperature = float(default_proxy["val"]["temperature"]) self._include_log_probs = bool(default_proxy.get("include_log_probs", True)) + self._include_routed_experts = bool(default_proxy.get("include_routed_experts", False)) @property def model_name(self) -> str: @@ -69,6 +70,8 @@ def prepare_body(self, body: dict[str, Any], mode: str) -> dict[str, Any]: } if self._include_log_probs: prepared["logprobs"] = True + if self._include_routed_experts: + prepared["return_routed_experts"] = True return prepared if mode == "val": prepared = { @@ -126,6 +129,7 @@ async def forward_request( response_body = ( response.json() if response.headers.get("content-type", "").startswith("application/json") else {} ) + routed_experts = _pop_routed_experts(response_body) _capture_event( rollout_id=rollout_id, @@ -137,6 +141,7 @@ async def forward_request( http_status=response.status_code, status=_status_from_http_status(response.status_code), retry_count=int(response.extensions.get("agl_retry_count", 0)), + routed_experts=routed_experts, ) return JSONResponse(content=response_body, status_code=response.status_code) finally: @@ -144,6 +149,16 @@ async def forward_request( await _dec_inflight(pause_state) +def _pop_routed_experts(response_body: dict[str, Any]) -> Any: + choices = response_body.get("choices") + if not isinstance(choices, list): + return None + for choice in choices: + if isinstance(choice, dict) and (routed_experts := choice.pop("routed_experts", None)) is not None: + return routed_experts + return None + + async def _send_upstream_with_retries( *, client: httpx.AsyncClient, @@ -216,6 +231,7 @@ def _capture_event( http_status: int, status: str, retry_count: int, + routed_experts: Any = None, ) -> None: record_event( rollout_id, @@ -231,6 +247,7 @@ def _capture_event( "http_status": http_status, "status": status, "retry_count": retry_count, + "routed_experts": routed_experts, "usage": _extract_usage(response_body), "finish_reason": _extract_finish_reason(response_body), }, diff --git a/agentlightning/server/routes/events.py b/agentlightning/server/routes/events.py index 291023ef0..5afad7974 100644 --- a/agentlightning/server/routes/events.py +++ b/agentlightning/server/routes/events.py @@ -21,11 +21,39 @@ def _not_found(rollout_id: str) -> HTTPException: return HTTPException(status_code=404, detail=f"Rollout not found: {rollout_id}") +def _drop_superseded_routed_experts(events: list[Event], event_type: str, data: dict[str, Any]) -> None: + """Keep only the last route tensor in each mergeable trajectory.""" + if event_type != "model_request" or data.get("routed_experts") is None: + return + + current = _trim_model_request(data) + prompt_ids = current["prompt_token_ids"] + for index in range(len(events) - 1, -1, -1): + previous = events[index] + if previous.event_type != "model_request": + continue + previous_data = _trim_model_request(previous.data) + previous_context = previous_data["prompt_token_ids"] + previous_data["response_token_ids"] + same_prompt = prompt_ids == previous_data["prompt_token_ids"] + extends_context = prompt_ids[: len(previous_context)] == previous_context + if (same_prompt or extends_context) and "routed_experts" in previous.data: + data_without_routes = dict(previous.data) + data_without_routes.pop("routed_experts") + events[index] = previous.model_copy(update={"data": data_without_routes}) + return + + def record_event(rollout_id: str, attempt_id: str, event_type: str, data: dict[str, Any]) -> Event: """Append a single event for an existing rollout.""" if rollout_id not in _rollouts: raise _not_found(rollout_id) + rid_events = _events[rollout_id] + if attempt_id not in rid_events: + rid_events[attempt_id] = [] + attempt_events = rid_events[attempt_id] + _drop_superseded_routed_experts(attempt_events, event_type, data) + event = Event( event_type=event_type, rollout_id=rollout_id, @@ -34,10 +62,7 @@ def record_event(rollout_id: str, attempt_id: str, event_type: str, data: dict[s data=data, ) - rid_events = _events[rollout_id] - if attempt_id not in rid_events: - rid_events[attempt_id] = [] - rid_events[attempt_id].append(event) + attempt_events.append(event) return event @@ -109,6 +134,7 @@ def _trim_model_request(data: dict[str, Any]) -> dict[str, Any]: prompt_token_ids: list[int] = [] response_token_ids: list[int] = [] response_log_probs: list[float] | None = None + routed_experts = data.get("routed_experts") if isinstance(resp, dict): prompt_token_ids = resp.get("prompt_token_ids", []) @@ -136,6 +162,8 @@ def _trim_model_request(data: dict[str, Any]) -> dict[str, Any]: "response_log_probs": response_log_probs, "server": {"model": srv.get("model"), "version": srv.get("version")}, } + if routed_experts is not None: + trimmed["routed_experts"] = routed_experts for key in ("http_status", "status"): if key in data: trimmed[key] = data[key] @@ -169,6 +197,14 @@ def _to_triplet_format(event: Event) -> Event: return event +def _without_routed_experts(event: Event) -> Event: + if event.event_type != "model_request" or "routed_experts" not in event.data: + return event + data = dict(event.data) + data.pop("routed_experts") + return event.model_copy(update={"data": data}) + + def _dedupe_model_requests_by_prompt_token_ids(events: list[Event]) -> list[Event]: """Keep the last request for each valid, non-empty prompt-token key.""" last_index_by_prompt: dict[tuple[int, ...], int] = {} @@ -202,6 +238,7 @@ async def query_events( rollout_id: str, event_type: str | None = None, format: str | None = Query(None, description="Set to 'triplet' to trim events for RL training"), + include_routed_experts: bool = True, ) -> list[Event]: """Query events for the default rollout attempt.""" events = _query_events( @@ -211,4 +248,6 @@ async def query_events( if format == "triplet": events = [_to_triplet_format(e) for e in events] events = _dedupe_model_requests_by_prompt_token_ids(events) + if not include_routed_experts: + events = [_without_routed_experts(event) for event in events] return events diff --git a/agentlightning/verl/agl_rollout_manager.py b/agentlightning/verl/agl_rollout_manager.py index 709459220..cf0c3c3c3 100644 --- a/agentlightning/verl/agl_rollout_manager.py +++ b/agentlightning/verl/agl_rollout_manager.py @@ -133,6 +133,14 @@ def _to_native(obj: Any) -> Any: return obj +def _without_routed_experts(event: Event) -> Event: + if event.event_type != "model_request" or "routed_experts" not in event.data: + return event + data = dict(event.data) + data.pop("routed_experts") + return event.model_copy(update={"data": data}) + + # [multimodal-patch] Extract image URLs from OpenAI-style chat messages, in order of # appearance (ported from agent-lightning v0.3.0 TripletAdapter.extract_prompt_image_urls). def _extract_image_urls_from_messages(messages: Any) -> list[str]: @@ -339,9 +347,22 @@ def _record_lifecycle_timestamps(enqueued_rollout: EnqueuedRollout, rollout: Rol if state in TERMINAL_STATES: enqueued_rollout.finished_at = updated_at - def _get_events(self, rollout_id: str, *, event_type: str | None = None, format: str | None = None) -> list[Event]: + def _get_events( + self, + rollout_id: str, + *, + event_type: str | None = None, + format: str | None = None, + include_routed_experts: bool = True, + ) -> list[Event]: params = { - key: value for key, value in {"event_type": event_type, "format": format}.items() if value is not None + key: value + for key, value in { + "event_type": event_type, + "format": format, + "include_routed_experts": include_routed_experts, + }.items() + if value is not None } response = self.client.get(f"/api/rollouts/{rollout_id}/events", params=params) response.raise_for_status() @@ -407,7 +428,7 @@ def _create_rollouts(self, data: dict[str, Any], *, is_train: bool) -> list[Enqu ] def _fetch_rollout_events(self, rollout_id: str) -> tuple[list[Event], list[Event]]: - raw_events = self._get_events(rollout_id) + raw_events = self._get_events(rollout_id, include_routed_experts=False) triplet_events = self._get_events(rollout_id, format="triplet") return raw_events, triplet_events @@ -425,7 +446,7 @@ def _run_succeeded_hook(self, rollout: Rollout) -> None: return attempt_id = rollout.status.last_attempt_id or "unknown" trace_event_helper = _TraceEventHelper() - raw_events = self._get_events(rollout.rollout_id) + raw_events = self._get_events(rollout.rollout_id, include_routed_experts=False) events_by_attempt = self._events_by_attempt(raw_events, attempt_id) try: self._hooks.on_succeeded(rollout, events_by_attempt, trace_event_helper) @@ -466,6 +487,7 @@ def _build_completed_rollout(self, enqueued_rollout: EnqueuedRollout, rollout: R response={ "token_ids": response_token_ids, "log_probs": data.get("response_log_probs"), + "routed_experts": data.get("routed_experts"), }, reward=None, metadata={"server": data.get("server", {})}, @@ -494,6 +516,7 @@ def _build_completed_rollout(self, enqueued_rollout: EnqueuedRollout, rollout: R finished_at = enqueued_rollout.finished_at if finished_at is None: finished_at = rollout.status.updated_at + diagnostic_triplet_events = [_without_routed_experts(event).model_dump() for event in triplet_events] return CompletedRollout( rollout_id=enqueued_rollout.rollout_id, data_id=enqueued_rollout.data_id, @@ -507,7 +530,7 @@ def _build_completed_rollout(self, enqueued_rollout: EnqueuedRollout, rollout: R triplets=triplets, metadata=metadata, events=[event.model_dump() for event in raw_events], - triplet_events=[event.model_dump() for event in triplet_events], + triplet_events=diagnostic_triplet_events, rollout_state=rollout.status.state, error_message=rollout.status.error_message, ) diff --git a/agentlightning/verl/rollout_adapter.py b/agentlightning/verl/rollout_adapter.py index 2c332a01e..18f2e1cca 100644 --- a/agentlightning/verl/rollout_adapter.py +++ b/agentlightning/verl/rollout_adapter.py @@ -185,6 +185,34 @@ def get_right_padded_ids_and_attention_mask( return ids + [pad_token_id] * pad_len, [1] * seq_len + [0] * pad_len +def _build_routed_experts_batch( + rows: list[tuple[str, int, int, int]], + max_prompt_length: int, + max_response_length: int, + device: torch.device, +) -> torch.Tensor: + decoded = [np.load(io.BytesIO(base64.b64decode(payload)), allow_pickle=False) for payload, _, _, _ in rows] + batch = torch.zeros( + (len(rows), max_prompt_length + max_response_length, *decoded[0].shape[1:]), + dtype=torch.uint8, + device=device, + ) + for index, (routes, (_, original_prompt_length, prompt_length, response_length)) in enumerate( + zip(decoded, rows, strict=True) + ): + if len(routes) < original_prompt_length + response_length - 1: + raise RuntimeError("R3 routed_experts is shorter than its token sequence") + routes = torch.from_numpy(routes).to(device=device, dtype=torch.uint8) + prompt_routes = min(prompt_length, len(routes)) + prompt_start = max_prompt_length - prompt_length + batch[index, prompt_start : prompt_start + prompt_routes] = routes[:prompt_routes] + response_routes = min(response_length, max(len(routes) - original_prompt_length, 0)) + batch[index, max_prompt_length : max_prompt_length + response_routes] = routes[ + original_prompt_length : original_prompt_length + response_routes + ] + return batch + + # --------------------------------------------------------------------------- # [multimodal-patch] Multimodal (image) support for mrope VLM training. # Mirrors verl 0.8.0's AgentLoopWorker._compute_multi_modal_inputs / @@ -349,21 +377,25 @@ def __init__( *, max_prompt_length: int, max_response_length: int, + max_total_length: int | None = None, device: torch.device, pad_token_id: int, reward_fillna_value: float = 0.0, trace_aggregator_level: str = "transition", tokenizer: Any | None = None, processor: Any | None = None, # [multimodal-patch] HF processor; None keeps text-only behavior + require_routed_experts: bool = False, ) -> None: self.max_prompt_length = max_prompt_length self.max_response_length = max_response_length + self.max_total_length = max_total_length self.device = device self.pad_token_id = pad_token_id self.reward_fillna_value = reward_fillna_value self.trace_aggregator_level = trace_aggregator_level self.tokenizer = tokenizer self.processor = processor # [multimodal-patch] + self.require_routed_experts = require_routed_experts def get_train_data_batch( self, @@ -409,6 +441,7 @@ def get_train_data_batch( turn_index_list: list[int] = [] is_drop_list: list[bool] = [] response_log_probs_list: list[list[float] | None] = [] + routed_experts_rows: list[tuple[str, int, int, int] | None] = [] image_urls_list: list[list[str] | None] = [] # [multimodal-patch] per kept training row n_trunc_sample_because_of_response = 0 n_skipped_empty_training_rows = 0 @@ -426,21 +459,26 @@ def append_training_row( reward: float, response_mask: list[int] | None = None, response_log_probs: list[float] | None = None, + routed_experts: str | None = None, image_urls: list[str] | None = None, # [multimodal-patch] ) -> None: nonlocal n_skipped_empty_training_rows, n_trunc_sample_because_of_response + original_prompt_length = len(prompt_ids) if len(prompt_ids) > self.max_prompt_length: prompt_ids = prompt_ids[: self.max_prompt_length] is_drop = True else: is_drop = False - if len(response_ids) > self.max_response_length: - response_ids = response_ids[: self.max_response_length] + response_limit = self.max_response_length + if self.max_total_length is not None: + response_limit = min(response_limit, max(self.max_total_length - len(prompt_ids), 0)) + if len(response_ids) > response_limit: + response_ids = response_ids[:response_limit] if response_mask is not None: - response_mask = response_mask[: self.max_response_length] + response_mask = response_mask[:response_limit] if response_log_probs is not None: - response_log_probs = response_log_probs[: self.max_response_length] + response_log_probs = response_log_probs[:response_limit] n_trunc_sample_because_of_response += 1 if response_log_probs is not None and len(response_log_probs) != len(response_ids): @@ -450,6 +488,8 @@ def append_training_row( if train_token_count == 0: n_skipped_empty_training_rows += 1 return + if self.require_routed_experts and routed_experts is None: + raise RuntimeError(f"R3 requires routed_experts for rollout {rollout_id}") one_input_ids, one_input_attention_mask = get_left_padded_ids_and_attention_mask( prompt_ids, self.max_prompt_length, self.pad_token_id @@ -470,6 +510,11 @@ def append_training_row( response_mask_list.append(one_response_mask) response_log_probs_list.append(response_log_probs) + routed_experts_rows.append( + (routed_experts, original_prompt_length, len(prompt_ids), len(response_ids)) + if routed_experts is not None + else None + ) reward_list.append(reward) data_id_list.append(data_id) @@ -501,6 +546,7 @@ def append_training_row( response_ids=response_ids, reward=final_reward, response_log_probs=log_probs, + routed_experts=triplet.response.get("routed_experts"), image_urls=triplet.image_urls, # [multimodal-patch] ) continue @@ -512,6 +558,7 @@ def append_training_row( current_context = current_prompt_ids + current_response_ids current_response_mask = [1] * len(current_response_ids) current_response_log_probs: list[float] | None = first_triplet.response["log_probs"] + current_routed_experts = first_triplet.response.get("routed_experts") response_len_per_turn_list.append(len(current_response_ids)) merged_group_count = 0 @@ -537,6 +584,7 @@ def append_training_row( else: current_response_log_probs += list(log_probs) current_context = next_context + current_routed_experts = triplet.response.get("routed_experts") continue if len(merge_mismatch_rows) < _TRACE_MERGE_MISMATCH_WANDB_LIMIT: @@ -568,6 +616,7 @@ def append_training_row( reward=final_reward, response_mask=current_response_mask, response_log_probs=current_response_log_probs, + routed_experts=current_routed_experts, ) merged_group_count += 1 @@ -577,6 +626,7 @@ def append_training_row( current_response_ids = list(response_ids) current_response_mask = [1] * len(response_ids) current_response_log_probs = log_probs + current_routed_experts = triplet.response.get("routed_experts") append_training_row( rollout_id=rollout.rollout_id, @@ -587,6 +637,7 @@ def append_training_row( reward=final_reward, response_mask=current_response_mask, response_log_probs=current_response_log_probs, + routed_experts=current_routed_experts, ) merged_group_count += 1 @@ -747,6 +798,13 @@ def append_training_row( if log_probs is not None ] batch_dict["rollout_log_probs"] = torch.tensor(padded_log_probs_list, dtype=torch.float32).to(self.device) + if routed_experts_rows and all(row is not None for row in routed_experts_rows): + batch_dict["routed_experts"] = _build_routed_experts_batch( + [row for row in routed_experts_rows if row is not None], + self.max_prompt_length, + self.max_response_length, + self.device, + ) batch = TensorDict(batch_dict, batch_size=n_sample) # type: ignore[arg-type] data_proto = DataProto(batch=batch) diff --git a/agentlightning/verl/trainer.py b/agentlightning/verl/trainer.py index 57575e7ba..8b1d3a827 100644 --- a/agentlightning/verl/trainer.py +++ b/agentlightning/verl/trainer.py @@ -407,10 +407,14 @@ def _rollout(self, gen_batch: DataProto, is_train: bool) -> tuple[DataProto, dic if str(level).startswith("trajectory") else self.config.data.max_response_length ) + max_total_length = ( + trace_aggregator.get("trajectory_max_total_length") if str(level).startswith("trajectory") else None + ) pad_token_id = self.tokenizer.pad_token_id if self.tokenizer.pad_token_id is not None else 0 rollout_adapter = RolloutAdapter( max_prompt_length=max_prompt_length, max_response_length=max_response_length, + max_total_length=max_total_length, device=torch.device("cpu"), pad_token_id=pad_token_id, reward_fillna_value=self.config.agentlightning.reward_fillna_value, @@ -419,6 +423,14 @@ def _rollout(self, gen_batch: DataProto, is_train: bool) -> tuple[DataProto, dic # [multimodal-patch] RayPPOTrainer stores the processor from entrypoint; forwarding it # enables pixel_values + mrope position ids for image-bearing training rows. processor=getattr(self, "processor", None), + require_routed_experts=( + OmegaConf.select( + self.config, + "actor_rollout_ref.actor.megatron.router_replay.mode", + default="disabled", + ) + == "R3" + ), ) if is_train: diff --git a/docs/76-example-coding-agent-moe.md b/docs/76-example-coding-agent-moe.md new file mode 100644 index 000000000..4b895e4a2 --- /dev/null +++ b/docs/76-example-coding-agent-moe.md @@ -0,0 +1,56 @@ +# Coding Agent: MoE + +| GPU | Model | Actor | Rollout | Router replay | +|---|---|---|---|---| +| 4× B200 | `Qwen/Qwen3.5-35B-A3B` | Megatron | vLLM | R3 | + +This is the MoE variant of the [Coding Agent](75-example-coding-agent.md) example. It reuses the same +SWE-smith data, Kubernetes controller, repository images, agent, and reward. The existing +`Qwen/Qwen3.5-9B` FSDP path remains available through `examples/swe_smith/run.sh`. + +## Why R3 + +An MoE token can select different experts during rollout and training. R3 records vLLM's rollout +routing and passes it through the Agent Lightning event, triplet, and `DataProto` pipeline so the +Megatron actor update replays the same experts. + +The MoE launcher enables both sides: + +```text +actor_rollout_ref.rollout.enable_rollout_routing_replay=true +actor_rollout_ref.actor.megatron.router_replay.mode=R3 +``` + +## Run + +Use the MoE wrapper for all three roles: + +```bash +# Machine B +export AGL_SERVER_PUBLIC_HOST= +export AGL_KEY= +examples/swe_smith/run_moe.sh server + +# Machine A +export AGL_SERVER_PUBLIC_HOST= +export AGL_KEY= +export AGL_NAMESPACE=agents +examples/swe_smith/run_moe.sh controller + +# Machine B +export AGL_KEY= +examples/swe_smith/run_moe.sh trainer +``` + +`run_moe.sh` selects `train_smith_agent_moe.py`, requests routed experts from the Gateway, and keeps +the standard 9B launcher unchanged. R3 route compaction requires the default trajectory aggregator; +do not override it to `transition`. Other additional arguments are VERL dotlist overrides: + +```bash +examples/swe_smith/run_moe.sh trainer \ + trainer.total_training_steps=100 \ + actor_rollout_ref.rollout.n=4 +``` + +The default topology is PP=1, TP=1, EP=4, and ETP=1 with parameter, optimizer, and gradient offload. +Treat it as a four-B200 starting point and tune batch sizes for the available memory. diff --git a/docs/README.md b/docs/README.md index 2d0609de6..a9d83ca7a 100644 --- a/docs/README.md +++ b/docs/README.md @@ -42,4 +42,5 @@ For the legacy Agent Lightning releases earlier than v1.0, see the [`v0.x` code | [Search-R1](65-example-search-r1.md) | Train a multi-turn retrieval and reasoning agent. | | [LLM-in-Sandbox](70-example-llm-in-sandbox.md) | Train a general agent with computer and code execution tools. | | [Coding Agent](75-example-coding-agent.md) | Train a coding agent using repository tests as feedback. | +| [Coding Agent: MoE](76-example-coding-agent-moe.md) | Train Qwen3.5-35B-A3B with Megatron and R3. | | [Multimodal QA](80-example-multimodal-qa.md) | Train a vision-language model on synthetic image QA. | diff --git a/examples/swe_smith/README.md b/examples/swe_smith/README.md index 66107c2ff..2a63da9dc 100644 --- a/examples/swe_smith/README.md +++ b/examples/swe_smith/README.md @@ -1,3 +1,5 @@ # Coding Agent -See the [Coding Agent documentation](../../docs/75-example-coding-agent.md) for data filtering, distributed setup, repository image preparation, training, and reward-hacking protections. +See the [Coding Agent documentation](../../docs/75-example-coding-agent.md) for the Qwen3.5-9B +example, or [Coding Agent: MoE](../../docs/76-example-coding-agent-moe.md) for Qwen3.5-35B-A3B with +Megatron and R3. diff --git a/examples/swe_smith/run.sh b/examples/swe_smith/run.sh index 386db777a..79f14fbb5 100755 --- a/examples/swe_smith/run.sh +++ b/examples/swe_smith/run.sh @@ -18,6 +18,8 @@ EXAMPLE_DIR="examples/swe_smith" AGL_SERVER_PORT="${AGL_SERVER_PORT:-8080}" AGL_KEY="${AGL_KEY:-dummy}" AGL_MODEL_NAME="${AGL_MODEL_NAME:-Qwen/Qwen3-8B}" +AGL_TRAIN_SCRIPT="${AGL_TRAIN_SCRIPT:-$EXAMPLE_DIR/train_smith_agent.py}" +AGL_INCLUDE_ROUTED_EXPERTS="${AGL_INCLUDE_ROUTED_EXPERTS:-false}" AGL_NAMESPACE="${AGL_NAMESPACE:-default}" PUBLIC_HOST="${AGL_SERVER_PUBLIC_HOST:-0.0.0.0}" SERVER_URL="http://${PUBLIC_HOST}:${AGL_SERVER_PORT}" @@ -44,7 +46,8 @@ if [ "$ROLE" = "server" ]; then port="$AGL_SERVER_PORT" \ host="${AGL_SERVER_BIND:-0.0.0.0}" \ key="$AGL_KEY" \ - default_proxy.model_name="$AGL_MODEL_NAME" & + default_proxy.model_name="$AGL_MODEL_NAME" \ + default_proxy.include_routed_experts="$AGL_INCLUDE_ROUTED_EXPERTS" & SERVER_PID=$! cleanup() { if kill -0 "$SERVER_PID" 2>/dev/null; then @@ -80,7 +83,7 @@ elif [ "$ROLE" = "trainer" ]; then exit 1 fi echo "=== Running SWE-smith training ===" - python "$EXAMPLE_DIR/train_smith_agent.py" \ + python "$AGL_TRAIN_SCRIPT" \ --agl-base-url "http://localhost:$AGL_SERVER_PORT" \ --agl-key "$AGL_KEY" \ --train-dataset-path "$TRAIN_DATASET_PATH" \ diff --git a/examples/swe_smith/run_moe.sh b/examples/swe_smith/run_moe.sh new file mode 100755 index 000000000..c59efcc0a --- /dev/null +++ b/examples/swe_smith/run_moe.sh @@ -0,0 +1,9 @@ +#!/usr/bin/env bash +# Copyright (c) Microsoft. All rights reserved. + +set -euo pipefail +EXAMPLE_DIR="$(cd "$(dirname "$0")" && pwd)" +export AGL_MODEL_NAME="${AGL_MODEL_NAME:-Qwen/Qwen3.5-35B-A3B}" +export AGL_TRAIN_SCRIPT="$EXAMPLE_DIR/train_smith_agent_moe.py" +export AGL_INCLUDE_ROUTED_EXPERTS=true +exec "$EXAMPLE_DIR/run.sh" "$@" diff --git a/examples/swe_smith/train_smith_agent_moe.py b/examples/swe_smith/train_smith_agent_moe.py new file mode 100755 index 000000000..31eeb6e70 --- /dev/null +++ b/examples/swe_smith/train_smith_agent_moe.py @@ -0,0 +1,142 @@ +#!/usr/bin/env python3 +# Copyright (c) Microsoft. All rights reserved. + +from __future__ import annotations + +import argparse +from collections.abc import Sequence +from typing import cast + +from omegaconf import DictConfig, OmegaConf +from train_smith_agent import EXAMPLE_DIR, load_split_file +from train_smith_agent import build_config as build_fsdp_config + +MODEL = "Qwen/Qwen3.5-35B-A3B" + + +def build_config( + *, + model: str = MODEL, + agl_base_url: str = "http://localhost:8080", + agl_key: str = "", + run_name: str | None = None, + config_overrides: Sequence[str] = (), +) -> DictConfig: + config = build_fsdp_config( + model=model, + agl_base_url=agl_base_url, + agl_key=agl_key, + config_overrides=(), + ) + + config.algorithm.update( + { + "enable_rollout_level_advantage": True, + "enable_per_rollout_mean_loss": True, + "enable_cispo_loss": False, + "enable_rollout_level_advantage_scale": False, + } + ) + config.actor_rollout_ref.rollout.update( + { + "enable_rollout_routing_replay": True, + "calculate_log_probs": True, + "log_prob_use_dynamic_bsz": False, + "max_model_len": 81920, + } + ) + + actor = config.actor_rollout_ref.actor + actor.pop("fsdp_config", None) + actor.update( + { + "strategy": "megatron", + "model_engine": "megatron", + "use_dynamic_bsz": False, + "loss_agg_mode": "seq-mean-token-sum", + "policy_loss": {"loss_mode": "per_rollout_mean"}, + "megatron": { + "pipeline_model_parallel_size": 1, + "tensor_model_parallel_size": 1, + "expert_model_parallel_size": 4, + "expert_tensor_parallel_size": 1, + "param_offload": True, + "optimizer_offload": True, + "grad_offload": True, + "use_mbridge": True, + "router_replay": {"mode": "R3"}, + "override_transformer_config": { + "moe_enable_deepep": True, + "moe_token_dispatcher_type": "flex", + "moe_router_dtype": "fp32", + "recompute_method": "uniform", + "recompute_granularity": "full", + "recompute_num_layers": 1, + "moe_permute_fusion": False, + }, + }, + } + ) + + ref = config.actor_rollout_ref.ref + ref.pop("fsdp_config", None) + ref.update( + { + "log_prob_use_dynamic_bsz": False, + "megatron": { + "pipeline_model_parallel_size": 1, + "tensor_model_parallel_size": 1, + "expert_model_parallel_size": 4, + "expert_tensor_parallel_size": 1, + "param_offload": True, + }, + } + ) + + config.trainer.update( + { + "val_before_train": False, + "balance_batch": True, + "experiment_name": f"swe_smith_qwen35_35b_a3b_megatron_r3{f'_{run_name}' if run_name else ''}", + } + ) + config.agentlightning.trace_aggregator.update( + { + "level": "trajectory", + "trajectory_max_prompt_length": 24576, + "trajectory_max_response_length": 81920, + "trajectory_max_total_length": 81920, + } + ) + config.agentlightning.async_rollout.update({"enabled": False, "async_train_batch_size": None}) + return cast(DictConfig, OmegaConf.merge(config, OmegaConf.from_dotlist(list(config_overrides)))) + + +def main() -> None: + parser = argparse.ArgumentParser(description="Train Qwen3.5-35B-A3B with Megatron and R3") + parser.add_argument("--train-dataset-path", default=str(EXAMPLE_DIR / "train_dataset_mixed.jsonl")) + parser.add_argument("--val-dataset-path", default=str(EXAMPLE_DIR / "val_dataset_filtered.jsonl")) + parser.add_argument("--max-val-instances", type=int) + parser.add_argument("--model", default=MODEL) + parser.add_argument("--agl-base-url", default="http://localhost:8080") + parser.add_argument("--agl-key", default="") + parser.add_argument("--run-name") + args, overrides = parser.parse_known_args() + + train_dataset = load_split_file(args.train_dataset_path) + val_dataset = load_split_file(args.val_dataset_path, max_instances=args.max_val_instances) + config = build_config( + model=args.model, + agl_base_url=args.agl_base_url, + agl_key=args.agl_key, + run_name=args.run_name, + config_overrides=overrides, + ) + + from agentlightning.verl.entrypoint import run_ppo + + run_ppo(config=config, train_dataset=train_dataset, val_dataset=val_dataset) + + +if __name__ == "__main__": + main() diff --git a/mkdocs.yml b/mkdocs.yml index bd3e6fec0..f8a3c2f77 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -112,6 +112,7 @@ nav: - Search-R1: 65-example-search-r1.md - LLM-in-Sandbox: 70-example-llm-in-sandbox.md - Coding Agent: 75-example-coding-agent.md + - Coding Agent (MoE): 76-example-coding-agent-moe.md - Multimodal QA: 80-example-multimodal-qa.md extra_css: