Skip to content
Merged
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
2 changes: 1 addition & 1 deletion .github/workflows/cicd-main.yml
Original file line number Diff line number Diff line change
Expand Up @@ -301,7 +301,7 @@ jobs:
build-args: |
MAX_JOBS=4
NEMO_RL_COMMIT=${{ needs.pre-flight.outputs.test_sha }}
${{ needs.org-member-pre-flight.outputs.is_member != 'true' && 'SKIP_SGLANG_BUILD=1' || '' }}
SKIP_SGLANG_BUILD=1

update-uv-cache:
name: Update uv build cache
Expand Down
57 changes: 49 additions & 8 deletions nemo_rl/models/generation/vllm/vllm_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -122,6 +122,7 @@ def _maybe_process_fp8_kv_cache(self) -> None:
return

# FP8 KV cache: process KV scales after weight loading
from vllm.config import set_current_vllm_config
from vllm.model_executor.model_loader.utils import (
process_weights_after_loading,
)
Expand All @@ -130,11 +131,12 @@ def _maybe_process_fp8_kv_cache(self) -> None:
target_device = next(self.model_runner.model.parameters()).device

# Call process_weights_after_loading to handle KV scales
process_weights_after_loading(
self.model_runner.model,
self.model_runner.model_config,
target_device,
)
with set_current_vllm_config(self.model_runner.vllm_config):
process_weights_after_loading(
self.model_runner.model,
self.model_runner.model_config,
target_device,
)

