diff --git a/docs/en/advanced/selected-logprob-provider.md b/docs/en/advanced/selected-logprob-provider.md index ade506eac..f49190f94 100644 --- a/docs/en/advanced/selected-logprob-provider.md +++ b/docs/en/advanced/selected-logprob-provider.md @@ -22,3 +22,19 @@ loaded or raises `SelectedLogprobProviderUnavailable`. `strict` rejects that case. Provider exceptions and invalid result shapes always fail the run. A strict provider result must include non-empty `backend_id` and `contract_id` and remain connected to autograd when logits require gradients. + +For the RL-Kernel WS2 provider, the validated launch contract is: + +```bash +--tensor-model-parallel-size 2 \ +--context-parallel-size 2 \ +--rollout-top-p 1.0 \ +--selected-logprob-provider rl_engine.integrations.vime.logp.provider \ +--selected-logprob-provider-mode strict +``` + +The provider owns the TP vocabulary reduction only. Vime continues to own CP +token-row layout, response extraction, and PPO/GRPO loss composition. The +provider does not claim attention or FFN train/rollout consistency; those +claims require runtime readback from both Megatron and vLLM and are reported by +the RL-Kernel validation example. diff --git a/scripts/run-qwen3-8B-rlkernel-tp2-cp2.sh b/scripts/run-qwen3-8B-rlkernel-tp2-cp2.sh new file mode 100644 index 000000000..47023f93b --- /dev/null +++ b/scripts/run-qwen3-8B-rlkernel-tp2-cp2.sh @@ -0,0 +1,167 @@ +#!/usr/bin/env bash +# Qwen3-8B GRPO smoke/validation run with the RL-Kernel selected-logprob +# provider. Vime remains the launcher; RL-Kernel owns the provider and its +# contract. This script intentionally keeps the framework-side change small. + +set -euo pipefail + +VIME_ROOT="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")/.." && pwd)" +RL_KERNEL_ROOT="${RL_KERNEL_ROOT:-${VIME_ROOT}/../RL-Kernel}" + +if [[ ! -f "${RL_KERNEL_ROOT}/rl_engine/integrations/vime/logp.py" ]]; then + echo "RL_KERNEL_ROOT must point to an RL-Kernel checkout containing the Vime provider" >&2 + exit 2 +fi + +export PYTHONUNBUFFERED=1 +export VIME_RL_KERNEL_STRICT="${VIME_RL_KERNEL_STRICT:-1}" +MEGATRON_ROOT="${MEGATRON_ROOT:-/root/Megatron-LM}" +export PYTHONPATH="${RL_KERNEL_ROOT}:${VIME_ROOT}:${MEGATRON_ROOT}:${PYTHONPATH:-}" + +# The provider is a vocab-parallel TP implementation. CP owns token rows and +# must not be used as a vocabulary reduction group. +TP_SIZE="${TP_SIZE:-2}" +CP_SIZE="${CP_SIZE:-2}" +ACTOR_GPUS="${ACTOR_GPUS:-4}" +ROLLOUT_GPUS="${ROLLOUT_GPUS:-4}" +NUM_GPUS="${NUM_GPUS:-8}" +ROLLOUT_GPUS_PER_ENGINE="${ROLLOUT_GPUS_PER_ENGINE:-2}" +ROLLOUT_TOP_P="${ROLLOUT_TOP_P:-1.0}" +COLOCATE="${COLOCATE:-0}" + +if [[ "${NUM_GPUS}" != "8" || "${ACTOR_GPUS}" != "4" || "${ROLLOUT_GPUS}" != "4" ]]; then + echo "This validation entry point requires an 8-GPU node with 4 actor GPUs and 4 rollout GPUs" >&2 + exit 2 +fi +if [[ "${COLOCATE}" != "0" && "${COLOCATE}" != "1" ]]; then + echo "COLOCATE must be 0 (default, disjoint train/rollout GPUs) or 1" >&2 + exit 2 +fi + +if [[ "${TP_SIZE}" != "2" || "${CP_SIZE}" != "2" ]]; then + echo "This validation entry point is intentionally fixed to TP=2, CP=2" >&2 + exit 2 +fi +if [[ "${ROLLOUT_TOP_P}" != "1.0" ]]; then + echo "RL-Kernel strict selected-logprob validation requires ROLLOUT_TOP_P=1.0" >&2 + exit 2 +fi + +source "${VIME_ROOT}/scripts/models/qwen3-8B.sh" + +MODEL_ROOT="${MODEL_ROOT:-/root/Qwen3-8B}" +TORCH_DIST_ROOT="${TORCH_DIST_ROOT:-/root/Qwen3-8B_torch_dist}" +VIME_CKPT="${VIME_CKPT:-/root/Qwen3-8B_vime_rlkernel_tp2_cp2}" +PROMPT_DATA="${PROMPT_DATA:-/root/dapo-math-17k/dapo-math-17k.jsonl}" + +if ! command -v nvidia-smi >/dev/null 2>&1; then + echo "nvidia-smi is required; refusing to run the CUDA validation on an unknown device" >&2 + exit 3 +fi +GPU_NAMES="$(nvidia-smi --query-gpu=name --format=csv,noheader 2>/dev/null || true)" +GPU_COUNT="$(printf '%s\n' "${GPU_NAMES}" | sed '/^$/d' | wc -l | tr -d ' ')" +if [[ "${GPU_COUNT}" != "${NUM_GPUS}" ]]; then + echo "Expected ${NUM_GPUS} visible GPUs, found ${GPU_COUNT}" >&2 + printf '%s\n' "${GPU_NAMES}" >&2 + exit 3 +fi +if [[ "${GPU_REQUIRE_H100:-1}" == "1" ]] && ! printf '%s\n' "${GPU_NAMES}" | grep -q 'H100'; then + echo "Expected H100 GPUs; refusing to run on a different GPU class" >&2 + printf '%s\n' "${GPU_NAMES}" >&2 + exit 3 +fi +python3 - <<'PY' +import torch + +if not torch.cuda.is_available() or torch.cuda.device_count() != 8: + raise SystemExit("PyTorch must expose 8 CUDA devices for this validation") +PY +for required_path in "${MODEL_ROOT}" "${TORCH_DIST_ROOT}" "${PROMPT_DATA}" "${MEGATRON_ROOT}"; do + if [[ ! -e "${required_path}" ]]; then + echo "Required runtime path does not exist: ${required_path}" >&2 + exit 3 + fi +done +python3 - <<'PY' +from rl_engine.integrations.vime.logp import provider +print(f"RL-Kernel provider import OK: {provider.__module__}.{provider.__name__}") +PY + +CKPT_ARGS=( + --hf-checkpoint "${MODEL_ROOT}" + --ref-load "${TORCH_DIST_ROOT}" + --load "${VIME_CKPT}" + --save "${VIME_CKPT}" + --save-interval 100000 +) + +ROLLOUT_ARGS=( + --prompt-data "${PROMPT_DATA}" + --input-key prompt + --label-key label + --apply-chat-template + --rollout-shuffle + --rm-type deepscaler + --num-rollout "${NUM_ROLLOUT:-1}" + --rollout-batch-size "${ROLLOUT_BATCH_SIZE:-8}" + --n-samples-per-prompt "${N_SAMPLES_PER_PROMPT:-2}" + --rollout-max-response-len "${MAX_RESPONSE_LEN:-1024}" + --rollout-temperature 1.0 + --rollout-top-p "${ROLLOUT_TOP_P}" + --global-batch-size "${GLOBAL_BATCH_SIZE:-16}" + --balance-data +) + +PARALLEL_ARGS=( + --tensor-model-parallel-size "${TP_SIZE}" + --context-parallel-size "${CP_SIZE}" + --pipeline-model-parallel-size 1 + --sequence-parallel + --expert-model-parallel-size 1 + --expert-tensor-parallel-size 1 + --use-dynamic-batch-size + --max-tokens-per-gpu "${MAX_TOKENS_PER_GPU:-2048}" +) + +RL_KERNEL_ARGS=( + --selected-logprob-provider rl_engine.integrations.vime.logp.provider + --selected-logprob-provider-mode strict + --custom-megatron-init-path rl_engine.integrations.megatron_runtime.initialize_from_environment +) + +MISC_ARGS=( + --attention-dropout 0.0 + --hidden-dropout 0.0 + --attention-softmax-in-fp32 + --attention-backend flash + --no-gradient-accumulation-fusion + --rollout-num-gpus-per-engine "${ROLLOUT_GPUS_PER_ENGINE}" + --vllm-gpu-memory-utilization "${VLLM_GPU_MEMORY_UTILIZATION:-0.4}" +) + +ray stop --force || true +ray start --head --node-ip-address "${MASTER_ADDR:-127.0.0.1}" \ + --num-gpus "${NUM_GPUS}" --disable-usage-stats \ + --dashboard-host=0.0.0.0 --dashboard-port="${RAY_DASHBOARD_PORT:-8265}" + +TRAIN_LAYOUT_ARGS=() +if [[ "${COLOCATE}" == "1" ]]; then + TRAIN_LAYOUT_ARGS+=(--colocate) +else + TRAIN_LAYOUT_ARGS+=(--megatron-to-hf-mode bridge) +fi + +ray job submit --address="http://127.0.0.1:${RAY_DASHBOARD_PORT:-8265}" \ + --working-dir "${VIME_ROOT}" \ + -- python3 train.py \ + --train-backend megatron \ + --actor-num-nodes 1 \ + --actor-num-gpus-per-node "${ACTOR_GPUS}" \ + --rollout-num-gpus "${ROLLOUT_GPUS}" \ + "${TRAIN_LAYOUT_ARGS[@]}" \ + "${MODEL_ARGS[@]}" \ + "${CKPT_ARGS[@]}" \ + "${ROLLOUT_ARGS[@]}" \ + "${PARALLEL_ARGS[@]}" \ + "${RL_KERNEL_ARGS[@]}" \ + "${MISC_ARGS[@]}" diff --git a/tests/test_logprob_response_spans.py b/tests/test_logprob_response_spans.py index d7ea13cd4..5aa176ddd 100644 --- a/tests/test_logprob_response_spans.py +++ b/tests/test_logprob_response_spans.py @@ -5,7 +5,13 @@ import torch from megatron.core import mpu -from vime.backends.megatron_utils.loss import _build_topp_keep_mask, get_rollout_top_p_logprob_kwargs +from vime.backends.megatron_utils.loss import ( + _build_topp_keep_mask, + _maybe_capture_log_probs, + drain_captured_log_probs, + enable_log_prob_capture, + get_rollout_top_p_logprob_kwargs, +) NUM_GPUS = 0 @@ -97,5 +103,20 @@ def test_top_p_mask_aligns_with_cp1_response_rows(monkeypatch): assert masked_rows == {2: [13], 3: [14], 5: [21], 6: [22], 7: [23]} +@pytest.mark.unit +def test_logprob_capture_uses_partition_keys_and_detaches_values(): + enable_log_prob_capture() + first = torch.tensor([1.0, 2.0], requires_grad=True) + second = torch.tensor([3.0], requires_grad=True) + + _maybe_capture_log_probs({"partition": [7, 3]}, [first, second]) + captured = drain_captured_log_probs() + + assert set(captured) == {3, 7} + torch.testing.assert_close(captured[7], first) + torch.testing.assert_close(captured[3], second) + assert not captured[7].requires_grad + + if __name__ == "__main__": raise SystemExit(pytest.main([__file__])) diff --git a/tests/test_megatron_argument_validation.py b/tests/test_megatron_argument_validation.py index a9f87c8e9..96422acb6 100644 --- a/tests/test_megatron_argument_validation.py +++ b/tests/test_megatron_argument_validation.py @@ -336,6 +336,8 @@ def add_cli_args(parser, **_kwargs): module.get_vime_extra_args_provider()(parser) args = parser.parse_args( [ + "--rollout-batch-size", + "1", "--selected-logprob-provider", "rl_engine.integrations.vime.logp.provider", "--selected-logprob-provider-mode", @@ -347,5 +349,17 @@ def add_cli_args(parser, **_kwargs): assert args.selected_logprob_provider_mode == "strict" +@pytest.mark.unit +def test_strict_selected_logprob_provider_rejects_top_p_replay(monkeypatch): + module = load_arguments_module(monkeypatch) + args = argparse.Namespace( + selected_logprob_provider="rl_engine.integrations.vime.logp.provider", + selected_logprob_provider_mode="strict", + rollout_top_p=0.9, + ) + with pytest.raises(ValueError, match="rollout-top-p 1.0"): + module._validate_selected_logprob_provider_args(args) + + if __name__ == "__main__": raise SystemExit(pytest.main([__file__])) diff --git a/tests/test_selected_logprob_provider.py b/tests/test_selected_logprob_provider.py index c9ac8d480..9f878d401 100644 --- a/tests/test_selected_logprob_provider.py +++ b/tests/test_selected_logprob_provider.py @@ -35,6 +35,27 @@ def _native(*args, **kwargs): return logits[:, :1], entropy +def _structural_request(**overrides) -> SelectedLogprobRequest: + values = dict( + logits=torch.randn(3, 5), + target_ids=torch.tensor([1, 2, 3]), + tensor_parallel_group=None, + context_parallel=ContextParallelLayout(world_size=1, rank=0, layout="single"), + with_entropy=False, + with_entropy_grad=False, + chunk_size=64, + hidden=torch.randn(3, 4), + lm_head_weight=torch.randn(5, 4), + lm_head_bias=torch.randn(5), + vocab_start_index=0, + global_vocab_size=5, + real_vocab_size=4, + temperature=torch.ones(3), + ) + values.update(overrides) + return SelectedLogprobRequest(**values) + + def _install_provider(monkeypatch, provider): module_name = "selected_logprob_provider_fixture" module = types.ModuleType(module_name) @@ -81,6 +102,50 @@ def provider(actual_request): torch.testing.assert_close(entropy, request.logits.sum(dim=-1)) +def test_structural_request_accepts_aligned_hidden_and_lm_head(): + request = _structural_request() + + assert request.hidden is not None and request.hidden.shape == (3, 4) + assert request.lm_head_weight is not None and request.lm_head_weight.shape == (5, 4) + + +@pytest.mark.parametrize( + ("overrides", "match"), + [ + ({"lm_head_weight": None}, "lm_head_weight"), + ({"hidden": torch.randn(2, 4)}, "hidden rows"), + ({"lm_head_weight": torch.randn(5, 6)}, "hidden width"), + ({"temperature": torch.tensor([1.0, 0.0, 1.0])}, "temperature must be positive"), + ], +) +def test_structural_request_rejects_incomplete_or_misaligned_inputs(overrides, match): + with pytest.raises(ValueError, match=match): + _structural_request(**overrides) + + +def test_provider_may_return_a_structural_result_from_an_external_package(monkeypatch): + request = _request() + + def provider(actual_request): + return SimpleNamespace( + selected_logprobs=actual_request.logits[:, :1], + entropy=None, + backend_id="external.structural", + contract_id="external.structural.v1", + provenance={"tp_reduction": "provider_owned"}, + ) + + path = _install_provider(monkeypatch, provider) + actual, entropy = compute_selected_logprobs( + args=SimpleNamespace(selected_logprob_provider=path, selected_logprob_provider_mode="strict"), + request=request, + native=_native, + ) + + assert entropy is None + torch.testing.assert_close(actual, request.logits[:, :1]) + + def test_auto_mode_only_falls_back_for_explicit_unavailability(monkeypatch): request = _request() calls = {"native": 0} diff --git a/tests/test_train_dump_utils.py b/tests/test_train_dump_utils.py new file mode 100644 index 000000000..a7cbf60fa --- /dev/null +++ b/tests/test_train_dump_utils.py @@ -0,0 +1,23 @@ +from types import SimpleNamespace + +import torch + +from vime.utils.train_dump_utils import save_debug_train_data + + +def test_debug_dump_adds_rank_and_replaces_atomically(tmp_path, monkeypatch): + monkeypatch.setattr(torch.distributed, "get_rank", lambda: 2) + args = SimpleNamespace(save_debug_train_data=str(tmp_path / "{rollout_id}.pt")) + + save_debug_train_data( + args, + rollout_id=5, + rollout_data={"log_probs": [torch.tensor([1.0])]}, + ) + + output = tmp_path / "5.rank2.pt" + payload = torch.load(output, weights_only=True) + assert payload["rollout_id"] == 5 + assert payload["rank"] == 2 + torch.testing.assert_close(payload["rollout_data"]["log_probs"][0], torch.tensor([1.0])) + assert list(tmp_path.glob(".*.tmp")) == [] diff --git a/vime/backends/megatron_utils/actor.py b/vime/backends/megatron_utils/actor.py index 8a757dfe0..d752c0654 100644 --- a/vime/backends/megatron_utils/actor.py +++ b/vime/backends/megatron_utils/actor.py @@ -36,7 +36,13 @@ from .data import DataIterator, get_data_iterator, log_perf_data, log_rollout_data from .hf_checkpoint_saver import save_hf_model_to_path from .initialize import init, is_megatron_main_rank -from .loss import compute_advantages_and_returns, get_log_probs_and_entropy, get_values +from .loss import ( + compute_advantages_and_returns, + drain_captured_log_probs, + enable_log_prob_capture, + get_log_probs_and_entropy, + get_values, +) from .model import forward_only, initialize_model_and_optimizer, save, train from .update_weight.common import named_params_and_buffers from .update_weight.update_weight_from_disk import UpdateWeightFromDisk @@ -504,20 +510,37 @@ def train_actor(self, rollout_id: int, rollout_data: RolloutBatch, external_data # Train if self.args.use_routing_replay: os.environ["ROUTING_REPLAY_STAGE"] = "replay_backward" - with timer("actor_train"): - train( - rollout_id, - self.model, - self.optimizer, - self.opt_param_scheduler, - data_iterator, - num_microbatches, - global_batch_sizes, - ) + capture_log_probs = self.args.save_debug_train_data is not None and "log_probs" not in rollout_data + if capture_log_probs: + enable_log_prob_capture() + try: + with timer("actor_train"): + train( + rollout_id, + self.model, + self.optimizer, + self.opt_param_scheduler, + data_iterator, + num_microbatches, + global_batch_sizes, + ) + finally: + if capture_log_probs: + captured = drain_captured_log_probs() + partition = rollout_data.get("partition") + if captured and partition is not None and all(int(pos) in captured for pos in partition): + rollout_data["log_probs"] = [captured[int(pos)] for pos in partition] + if self.args.save_debug_train_data is not None: + train_dump_utils.save_debug_train_data( + self.args, + rollout_id=rollout_id, + rollout_data=rollout_data, + ) self.prof.step(rollout_id=rollout_id) - train_dump_utils.save_debug_train_data(self.args, rollout_id=rollout_id, rollout_data=rollout_data) + if self.args.save_debug_train_data is None: + train_dump_utils.save_debug_train_data(self.args, rollout_id=rollout_id, rollout_data=rollout_data) if self.args.use_routing_replay: RoutingReplay.clear_all() diff --git a/vime/backends/megatron_utils/arguments.py b/vime/backends/megatron_utils/arguments.py index b7bdc1861..e6dcb1f70 100644 --- a/vime/backends/megatron_utils/arguments.py +++ b/vime/backends/megatron_utils/arguments.py @@ -69,10 +69,35 @@ def _is_moe_config(hf_config): ) +def _validate_selected_logprob_provider_args(args): + provider_path = str(getattr(args, "selected_logprob_provider", "") or "").strip() + provider_mode = str(getattr(args, "selected_logprob_provider_mode", "auto") or "auto").strip().lower() + if provider_mode == "strict" and not provider_path: + raise ValueError( + "--selected-logprob-provider-mode strict requires --selected-logprob-provider; " + "otherwise Vime has no provider to enforce" + ) + if provider_mode not in {"auto", "strict"}: + raise ValueError( + "selected_logprob_provider_mode must be 'auto' or 'strict', " + f"got {provider_mode!r}" + ) + if ( + provider_mode == "strict" + and provider_path == "rl_engine.integrations.vime.logp.provider" + and float(getattr(args, "rollout_top_p", 1.0)) != 1.0 + ): + raise ValueError( + "the RL-Kernel selected-logprob provider currently requires --rollout-top-p 1.0; " + "top-p replay masks are not part of its validated contract" + ) + + def validate_args(args): """Run megatron's own validate_args plus vime-specific megatron validations.""" _megatron_validate_args(args) + _validate_selected_logprob_provider_args(args) # always use varlen args.variable_seq_lengths = True diff --git a/vime/backends/megatron_utils/loss.py b/vime/backends/megatron_utils/loss.py index de8a58ad2..623f2d3c0 100644 --- a/vime/backends/megatron_utils/loss.py +++ b/vime/backends/megatron_utils/loss.py @@ -1,6 +1,6 @@ import warnings from argparse import Namespace -from collections.abc import Callable, Iterator +from collections.abc import Callable, Iterator, Mapping from typing import Any import torch @@ -31,11 +31,7 @@ get_sum_of_sample_mean, slice_log_prob_with_cp, ) -from .selected_logprob_provider import ( - ContextParallelLayout, - SelectedLogprobRequest, - compute_selected_logprobs, -) +from .selected_logprob_provider import ContextParallelLayout, SelectedLogprobRequest, compute_selected_logprobs ROLLOUT_TOP_P_TOKEN_KEYS = ( "rollout_top_p_token_ids", @@ -43,6 +39,31 @@ ) +_LOG_PROB_CAPTURE: dict[int, torch.Tensor] | None = None + + +def enable_log_prob_capture() -> None: + global _LOG_PROB_CAPTURE + _LOG_PROB_CAPTURE = {} + + +def drain_captured_log_probs() -> dict[int, torch.Tensor]: + global _LOG_PROB_CAPTURE + captured = _LOG_PROB_CAPTURE or {} + _LOG_PROB_CAPTURE = None + return captured + + +def _maybe_capture_log_probs(batch: RolloutBatch, log_probs: list[torch.Tensor]) -> None: + if _LOG_PROB_CAPTURE is None: + return + positions = batch.get("partition") + if not positions: + return + for position, log_prob in zip(positions, log_probs, strict=True): + _LOG_PROB_CAPTURE.setdefault(int(position), log_prob.detach().clone()) + + def get_rollout_top_p_logprob_kwargs(args: Namespace, batch: dict[str, Any]) -> dict[str, Any]: if args.rollout_top_p == 1.0: return {} @@ -490,6 +511,7 @@ def get_log_probs_and_entropy( non_loss_data: bool = True, top_p_token_ids: list[list[int]] | None = None, top_p_token_offsets: list[list[int]] | None = None, + linear_logp_context: Mapping[str, Any] | None = None, ) -> dict[str, list[torch.Tensor]]: """Compute per-token log-probabilities (and optionally entropy) on responses. @@ -501,7 +523,8 @@ def get_log_probs_and_entropy( log-probabilities; entropy is always computed from the unmasked logits. """ assert non_loss_data - assert logits.dtype == torch.float32, f"{logits.dtype}" + if logits.dtype not in (torch.float32, torch.float16, torch.bfloat16): + raise TypeError(f"selected-logprob logits must use fp32, fp16, or bf16; got {logits.dtype}") assert len(logits.shape) == 3, f"{logits.shape}" assert logits.size(0) == 1, f"{logits.shape}" logits = logits.squeeze(0) @@ -551,7 +574,28 @@ def get_log_probs_and_entropy( with_entropy_grad=with_entropy_grad, chunk_size=chunk_size, log_prob_keep_mask=top_p_keep_mask, - metadata={"logits_are_temperature_scaled": True}, + hidden=None if linear_logp_context is None else linear_logp_context.get("hidden"), + lm_head_weight=None if linear_logp_context is None else linear_logp_context.get("lm_head_weight"), + lm_head_bias=None if linear_logp_context is None else linear_logp_context.get("lm_head_bias"), + vocab_start_index=0 if linear_logp_context is None else int(linear_logp_context.get("vocab_start_index", 0)), + global_vocab_size=( + getattr(args, "padded_vocab_size", None) + if linear_logp_context is None + else linear_logp_context.get("global_vocab_size", getattr(args, "padded_vocab_size", None)) + ), + real_vocab_size=( + getattr(args, "vocab_size", None) + if linear_logp_context is None + else linear_logp_context.get("real_vocab_size", getattr(args, "vocab_size", None)) + ), + temperature=rollout_temperature, + metadata={ + "logits_are_temperature_scaled": True, + "real_vocab_size": getattr(args, "vocab_size", None), + "padded_vocab_size": getattr(args, "padded_vocab_size", None), + "tp_rank": mpu.get_tensor_model_parallel_rank(), + "tp_world_size": mpu.get_tensor_model_parallel_world_size(), + }, ) # The provider owns only selected-logprob math and its TP reduction. Vime # retains target construction, CP layout ownership, response extraction, @@ -910,6 +954,7 @@ def policy_loss_function( batch: RolloutBatch, logits: torch.Tensor, sum_of_sample_mean: Callable[[torch.Tensor], torch.Tensor], + linear_logp_context: Mapping[str, Any] | None = None, ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]: """Compute policy loss (PPO/GSPO) and metrics. @@ -948,10 +993,12 @@ def policy_loss_function( total_lengths=total_lengths, response_lengths=response_lengths, with_entropy=True, + linear_logp_context=linear_logp_context, **get_rollout_top_p_logprob_kwargs(args, batch), ) log_probs = log_probs_and_entropy["log_probs"] + _maybe_capture_log_probs(batch, log_probs) if not args.use_rollout_logprobs and not old_log_probs: old_log_probs = [log_prob.detach() for log_prob in log_probs] train_log_probs_for_tis = batch.get("log_probs") @@ -1098,10 +1145,17 @@ def policy_loss_function( loss += 0 * logits.sum() train_rollout_logprob_abs_diff = None + train_rollout_logprob_debug = None if "rollout_log_probs" in batch and batch["rollout_log_probs"]: rollout_log_probs = torch.cat(batch["rollout_log_probs"], dim=0) log_probs_to_compare = log_probs if args.use_rollout_logprobs else old_log_probs - train_rollout_logprob_abs_diff = sum_of_sample_mean((log_probs_to_compare - rollout_log_probs).abs()) + logprob_abs_diff = (log_probs_to_compare - rollout_log_probs).abs() + train_rollout_logprob_abs_diff = sum_of_sample_mean(logprob_abs_diff) + train_rollout_logprob_debug = { + "mismatch_count": torch.ne(log_probs_to_compare, rollout_log_probs).sum().to(torch.float32), + "max_abs_diff": logprob_abs_diff.max() if logprob_abs_diff.numel() else logprob_abs_diff.new_tensor(0.0), + "numel": logprob_abs_diff.new_tensor(float(logprob_abs_diff.numel())), + } reported_loss = { "loss": loss.clone().detach(), @@ -1113,6 +1167,14 @@ def policy_loss_function( if train_rollout_logprob_abs_diff is not None: reported_loss["train_rollout_logprob_abs_diff"] = train_rollout_logprob_abs_diff.clone().detach() + assert train_rollout_logprob_debug is not None + reported_loss["train_current_rollout_logprob_mismatch_count"] = ( + train_rollout_logprob_debug["mismatch_count"].clone().detach() + ) + reported_loss["train_current_rollout_logprob_max_abs_diff"] = ( + train_rollout_logprob_debug["max_abs_diff"].clone().detach() + ) + reported_loss["train_current_rollout_logprob_numel"] = train_rollout_logprob_debug["numel"].clone().detach() if args.use_kl_loss: reported_loss["kl_loss"] = kl_loss.clone().detach() @@ -1142,6 +1204,7 @@ def value_loss_function( batch: RolloutBatch, logits: torch.Tensor, sum_of_sample_mean: Callable[[torch.Tensor], torch.Tensor], + linear_logp_context: Mapping[str, Any] | None = None, ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]: """Compute clipped value loss and metrics. @@ -1199,6 +1262,7 @@ def sft_loss_function( batch: RolloutBatch, logits: torch.Tensor, sum_of_sample_mean: Callable[[torch.Tensor], torch.Tensor], + linear_logp_context: Mapping[str, Any] | None = None, ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]: """Compute supervised fine-tuning loss over response tokens. @@ -1226,6 +1290,7 @@ def sft_loss_function( total_lengths=total_lengths, response_lengths=response_lengths, with_entropy=False, + linear_logp_context=linear_logp_context, ) log_probs = log_probs_and_entropy["log_probs"] @@ -1250,6 +1315,7 @@ def loss_function( num_microbatches: int, step_global_batch_size: int, logits: torch.Tensor, + linear_logp_context: Mapping[str, Any] | None = None, ) -> tuple[torch.Tensor, int | torch.Tensor, dict[str, list[str] | torch.Tensor]]: """Dispatch to the configured loss and rescale for Megatron integration. @@ -1301,9 +1367,29 @@ def loss_function( raise ValueError(f"Unknown loss type: {args.loss_type}") if args.recompute_loss_function: - loss, log = checkpoint(func, args, batch, logits, sum_of_sample_mean, use_reentrant=False) + if args.loss_type in {"policy_loss", "sft_loss"}: + loss, log = checkpoint( + func, + args, + batch, + logits, + sum_of_sample_mean, + linear_logp_context=linear_logp_context, + use_reentrant=False, + ) + else: + loss, log = checkpoint(func, args, batch, logits, sum_of_sample_mean, use_reentrant=False) else: - loss, log = func(args, batch, logits, sum_of_sample_mean) + if args.loss_type in {"policy_loss", "sft_loss"}: + loss, log = func( + args, + batch, + logits, + sum_of_sample_mean, + linear_logp_context=linear_logp_context, + ) + else: + loss, log = func(args, batch, logits, sum_of_sample_mean) # With allgather-CP, some CP ranks may have no loss-contributing tokens (e.g., all # padding). Without this, gradient doesn't flow through their attention path, so diff --git a/vime/backends/megatron_utils/megatron_to_hf/qwen2.py b/vime/backends/megatron_utils/megatron_to_hf/qwen2.py index f7b72935c..0ea376a68 100644 --- a/vime/backends/megatron_utils/megatron_to_hf/qwen2.py +++ b/vime/backends/megatron_utils/megatron_to_hf/qwen2.py @@ -57,9 +57,15 @@ def convert_qwen2_to_hf(args, name, param): ] elif rest == "mlp.linear_fc2.weight": return [(f"model.layers.{layer_idx}.mlp.down_proj.weight", param)] - elif rest == "self_attention.linear_qkv.layer_norm_weight": + elif rest in ( + "input_layernorm.weight", + "self_attention.linear_qkv.layer_norm_weight", + ): return [(f"model.layers.{layer_idx}.input_layernorm.weight", param)] - elif rest == "mlp.linear_fc1.layer_norm_weight": + elif rest in ( + "pre_mlp_layernorm.weight", + "mlp.linear_fc1.layer_norm_weight", + ): return [(f"model.layers.{layer_idx}.post_attention_layernorm.weight", param)] # qk norm diff --git a/vime/backends/megatron_utils/model.py b/vime/backends/megatron_utils/model.py index 8642927ac..ff5917acb 100644 --- a/vime/backends/megatron_utils/model.py +++ b/vime/backends/megatron_utils/model.py @@ -41,6 +41,81 @@ logger = logging.getLogger(__name__) +def _unwrap_model_chunk(model_chunk): + while hasattr(model_chunk, "module"): + model_chunk = model_chunk.module + return model_chunk + + +def _install_linear_logp_capture(model_chunks: Sequence[DDP], args: Namespace) -> None: + from megatron.core.tensor_parallel.utils import VocabUtility + + for model_chunk in model_chunks: + model_module = _unwrap_model_chunk(model_chunk) + output_layer = getattr(model_module, "output_layer", None) + if output_layer is None or hasattr(output_layer, "_vime_linear_logp_capture_handle"): + continue + if not hasattr(output_layer, "output_size_per_partition") or not hasattr(output_layer, "tp_group"): + continue + + def capture(module, inputs, kwargs, *, owner=model_module, runtime_args=args): + hidden = inputs[0] if inputs else kwargs.get("input_") + weight = kwargs.get("weight") + if weight is None: + weight = getattr(module, "weight", None) + if not isinstance(hidden, torch.Tensor) or not isinstance(weight, torch.Tensor): + raise RuntimeError("Megatron final output layer did not expose hidden and LM-head weight") + if hidden.ndim == 3: + hidden_2d = hidden.transpose(0, 1).contiguous().reshape(-1, hidden.size(-1)) + elif hidden.ndim == 2: + hidden_2d = hidden + else: + raise RuntimeError(f"unsupported Megatron LM-head hidden shape: {tuple(hidden.shape)}") + tp_world = mpu.get_tensor_model_parallel_world_size() + tp_rank = mpu.get_tensor_model_parallel_rank() + global_vocab = int(getattr(runtime_args, "padded_vocab_size", weight.size(0) * tp_world)) + real_vocab = int(getattr(runtime_args, "vocab_size", global_vocab)) + if weight.size(0) * tp_world != global_vocab: + raise RuntimeError( + "strict linear_logp requires complete padded TP shards: " + f"local={weight.size(0)} * tp={tp_world} != padded={global_vocab}" + ) + vocab_start, _ = VocabUtility.vocab_range_from_per_partition_vocab_size(weight.size(0), tp_rank, tp_world) + if vocab_start != tp_rank * weight.size(0): + raise RuntimeError( + "strict linear_logp LM-head vocab offset is not rank-contiguous: " + f"got {vocab_start}, expected {tp_rank * weight.size(0)}" + ) + if real_vocab <= 0 or real_vocab > global_vocab: + raise RuntimeError( + f"invalid strict linear_logp vocab contract: real={real_vocab}, padded={global_vocab}" + ) + if hidden_2d.size(1) != weight.size(1): + raise RuntimeError("strict linear_logp hidden width does not match LM-head width") + bias = getattr(module, "bias", None) + if bias is not None and (bias.ndim != 1 or bias.size(0) != weight.size(0)): + raise RuntimeError("strict linear_logp LM-head bias does not match padded shard") + owner._vime_linear_logp_context = { + "hidden": hidden_2d, + "lm_head_weight": weight, + "lm_head_bias": bias, + "tp_group": getattr(module, "tp_group", None), + "vocab_start_index": int(vocab_start), + "global_vocab_size": global_vocab, + "real_vocab_size": real_vocab, + } + + handle = output_layer.register_forward_pre_hook(capture, with_kwargs=True) + output_layer._vime_linear_logp_capture_handle = handle + + +def _take_linear_logp_context(model_chunk): + owner = _unwrap_model_chunk(model_chunk) + context = getattr(owner, "_vime_linear_logp_context", None) + owner._vime_linear_logp_context = None + return context + + def _disable_tqdm_for_non_main_rank() -> bool: return not ( mpu.get_data_parallel_rank(with_context_parallel=True) == 0 @@ -292,6 +367,8 @@ def setup_model_and_optimizer( assert args.load is not None or args.pretrained_checkpoint is not None model = get_model(get_model_provider_func(args, role), ModelType.encoder_or_decoder) + if role == "actor": + _install_linear_logp_capture(model, args) # Optimizer kwargs = {} @@ -431,6 +508,7 @@ def forward_step( if batch["multimodal_train_inputs"] is not None: forward_kwargs.update(batch["multimodal_train_inputs"]) output_tensor = model(**forward_kwargs) + linear_logp_context = _take_linear_logp_context(model) output_kwargs = { "args": args, @@ -438,6 +516,7 @@ def forward_step( "total_lengths": total_lengths, "response_lengths": response_lengths, "with_entropy": args.use_rollout_entropy, + "linear_logp_context": linear_logp_context, } if use_rollout_top_p_replay: output_kwargs.update(get_rollout_top_p_logprob_kwargs(args, batch)) @@ -593,6 +672,7 @@ def forward_step(data_iterator: DataIterator, model: GPTModel, return_schedule_p "rollout_log_probs", "teacher_log_probs", "rollout_mask_sums", + *(["partition"] if args.save_debug_train_data is not None else []), ], ), args.data_pad_size_multiplier, @@ -648,10 +728,19 @@ def forward_step(data_iterator: DataIterator, model: GPTModel, return_schedule_p output_tensor = model(**forward_kwargs) + linear_logp_context = _take_linear_logp_context(model) + if os.environ.get("ENABLE_ROUTING_REPLAY", "0") == "1": os.environ["ROUTING_REPLAY_STAGE"] = old_stage - return output_tensor, partial(loss_function, args, batch, num_microbatches, step_global_batch_size) + return output_tensor, partial( + loss_function, + args, + batch, + num_microbatches, + step_global_batch_size, + linear_logp_context=linear_logp_context, + ) # Forward pass. forward_backward_func = get_forward_backward_func() @@ -906,7 +995,10 @@ def train( logging_utils.log(args, log_dict, step_key="train/step") if args.ci_test and "train/train_rollout_logprob_abs_diff" in log_dict: - assert log_dict["train/train_rollout_logprob_abs_diff"] <= 0.1, f"{log_dict=}" + assert ( + log_dict["train/train_rollout_logprob_abs_diff"] + <= args.ci_train_rollout_logprob_abs_diff_threshold + ), f"{log_dict=}" if args.ci_test and not args.ci_disable_kl_checker: if step_id == 0 and "train/ppo_kl" in log_dict and "train/pg_clipfrac" in log_dict: diff --git a/vime/backends/megatron_utils/selected_logprob_provider.py b/vime/backends/megatron_utils/selected_logprob_provider.py index 27d26729f..0876979cc 100644 --- a/vime/backends/megatron_utils/selected_logprob_provider.py +++ b/vime/backends/megatron_utils/selected_logprob_provider.py @@ -64,6 +64,13 @@ class SelectedLogprobRequest: chunk_size: int log_prob_keep_mask: torch.Tensor | None = None metadata: Mapping[str, Any] = field(default_factory=dict) + hidden: torch.Tensor | None = None + lm_head_weight: torch.Tensor | None = None + lm_head_bias: torch.Tensor | None = None + vocab_start_index: int = 0 + global_vocab_size: int | None = None + real_vocab_size: int | None = None + temperature: float | torch.Tensor | None = None def __post_init__(self) -> None: if self.logits.ndim != 2: @@ -79,6 +86,45 @@ def __post_init__(self) -> None: raise TypeError("selected-logprob target_ids must use an integer dtype") if self.log_prob_keep_mask is not None and self.log_prob_keep_mask.shape != self.logits.shape: raise ValueError("selected-logprob log_prob_keep_mask must match logits shape") + structural = (self.hidden, self.lm_head_weight, self.lm_head_bias) + if self.hidden is not None or self.lm_head_weight is not None: + if not isinstance(self.hidden, torch.Tensor): + raise ValueError("linear_logp request.hidden must be a torch.Tensor") + if not isinstance(self.lm_head_weight, torch.Tensor): + raise ValueError("linear_logp request.lm_head_weight must be a torch.Tensor") + if self.hidden.ndim != 2: + raise ValueError("linear_logp request.hidden must be [T, hidden]") + if self.lm_head_weight.ndim != 2: + raise ValueError("linear_logp request.lm_head_weight must be [V_local, hidden]") + if self.hidden.size(0) != self.logits.size(0): + raise ValueError("linear_logp hidden rows must match local logits rows") + if self.hidden.size(1) != self.lm_head_weight.size(1): + raise ValueError("linear_logp hidden width must match LM-head width") + if self.hidden.device != self.logits.device or self.lm_head_weight.device != self.logits.device: + raise ValueError("linear_logp tensors must share the logits device") + if self.lm_head_bias is not None: + if not isinstance(self.lm_head_bias, torch.Tensor): + raise ValueError("linear_logp request.lm_head_bias must be a torch.Tensor") + if self.lm_head_bias.ndim != 1 or self.lm_head_bias.size(0) != self.lm_head_weight.size(0): + raise ValueError("linear_logp LM-head bias must match local vocab width") + if self.lm_head_bias.device != self.logits.device: + raise ValueError("linear_logp LM-head bias must share the logits device") + for name, value in ( + ("global_vocab_size", self.global_vocab_size), + ("real_vocab_size", self.real_vocab_size), + ): + if value is not None and (isinstance(value, bool) or int(value) <= 0): + raise ValueError(f"linear_logp {name} must be positive") + if self.temperature is not None: + if isinstance(self.temperature, torch.Tensor): + if self.temperature.numel() not in (1, self.logits.size(0)): + raise ValueError("linear_logp temperature must be scalar or [T]") + if bool((self.temperature <= 0).any().item()): + raise ValueError("linear_logp temperature must be positive") + elif float(self.temperature) <= 0: + raise ValueError("linear_logp temperature must be positive") + elif any(value is not None for value in structural): + raise ValueError("linear_logp structural request fields must be supplied together") @dataclass(frozen=True) @@ -142,15 +188,18 @@ def compute_selected_logprobs( return _native(request, native) try: result = provider(request) - except SelectedLogprobProviderUnavailable as exc: + except Exception as exc: + if not _is_provider_unavailable(exc): + raise if mode == "strict": raise RuntimeError(f"selected-logprob provider {path!r} is unavailable: {exc}") from exc logger.warning("Selected-logprob provider %s is unavailable; using native path: %s", path, exc) return _native(request, native) - _validate_result(result, request, strict=mode == "strict") - _log_provider_identity(result) - return result.selected_logprobs, result.entropy + normalized = _normalize_result(result) + _validate_result(normalized, request, strict=mode == "strict") + _log_provider_identity(normalized) + return normalized.selected_logprobs, normalized.entropy def _load_provider(path: str) -> SelectedLogprobProvider: @@ -178,9 +227,58 @@ def _native( ) -def _validate_result(result: Any, request: SelectedLogprobRequest, *, strict: bool) -> None: - if not isinstance(result, SelectedLogprobResult): - raise TypeError("selected-logprob providers must return SelectedLogprobResult") +def _is_provider_unavailable(exc: Exception) -> bool: + """Keep provider packages independent from Vime's exception class.""" + + return isinstance(exc, SelectedLogprobProviderUnavailable) or bool( + getattr(exc, "selected_logprob_provider_unavailable", False) + ) + + +def _normalize_result(result: Any) -> SelectedLogprobResult: + """Accept Vime's dataclass or a structural result from an external package.""" + + if isinstance(result, SelectedLogprobResult): + return result + if isinstance(result, Mapping): + + def get(name): + return result[name] + + else: + + def get(name): + return getattr(result, name) + + try: + provenance = get("provenance") + except (AttributeError, KeyError) as exc: + raise TypeError( + "selected-logprob providers must return selected_logprobs, entropy, backend_id, " + "contract_id, and provenance" + ) from exc + try: + selected_logprobs = get("selected_logprobs") + entropy = get("entropy") + backend_id = get("backend_id") + contract_id = get("contract_id") + except (AttributeError, KeyError) as exc: + raise TypeError( + "selected-logprob providers must return selected_logprobs, entropy, backend_id, " + "contract_id, and provenance" + ) from exc + if not isinstance(provenance, Mapping): + raise TypeError("selected-logprob provider provenance must be a mapping") + return SelectedLogprobResult( + selected_logprobs=selected_logprobs, + entropy=entropy, + backend_id=backend_id, + contract_id=contract_id, + provenance=provenance, + ) + + +def _validate_result(result: SelectedLogprobResult, request: SelectedLogprobRequest, *, strict: bool) -> None: if result.selected_logprobs.shape != (request.logits.size(0), 1): raise ValueError( "selected-logprob provider returned invalid selected_logprobs shape " diff --git a/vime/backends/megatron_utils/update_weight/hf_weight_iterator_bridge.py b/vime/backends/megatron_utils/update_weight/hf_weight_iterator_bridge.py index ae30f0b6f..a6ccca5be 100644 --- a/vime/backends/megatron_utils/update_weight/hf_weight_iterator_bridge.py +++ b/vime/backends/megatron_utils/update_weight/hf_weight_iterator_bridge.py @@ -59,7 +59,12 @@ def get_hf_weight_chunks(self, megatron_local_weights, progress_desc: str = "Upd named_weights = self._bridge.export_hf_weights(self.model, cpu=False, conversion_tasks=conversion_tasks) def _streaming_quantized(): - for hf_param_name, weight, megatron_param_name in named_weights: + for item in named_weights: + if len(item) == 3: + hf_param_name, weight, megatron_param_name = item + else: + hf_param_name, weight = item + megatron_param_name = hf_param_name processed_weight = postprocess_hf_param( args=self.args, megatron_param_name=megatron_param_name, diff --git a/vime/backends/megatron_utils/update_weight/update_weight_from_distributed.py b/vime/backends/megatron_utils/update_weight/update_weight_from_distributed.py index 67da89378..edf4c1570 100644 --- a/vime/backends/megatron_utils/update_weight/update_weight_from_distributed.py +++ b/vime/backends/megatron_utils/update_weight/update_weight_from_distributed.py @@ -15,7 +15,14 @@ from ray import ObjectRef from ray.actor import ActorHandle from tqdm import tqdm -from vllm.distributed.weight_transfer.nccl_engine import NCCLTrainerSendWeightsArgs, NCCLWeightTransferEngine + +try: + from vllm.distributed.weight_transfer.nccl_engine import NCCLWeightTransferEngine + + NCCLTrainerSendWeightsArgs = None +except ImportError: # Colocated IPC runs do not need the distributed transfer backend. + NCCLTrainerSendWeightsArgs = None + NCCLWeightTransferEngine = None from vime.utils.distributed_utils import get_gloo_group @@ -499,6 +506,12 @@ def update_weights_from_distributed( The *group* is a vLLM ``PyNcclCommunicator`` from ``trainer_init`` in the Megatron trainer process. """ + if NCCLWeightTransferEngine is None: + raise RuntimeError( + "This vLLM build does not provide the NCCL weight-transfer API; " + "use colocated IPC weight transfer or install a vLLM build with that API." + ) + refs = [ engine.update_weights_from_distributed.remote( names=[name for name, _ in converted_named_tensors], @@ -513,10 +526,13 @@ def update_weights_from_distributed( (name, (param.data if hasattr(param, "data") else param).contiguous()) for name, param in converted_named_tensors ) - NCCLWeightTransferEngine.trainer_send_weights( - named_gpu_iter, - NCCLTrainerSendWeightsArgs(group=group, packed=True), - ) + if NCCLTrainerSendWeightsArgs is None: + NCCLWeightTransferEngine.trainer_send_weights(named_gpu_iter, group, packed=True) + else: + NCCLWeightTransferEngine.trainer_send_weights( + named_gpu_iter, + NCCLTrainerSendWeightsArgs(group=group, packed=True), + ) return refs diff --git a/vime/backends/vllm_utils/arguments.py b/vime/backends/vllm_utils/arguments.py index f23a8b942..d970d6584 100644 --- a/vime/backends/vllm_utils/arguments.py +++ b/vime/backends/vllm_utils/arguments.py @@ -1,7 +1,11 @@ import argparse from vllm.engine.arg_utils import AsyncEngineArgs -from vllm.utils.argparse_utils import FlexibleArgumentParser + +try: + from vllm.utils.argparse_utils import FlexibleArgumentParser +except ImportError: # vLLM < 0.10 exported this parser directly. + from vllm.utils import FlexibleArgumentParser from vllm_router.launch_router import RouterArgs from vime.utils.http_utils import _wrap_ipv6 @@ -83,6 +87,10 @@ def wrapper(*name_or_flags, **kwargs): new_flags.append(s) new_kwargs = kwargs.copy() + # ``deprecated`` is vLLM parser metadata, not an argparse keyword + # before Python 3.13. Vime owns the renamed CLI surface, so dropping + # the upstream warning metadata preserves behavior on both parsers. + new_kwargs.pop("deprecated", None) if "dest" in new_kwargs and isinstance(new_kwargs["dest"], str): if not new_kwargs["dest"].startswith("vllm_"): new_kwargs["dest"] = f"vllm_{new_kwargs['dest']}" @@ -98,10 +106,15 @@ def patched_add_argument_group(*g_args, **g_kwargs): parser.add_argument = _wrap_add_argument(old_add_argument) parser.add_argument_group = patched_add_argument_group - AsyncEngineArgs.add_cli_args(parser) - from vllm.entrypoints.openai.cli_args import FrontendArgs - - FrontendArgs.add_cli_args(parser) + try: + from vllm.entrypoints.openai.cli_args import FrontendArgs + except ImportError: + from vllm.entrypoints.openai.cli_args import make_arg_parser + + make_arg_parser(parser) + else: + AsyncEngineArgs.add_cli_args(parser) + FrontendArgs.add_cli_args(parser) parser.add_argument = old_add_argument parser.add_argument_group = old_add_argument_group diff --git a/vime/backends/vllm_utils/vllm_engine.py b/vime/backends/vllm_utils/vllm_engine.py index 10301c976..7d239457d 100644 --- a/vime/backends/vllm_utils/vllm_engine.py +++ b/vime/backends/vllm_utils/vllm_engine.py @@ -10,7 +10,11 @@ import cloudpickle import requests -from vllm.utils.system_utils import kill_process_tree + +try: + from vllm.utils.system_utils import kill_process_tree +except ImportError: # vLLM < 0.10 exported this helper directly. + from vllm.utils import kill_process_tree from vime.backends.vllm_utils.external import get_server_info from vime.ray.ray_actor import RayActor @@ -89,7 +93,11 @@ def _run_vllm_server(kwargs: dict, env: dict) -> None: from vllm.entrypoints.cli.serve import ServeSubcommand from vllm.entrypoints.openai.cli_args import make_arg_parser, validate_parsed_serve_args - from vllm.utils.argparse_utils import FlexibleArgumentParser + + try: + from vllm.utils.argparse_utils import FlexibleArgumentParser + except ImportError: # vLLM < 0.10 + from vllm.utils import FlexibleArgumentParser ns = argparse.Namespace(**kwargs) parser = make_arg_parser(FlexibleArgumentParser()) @@ -358,13 +366,23 @@ def init_weight_transfer_engine(self, payload: dict) -> dict: return self._make_request("init_weight_transfer_engine", payload) def start_weight_update(self, is_checkpoint_format: bool = False) -> dict: - return self._make_request("start_weight_update", {"is_checkpoint_format": is_checkpoint_format}) + try: + return self._make_request("start_weight_update", {"is_checkpoint_format": is_checkpoint_format}) + except requests.HTTPError as exc: + if exc.response is not None and exc.response.status_code == 404: + return {"ok": True, "compat_noop": True} + raise def start_draft_weight_update(self) -> dict: return self._make_request("start_draft_weight_update", {}) def finish_weight_update(self) -> dict: - return self._make_request("finish_weight_update", {}) + try: + return self._make_request("finish_weight_update", {}) + except requests.HTTPError as exc: + if exc.response is not None and exc.response.status_code == 404: + return {"ok": True, "compat_noop": True} + raise def pull_weights(self, target_version: int): return self._make_request( @@ -674,7 +692,19 @@ def _compute_server_args( def _vllm_server_field_names() -> frozenset[str]: """Return the vLLM fields accepted by CLI generation and config overrides.""" from vllm.engine.arg_utils import AsyncEngineArgs - from vllm.entrypoints.openai.cli_args import FrontendArgs + + try: + from vllm.entrypoints.openai.cli_args import FrontendArgs + except ImportError: + from vllm.entrypoints.openai.cli_args import make_arg_parser + + try: + from vllm.utils.argparse_utils import FlexibleArgumentParser + except ImportError: # vLLM < 0.10 + from vllm.utils import FlexibleArgumentParser + + parser = make_arg_parser(FlexibleArgumentParser(add_help=False)) + return frozenset(action.dest for action in parser._actions if action.dest != argparse.SUPPRESS) return frozenset(f.name for f in (*dataclasses.fields(AsyncEngineArgs), *dataclasses.fields(FrontendArgs))) diff --git a/vime/utils/arguments.py b/vime/utils/arguments.py index 59808deed..3bb795d64 100644 --- a/vime/utils/arguments.py +++ b/vime/utils/arguments.py @@ -1516,6 +1516,11 @@ def add_ci_arguments(parser): "--ci-disable-kl-checker", action="store_true", ) + parser.add_argument( + "--ci-train-rollout-logprob-abs-diff-threshold", + type=float, + default=0.1, + ) parser.add_argument( "--ci-save-grad-norm", type=str, diff --git a/vime/utils/megatron_bridge_utils.py b/vime/utils/megatron_bridge_utils.py index c87fb5b7b..6c745609a 100644 --- a/vime/utils/megatron_bridge_utils.py +++ b/vime/utils/megatron_bridge_utils.py @@ -5,6 +5,15 @@ except ImportError: unwrap_model = None +# Megatron Bridge newer releases do not classify the vocab output wrapper used +# by Vime; register its column-parallel semantics for HF export. +try: + from megatron.bridge.models.conversion.param_mapping import AutoMapping + + AutoMapping.register_module_type("LinearCrossEntropyModule", "column") +except (ImportError, AttributeError): + pass + def patch_hf_config_for_megatron_bridge(hf_config): configs = [] diff --git a/vime/utils/train_dump_utils.py b/vime/utils/train_dump_utils.py index 4b5cbc6a5..c9f6e4ba2 100644 --- a/vime/utils/train_dump_utils.py +++ b/vime/utils/train_dump_utils.py @@ -1,4 +1,5 @@ import logging +import os from pathlib import Path import torch @@ -9,14 +10,19 @@ def save_debug_train_data(args, *, rollout_id, rollout_data): if (path_template := args.save_debug_train_data) is not None: rank = torch.distributed.get_rank() - path = Path(path_template.format(rollout_id=rollout_id, rank=rank)) + rendered = path_template.format(rollout_id=rollout_id, rank=rank) + path = Path(rendered) + if "{rank}" not in path_template: + path = path.with_name(f"{path.stem}.rank{rank}{path.suffix}") logger.info(f"Save debug train data to {path}") path.parent.mkdir(parents=True, exist_ok=True) + temporary = path.with_name(f".{path.name}.{os.getpid()}.tmp") torch.save( dict( rollout_id=rollout_id, rank=rank, rollout_data=rollout_data, ), - path, + temporary, ) + os.replace(temporary, path)