@staticmethod
def _split_policy_and_draft_weights(
Expand All @@ -158,6 +160,42 @@ def _split_policy_and_draft_weights(
policy_weights.append((key, tensor))
return policy_weights, draft_weights

@staticmethod
def _trim_vocab_padding(
draft_model: torch.nn.Module,
draft_weights: list[tuple[str, torch.Tensor]],
) -> list[tuple[str, torch.Tensor]]:
"""Trim padded vocab dimensions from draft weights.

Megatron pads vocab to a multiple, but vLLM 0.20's autoloader
strictly asserts loaded_weight.shape[0] == org_vocab_size on
VocabParallelEmbedding layers. Each such layer may have a
different org_vocab_size (e.g. embed_tokens uses vocab_size
while lm_head uses draft_vocab_size), so we match each weight
to its target module by name.
"""
from vllm.model_executor.layers.vocab_parallel_embedding import (
VocabParallelEmbedding,
)

vocab_sizes: dict[str, int] = {}
for name, module in draft_model.named_modules():
if isinstance(module, VocabParallelEmbedding):
vocab_sizes[name] = module.org_vocab_size

if not vocab_sizes:
return draft_weights

trimmed = []
for key, tensor in draft_weights:
for mod_name, org_vocab_size in vocab_sizes.items():
leaf = mod_name.rsplit(".", 1)[-1]
if leaf in key and tensor.shape[0] > org_vocab_size:
tensor = tensor[:org_vocab_size]
break
trimmed.append((key, tensor))
return trimmed

def _load_draft_weights(
self, draft_weights: list[tuple[str, torch.Tensor]]
) -> None:
Expand All @@ -172,6 +210,7 @@ def _load_draft_weights(
"[draft] Received draft weights but vLLM drafter is unavailable; skipping draft update."
)
return
draft_weights = self._trim_vocab_padding(draft_model, draft_weights)
draft_model.load_weights(weights=draft_weights)

def _load_weights(self, weights):
Expand Down Expand Up @@ -217,13 +256,15 @@ def update_weights_via_ipc_zmq(self) -> bool:

if payload == IPCProtocol.COMPLETE:
# means the update is done
from vllm.config import set_current_vllm_config
from vllm.model_executor.model_loader.utils import (
process_weights_after_loading,
)

process_weights_after_loading(
self.model_runner.model, self.model_config, self.device
)
with set_current_vllm_config(self.model_runner.vllm_config):
process_weights_after_loading(
self.model_runner.model, self.model_config, self.device
)
self.zmq_socket.send(IPCProtocol.ACK.value.encode())
break

Expand Down
2 changes: 2 additions & 0 deletions nemo_rl/models/generation/vllm/vllm_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -259,6 +259,7 @@ def _patch_vllm_llama_eagle3_own_lm_head():
" self.lm_head = ParallelLMHead(\n"
" self.config.draft_vocab_size,\n"
" self.config.hidden_size,\n"
" quant_config=get_draft_quant_config(vllm_config),\n"
' prefix=maybe_prefix(prefix, "lm_head"),\n'
" )\n"
" self.logits_processor = LogitsProcessor(\n"
Expand All @@ -268,6 +269,7 @@ def _patch_vllm_llama_eagle3_own_lm_head():
" self.lm_head = ParallelLMHead(\n"
" self.config.draft_vocab_size,\n"
" self.config.hidden_size,\n"
" quant_config=get_draft_quant_config(vllm_config),\n"
' prefix=maybe_prefix(prefix, "lm_head"),\n'
" )\n"
" self.has_own_lm_head = (\n"
Expand Down
96 changes: 66 additions & 30 deletions nemo_rl/models/generation/vllm/vllm_worker_async.py
Original file line number Diff line number Diff line change
Expand Up @@ -317,9 +317,13 @@ def _setup_vllm_openai_api_server(self, app: FastAPI) -> FastAPI:
TokenizeCompletionRequest,
TokenizeResponse,
)
from vllm.entrypoints.serve.render.serving import (
OpenAIServingRender,
)
from vllm.entrypoints.serve.tokenize.serving import (
OpenAIServingTokenization,
)
from vllm.exceptions import VLLMValidationError
from vllm.tool_parsers.abstract_tool_parser import ToolParserManager
from vllm.v1.engine.async_llm import logger as vllm_async_llm_logger

Expand Down Expand Up @@ -358,7 +362,10 @@ def model_post_init(self, context):
return super().model_post_init(context)

class NeMoRLOpenAIServingMixin:
async def _preprocess_chat(
# vLLM 0.20 moved chat preprocessing from
# OpenAIServing._preprocess_chat to OpenAIServingRender.preprocess_chat,
# so this override now applies via the render subclass.
async def preprocess_chat(
self,
request,
messages,
Expand All @@ -367,85 +374,80 @@ async def _preprocess_chat(
default_template_kwargs,
tool_dicts=None,
tool_parser=None,
reasoning_parser=None,
*,
skip_mm_cache: bool = False,
):
# Materialize the message tool calls so we can deepcopy below.
for message in messages:
if message.get("tool_calls"):
message["tool_calls"] = list(message["tool_calls"])

# Deepcopy messages here since _preprocess_chat may be destructive.
messages_for_replace_prefix_tokens = deepcopy(messages)

# res is (conversation, [engine_prompt])
try:
res = await super()._preprocess_chat(
res = await super().preprocess_chat(
request=request,
messages=messages,
default_template=default_template,
default_template_content_format=default_template_content_format,
default_template_kwargs=default_template_kwargs,
tool_dicts=tool_dicts,
tool_parser=tool_parser,
reasoning_parser=reasoning_parser,
skip_mm_cache=skip_mm_cache,
)
except ValueError as e:
except (ValueError, VLLMValidationError) as e:
if "maximum context length" in str(e):
import logging

# Print a clean one-liner warning that max model length has been exceeded
# The exception is still raised, but later filtered out by the MaxContextLengthFilter
logging.getLogger(__name__).warning(
"Prompt exceeds max_model_len: %s", e
)
raise

if request.required_prefix_token_ids is None:
if (
not hasattr(request, "required_prefix_token_ids")
or request.required_prefix_token_ids is None
):
return res

# Find the last assistant message
last_assistant_message_idx = None
for i in reversed(range(len(messages_for_replace_prefix_tokens))):
if messages_for_replace_prefix_tokens[i]["role"] == "assistant":
last_assistant_message_idx = i
break

if last_assistant_message_idx is None:
# If there's no assistant message, we just use the entire thing.
messages_to_last_assistant_message = (
messages_for_replace_prefix_tokens
)
else:
# Include the last assistant message itself.
messages_to_last_assistant_message = (
messages_for_replace_prefix_tokens[
: last_assistant_message_idx + 1
]
)

# For the prefix token calculation, we need add_generation_prompt=False
# to get tokens up to (and including) the last assistant message only.
# add_generation_prompt is a field on the request that gets embedded
# into ChatParams via build_chat_params().
modified_request = request.model_copy(
update={"add_generation_prompt": False}
)

# Call the actual preprocess chat subroutine so we don't miss anything. Whatever they do is whatever we do since we literally do what they do.
corresponding_res = await super()._preprocess_chat(
corresponding_res = await super().preprocess_chat(
request=modified_request,
messages=messages_to_last_assistant_message,
default_template=default_template,
default_template_content_format=default_template_content_format,
default_template_kwargs=default_template_kwargs,
tool_dicts=tool_dicts,
tool_parser=tool_parser,
reasoning_parser=reasoning_parser,
skip_mm_cache=skip_mm_cache,
)
actual_corresponding_token_ids = corresponding_res[1][0][
"prompt_token_ids"
]

engine_prompt = res[1][
0
] # We need to modify engine_prompt.prompt_token_ids
engine_prompt = res[1][0]

final_prompt_token_ids = _replace_prefix_tokens(
tokenizer=self.renderer.tokenizer,
Expand All @@ -468,23 +470,41 @@ class NeMoRLChatCompletionRequest(
):
required_prefix_token_ids: Optional[List[int]] = None

# This MRO is necessary i.e. NeMoRLOpenAIServingMixin > OpenAIServingChat
class NeMoRLOpenAIServingChat(NeMoRLOpenAIServingMixin, OpenAIServingChat):
# vLLM 0.20 routes both /v1/chat/completions and /tokenize through
# OpenAIServingRender.preprocess_chat, so the prefix-token override
# belongs on the render subclass.
class NeMoRLOpenAIServingChat(OpenAIServingChat):
pass

class NeMoRLOpenAIServingRender(NeMoRLOpenAIServingMixin, OpenAIServingRender):
pass

serving_chat_default_kwargs = dict(
response_role="assistant",
request_logger=None,
chat_template=None,
chat_template_content_format="auto",
enable_auto_tools=True,
)
serving_chat_kwargs = serving_chat_default_kwargs | self.cfg["vllm_cfg"].get(
"http_server_serving_chat_kwargs", dict()
)
openai_serving_render = NeMoRLOpenAIServingRender(
model_config=engine_client.model_config,
renderer=engine_client.renderer,
model_registry=openai_serving_models.registry,
request_logger=serving_chat_kwargs["request_logger"],
chat_template=serving_chat_kwargs["chat_template"],
chat_template_content_format=serving_chat_kwargs[
"chat_template_content_format"
],
enable_auto_tools=serving_chat_kwargs["enable_auto_tools"],
)
serving_chat_kwargs.update(
dict(
engine_client=engine_client,
models=openai_serving_models,
openai_serving_render=openai_serving_render,
return_tokens_as_token_ids=True,
)
)
Expand All @@ -509,9 +529,25 @@ async def create_chat_completion(
assert request.temperature == generation_config["temperature"]
assert request.top_p == generation_config["top_p"]

generator = await openai_serving_chat.create_chat_completion(
request, raw_request
)
try:
generator = await openai_serving_chat.create_chat_completion(
request, raw_request
)
except VLLMValidationError as e:
# vLLM 0.20 raises VLLMValidationError for prompts exceeding
# max_model_len during tokenization, instead of returning an
# ErrorResponse. Convert to HTTP 400 so the Gym proxy can
# detect context-length overflow and handle it gracefully.
return JSONResponse(
content={
"error": {
"message": str(e),
"type": "invalid_request_error",
"code": 400,
}
},
status_code=400,
)

if isinstance(generator, ErrorResponse):
return JSONResponse(
Expand All @@ -537,10 +573,9 @@ class NeMoRLTokenizeChatRequest(
TokenizeCompletionRequest, NeMoRLTokenizeChatRequest
]

# This MRO is necessary i.e. NeMoRLOpenAIServingMixin > OpenAIServingTokenization
class NeMoRLOpenAIServingTokenization(
NeMoRLOpenAIServingMixin, OpenAIServingTokenization
):
# Tokenize path delegates to OpenAIServingRender.preprocess_chat in
# vLLM 0.20, where the prefix-token override lives.
class NeMoRLOpenAIServingTokenization(OpenAIServingTokenization):
pass

serving_tokenization_kwargs = dict(
Expand All @@ -551,6 +586,7 @@ class NeMoRLOpenAIServingTokenization(
],
engine_client=serving_chat_kwargs["engine_client"],
models=serving_chat_kwargs["models"],
openai_serving_render=openai_serving_render,
)
openai_serving_tokenization = NeMoRLOpenAIServingTokenization(
**serving_tokenization_kwargs
Expand Down
7 changes: 5 additions & 2 deletions nemo_rl/models/policy/workers/megatron_policy_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -450,8 +450,11 @@ def train(
"all_mb_metrics": mb_metrics,
"grad_norm": torch.tensor([grad_norm]),
}
# Collect MoE aux metrics averaged across microbatches
num_moe_experts = getattr(self.model.config, "num_moe_experts", None)
# Read "config" via getattr-by-string so the token stays out of
# train.__code__.co_names; with torch 2.11 cloudpickle otherwise
# matches torch.distributed.config (a non-pickleable ConfigModuleInstance).
model_config = getattr(self.model, "config", None)
num_moe_experts = getattr(model_config, "num_moe_experts", None)
if num_moe_experts is not None and num_moe_experts > 1:
moe_loss_scale = 1.0 / max(1, total_num_microbatches)
moe_metrics = get_moe_metrics(
Expand Down
Loading
Loading