From ce9fabaf63ee3361da9836c7f5c5255526cdc903 Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Thu, 9 Jul 2026 11:52:50 +0800 Subject: [PATCH 01/50] feat(speco): add separate multi-GPU draft model training --- .../run_qwen3-8b_drafter_separate_training.sh | 135 +++++++ tests/unit/test_draft_feature_store.py | 62 +++ tests/unit/test_draft_train_launcher.py | 80 ++++ verl_speco/config/draft_trainer.yaml | 49 +++ verl_speco/draft_train.py | 32 ++ verl_speco/draft_train_launcher.py | 180 +++++++++ verl_speco/trainer/base_trainer.py | 53 +++ verl_speco/trainer/draft_dataset.py | 55 +++ verl_speco/trainer/draft_training_loop.py | 153 ++++++++ verl_speco/trainer/feature_store.py | 365 ++++++++++++++++++ verl_speco/trainer/speco_ray_trainer.py | 8 + verl_speco/workers/speco_worker.py | 118 ++++++ 12 files changed, 1290 insertions(+) create mode 100644 examples/run_qwen3-8b_drafter_separate_training.sh create mode 100644 tests/unit/test_draft_feature_store.py create mode 100644 tests/unit/test_draft_train_launcher.py create mode 100644 verl_speco/config/draft_trainer.yaml create mode 100644 verl_speco/draft_train.py create mode 100644 verl_speco/draft_train_launcher.py create mode 100644 verl_speco/trainer/draft_dataset.py create mode 100644 verl_speco/trainer/draft_training_loop.py create mode 100644 verl_speco/trainer/feature_store.py diff --git a/examples/run_qwen3-8b_drafter_separate_training.sh b/examples/run_qwen3-8b_drafter_separate_training.sh new file mode 100644 index 00000000..c05208be --- /dev/null +++ b/examples/run_qwen3-8b_drafter_separate_training.sh @@ -0,0 +1,135 @@ +set -x + +# GPU smoke example for separate EAGLE3 draft model training. +# +# Stage 1 collects draft-training features from a short PPO/vLLM run without +# training the drafter inside PPO. Stage 2 launches independent multi-GPU draft +# training with python -m verl_speco.draft_train_launcher, which internally +# starts torch.distributed.run. +# +# Usage: +# bash examples/run_qwen3-8b_drafter_eagle3_vllm_separate_multigpu.sh +# RUN_STAGE=collect bash examples/run_qwen3-8b_drafter_eagle3_vllm_separate_multigpu.sh +# RUN_STAGE=train bash examples/run_qwen3-8b_drafter_eagle3_vllm_separate_multigpu.sh + +project_name='verl_grpo_example_eagle3_drafter' +exp_name='qwen3_8b_eagle3_separate_drafter_vllm_gpu' + +gen_tp=2 +train_sp=1 +ppo_gpus_per_node=8 +draft_train_gpus_per_node=8 + +MODEL_PATH=/path/to/model +CKPTS_DIR=/path/to/checkpoint +TRAIN_FILE=/path/to/train_file +TEST_FILE=/path/to/test_file +DRAFTER_PATH=/path/to/vllm-compatible-eagle3-drafter +FEATURE_STORE_DIR=/path/to/speco/eagle3_features +DRAFT_CKPTS_DIR=/path/to/speco/eagle3_draft_ckpts + +RUN_STAGE=${RUN_STAGE:-both} + +if [ "${RUN_STAGE}" = "both" ] || [ "${RUN_STAGE}" = "collect" ]; then +PYTHONUNBUFFERED=1 python3 -m verl_speco.main \ + algorithm.adv_estimator=grpo \ + data.train_files=${TRAIN_FILE} \ + data.val_files=${TEST_FILE} \ + data.train_batch_size=64 \ + data.max_prompt_length=512 \ + data.max_response_length=8192 \ + data.filter_overlong_prompts=True \ + data.filter_overlong_prompts_workers=256 \ + data.truncation='error' \ + actor_rollout_ref.model.path=${MODEL_PATH} \ + actor_rollout_ref.actor.optim.lr=1e-6 \ + actor_rollout_ref.model.use_remove_padding=True \ + actor_rollout_ref.actor.ppo_mini_batch_size=16 \ + actor_rollout_ref.actor.ppo_micro_batch_size_per_gpu=16 \ + actor_rollout_ref.actor.use_kl_loss=True \ + actor_rollout_ref.actor.kl_loss_coef=0.001 \ + actor_rollout_ref.actor.kl_loss_type=low_var_kl \ + actor_rollout_ref.actor.entropy_coeff=0 \ + actor_rollout_ref.actor.calculate_entropy=False \ + actor_rollout_ref.model.enable_gradient_checkpointing=True \ + actor_rollout_ref.actor.fsdp_config.param_offload=True \ + actor_rollout_ref.actor.fsdp_config.optimizer_offload=True \ + actor_rollout_ref.rollout.log_prob_micro_batch_size_per_gpu=16 \ + actor_rollout_ref.rollout.tensor_model_parallel_size=${gen_tp} \ + actor_rollout_ref.actor.ulysses_sequence_parallel_size=${train_sp} \ + actor_rollout_ref.ref.ulysses_sequence_parallel_size=${train_sp} \ + actor_rollout_ref.ref.log_prob_use_dynamic_bsz=True \ + actor_rollout_ref.actor.use_dynamic_bsz=True \ + actor_rollout_ref.rollout.log_prob_use_dynamic_bsz=True \ + actor_rollout_ref.rollout.name=vllm \ + actor_rollout_ref.rollout.gpu_memory_utilization=0.4 \ + actor_rollout_ref.rollout.n=5 \ + actor_rollout_ref.ref.log_prob_micro_batch_size_per_gpu=16 \ + actor_rollout_ref.ref.fsdp_config.param_offload=True \ + actor_rollout_ref.rollout.drafter.enable=True \ + actor_rollout_ref.rollout.drafter.enable_drafter_training=True \ + actor_rollout_ref.rollout.drafter.model_path=${DRAFTER_PATH} \ + actor_rollout_ref.rollout.drafter.speculative_algorithm="EAGLE3" \ + actor_rollout_ref.rollout.drafter.training.mode=collect_only \ + actor_rollout_ref.rollout.drafter.training.feature_store.type=torch_shard \ + actor_rollout_ref.rollout.drafter.training.feature_store.path=${FEATURE_STORE_DIR} \ + actor_rollout_ref.rollout.drafter.training.feature_store.max_samples_per_shard=256 \ + actor_rollout_ref.rollout.drafter.training.feature_store.flush_interval_steps=1 \ + actor_rollout_ref.rollout.drafter.training.collect_hidden_states_from_sgl=True \ + actor_rollout_ref.rollout.drafter.training.collect_hidden_states_from_old_logprob=True \ + actor_rollout_ref.rollout.drafter.training.old_logprob_hidden_capture_impl=forward_hook \ + actor_rollout_ref.rollout.drafter.training.use_logits=False \ + actor_rollout_ref.rollout.drafter.rollout.spec_steps=3 \ + actor_rollout_ref.rollout.drafter.rollout.spec_topk=1 \ + actor_rollout_ref.rollout.drafter.rollout.spec_verify_tokens=4 \ + actor_rollout_ref.rollout.drafter.training.step=20 \ + actor_rollout_ref.rollout.drafter.training.collect_interval_steps=1 \ + actor_rollout_ref.rollout.drafter.training.training_interval_steps=1 \ + actor_rollout_ref.rollout.drafter.training.publish_interval_steps=0 \ + actor_rollout_ref.rollout.drafter.training.publish_async=False \ + actor_rollout_ref.rollout.drafter.training.publish_dtype=bf16 \ + actor_rollout_ref.rollout.load_format="auto" \ + actor_rollout_ref.actor.strategy=fsdp2 \ + algorithm.use_kl_in_reward=False \ + trainer.val_before_train=False \ + trainer.critic_warmup=0 \ + trainer.logger='["console"]' \ + trainer.project_name=${project_name} \ + trainer.experiment_name=${exp_name}_collect \ + trainer.n_gpus_per_node=${ppo_gpus_per_node} \ + trainer.nnodes=1 \ + trainer.default_local_dir=${CKPTS_DIR} \ + trainer.save_freq=20 \ + trainer.test_freq=5 \ + trainer.total_epochs=1 $@ +fi + +if [ "${RUN_STAGE}" = "both" ] || [ "${RUN_STAGE}" = "train" ]; then +PYTHONUNBUFFERED=1 python3 -m verl_speco.draft_train_launcher \ + speco.draft_training.num_gpus_per_node=${draft_train_gpus_per_node} \ + speco.draft_training.nnodes=1 \ + speco.draft_training.standalone=True \ + actor_rollout_ref.model.path=${MODEL_PATH} \ + actor_rollout_ref.actor.strategy=fsdp2 \ + actor_rollout_ref.actor.fsdp_config.param_offload=True \ + actor_rollout_ref.actor.fsdp_config.optimizer_offload=True \ + actor_rollout_ref.rollout.tensor_model_parallel_size=1 \ + actor_rollout_ref.rollout.drafter.enable=True \ + actor_rollout_ref.rollout.drafter.enable_drafter_training=True \ + actor_rollout_ref.rollout.drafter.model_path=${DRAFTER_PATH} \ + actor_rollout_ref.rollout.drafter.checkpoint_path=${DRAFT_CKPTS_DIR} \ + actor_rollout_ref.rollout.drafter.speculative_algorithm="EAGLE3" \ + actor_rollout_ref.rollout.drafter.training.mode=offline \ + actor_rollout_ref.rollout.drafter.training.max_steps=10 \ + actor_rollout_ref.rollout.drafter.training.save_interval_steps=5 \ + actor_rollout_ref.rollout.drafter.training.batch_size_per_gpu=2 \ + actor_rollout_ref.rollout.drafter.training.lr=1e-6 \ + actor_rollout_ref.rollout.drafter.training.lr_warmup_steps=0 \ + actor_rollout_ref.rollout.drafter.training.warmup_style=constant \ + actor_rollout_ref.rollout.drafter.training.use_logits=False \ + actor_rollout_ref.rollout.drafter.training.feature_store.type=torch_shard \ + actor_rollout_ref.rollout.drafter.training.feature_store.path=${FEATURE_STORE_DIR} \ + actor_rollout_ref.rollout.drafter.training.feature_store.shuffle=True \ + actor_rollout_ref.rollout.drafter.training.feature_store.repeat=True \ + actor_rollout_ref.rollout.drafter.training.feature_store.strict_schema=True $@ +fi diff --git a/tests/unit/test_draft_feature_store.py b/tests/unit/test_draft_feature_store.py new file mode 100644 index 00000000..576c9bc3 --- /dev/null +++ b/tests/unit/test_draft_feature_store.py @@ -0,0 +1,62 @@ +import pytest + +torch = pytest.importorskip("torch") + +from verl_speco.trainer.draft_dataset import DraftFeatureDataLoader, DraftFeatureDataLoaderConfig +from verl_speco.trainer.feature_store import DraftFeatureSample, TorchShardFeatureStore + + +def _sample(index: int = 0) -> DraftFeatureSample: + input_ids = torch.tensor([1, 2, 3, 4], dtype=torch.long) + index + loss_mask = torch.tensor([0, 1, 1, 0], dtype=torch.float32) + hidden_states = torch.randn(4, 8, dtype=torch.float32) + last_hidden_states = torch.randn(4, 4, dtype=torch.float32) + return DraftFeatureSample( + algorithm="EAGLE3", + input_ids=input_ids, + loss_mask=loss_mask, + hidden_states=hidden_states, + last_hidden_states=last_hidden_states, + metadata={ + "source": "unit", + "global_step": index, + "hidden_states_layout": "eagle3_aux_plus_last", + "sequence_length": 4, + "loss_tokens": 2, + }, + ) + + +def test_torch_shard_feature_store_roundtrip(tmp_path): + store = TorchShardFeatureStore(tmp_path, max_samples_per_shard=2) + store.write_many([_sample(0), _sample(1)]) + store.close() + + reader = TorchShardFeatureStore(tmp_path, read_only=True) + keys = list(reader.iter_keys(shuffle=False)) + assert len(keys) == 2 + loaded = reader.read(keys[0]) + assert loaded.algorithm == "EAGLE3" + assert torch.equal(loaded.input_ids, torch.tensor([1, 2, 3, 4])) + assert loaded.metadata["hidden_states_layout"] == "eagle3_aux_plus_last" + assert reader.get_metadata()["num_samples"] == 2 + + +def test_draft_feature_dataloader_slices_keys_by_rank(tmp_path): + store = TorchShardFeatureStore(tmp_path, max_samples_per_shard=4) + store.write_many([_sample(i) for i in range(4)]) + store.close() + + rank0 = DraftFeatureDataLoader( + TorchShardFeatureStore(tmp_path, read_only=True), + DraftFeatureDataLoaderConfig(batch_size=8, rank=0, world_size=2, shuffle=False, repeat=False), + ) + rank1 = DraftFeatureDataLoader( + TorchShardFeatureStore(tmp_path, read_only=True), + DraftFeatureDataLoaderConfig(batch_size=8, rank=1, world_size=2, shuffle=False, repeat=False), + ) + + rank0_ids = [int(sample.input_ids[0].item()) for batch in rank0 for sample in batch] + rank1_ids = [int(sample.input_ids[0].item()) for batch in rank1 for sample in batch] + assert rank0_ids == [1, 3] + assert rank1_ids == [2, 4] diff --git a/tests/unit/test_draft_train_launcher.py b/tests/unit/test_draft_train_launcher.py new file mode 100644 index 00000000..d701c9e5 --- /dev/null +++ b/tests/unit/test_draft_train_launcher.py @@ -0,0 +1,80 @@ +from __future__ import annotations + +import pytest + +from verl_speco.draft_train_launcher import ( + build_torch_distributed_command, + resolve_launch_config, +) + + +def test_launcher_resolves_python_friendly_gpu_count_override() -> None: + config = resolve_launch_config( + [ + "speco.draft_training.num_gpus_per_node=8", + "actor_rollout_ref.rollout.drafter.model_path=/draft", + ] + ) + + assert config.nproc_per_node == "8" + assert config.nnodes == "1" + assert config.standalone is True + + command = build_torch_distributed_command(config, ["foo=bar"], python_executable="python") + + assert command[:6] == [ + "python", + "-m", + "torch.distributed.run", + "--nnodes=1", + "--nproc_per_node=8", + "--standalone", + ] + assert command[-3:] == ["-m", "verl_speco.draft_train", "foo=bar"] + + +def test_launcher_resolves_multinode_settings() -> None: + config = resolve_launch_config( + [ + "speco.draft_training.nproc_per_node=4", + "speco.draft_training.nnodes=2", + "speco.draft_training.node_rank=1", + "speco.draft_training.master_addr=10.0.0.1", + "speco.draft_training.master_port=29511", + "speco.draft_training.standalone=false", + ] + ) + + command = build_torch_distributed_command(config, [], python_executable="python") + + assert "--nnodes=2" in command + assert "--nproc_per_node=4" in command + assert "--node_rank=1" in command + assert "--master_addr=10.0.0.1" in command + assert "--master_port=29511" in command + assert "--standalone" not in command + + +def test_launcher_uses_explicit_port_without_standalone() -> None: + config = resolve_launch_config( + [ + "speco.draft_training.num_gpus_per_node=8", + "speco.draft_training.master_port=29511", + ] + ) + + command = build_torch_distributed_command(config, [], python_executable="python") + + assert "--nproc_per_node=8" in command + assert "--master_port=29511" in command + assert "--standalone" not in command + + +def test_launcher_rejects_standalone_multinode() -> None: + with pytest.raises(ValueError, match="standalone=true requires nnodes=1"): + resolve_launch_config( + [ + "speco.draft_training.nnodes=2", + "speco.draft_training.standalone=true", + ] + ) diff --git a/verl_speco/config/draft_trainer.yaml b/verl_speco/config/draft_trainer.yaml new file mode 100644 index 00000000..2d5ee2e0 --- /dev/null +++ b/verl_speco/config/draft_trainer.yaml @@ -0,0 +1,49 @@ +# SPECO standalone draft trainer overlay. +# +# This config keeps the existing drafter backend shape while adding launcher +# and feature-store settings for independent draft model training. + +defaults: + - speco_trainer + - _self_ + +speco: + draft_training: + enable: true + launcher: python + nproc_per_node: 1 + nnodes: 1 + node_rank: null + master_addr: null + master_port: null + standalone: true + mode: feature_store + use_ray_controller: false + +actor_rollout_ref: + rollout: + drafter: + enable: true + enable_drafter_training: true + training: + mode: offline + max_steps: 1000 + save_interval_steps: 100 + save_final_checkpoint: true + eval_interval_steps: 0 + gradient_accumulation_steps: 1 + seed: 0 + nproc_per_node: ${speco.draft_training.nproc_per_node} + nnodes: ${speco.draft_training.nnodes} + node_rank: ${speco.draft_training.node_rank} + master_addr: ${speco.draft_training.master_addr} + master_port: ${speco.draft_training.master_port} + standalone: ${speco.draft_training.standalone} + feature_store: + type: torch_shard + path: null + max_samples_per_shard: 1024 + shuffle: true + repeat: true + prefetch_depth: 2 + strict_schema: true diff --git a/verl_speco/draft_train.py b/verl_speco/draft_train.py new file mode 100644 index 00000000..63838706 --- /dev/null +++ b/verl_speco/draft_train.py @@ -0,0 +1,32 @@ +"""Standalone SPECO draft model training entrypoint. + +The user-facing launcher is ``python -m verl_speco.draft_train_launcher``. It +starts this module through PyTorch distributed launch so each rank can +participate in draft model training. +""" + +from __future__ import annotations + +import logging + +import hydra + +from verl_speco.trainer.draft_training_loop import log_resolved_config, run_standalone_draft_training + + +logger = logging.getLogger(__name__) + + +@hydra.main(config_path="config", config_name="draft_trainer", version_base=None) +def main(config): + """Run standalone draft model training.""" + + logging.basicConfig(level=logging.INFO) + log_resolved_config(config) + result = run_standalone_draft_training(config) + if result.get("rank", 0) == 0: + logger.warning("Standalone SPECO draft training finished: %s", result) + + +if __name__ == "__main__": + main() diff --git a/verl_speco/draft_train_launcher.py b/verl_speco/draft_train_launcher.py new file mode 100644 index 00000000..700ebb84 --- /dev/null +++ b/verl_speco/draft_train_launcher.py @@ -0,0 +1,180 @@ +"""User-friendly launcher for standalone SPECO draft model training. + +This module lets examples keep the familiar ``python -m ...`` shape while still +starting one distributed training process per local device. It delegates to +PyTorch's distributed launcher instead of requiring users to type ``torchrun`` +directly. +""" + +from __future__ import annotations + +import argparse +import shlex +import subprocess +import sys +from dataclasses import dataclass +from typing import Iterable + + +_NPROC_KEYS = ( + "speco.draft_training.nproc_per_node", + "speco.draft_training.num_gpus_per_node", + "actor_rollout_ref.rollout.drafter.training.nproc_per_node", + "actor_rollout_ref.rollout.drafter.training.num_gpus_per_node", +) +_NNODES_KEYS = ( + "speco.draft_training.nnodes", + "speco.draft_training.num_nodes", + "actor_rollout_ref.rollout.drafter.training.nnodes", + "actor_rollout_ref.rollout.drafter.training.num_nodes", +) +_NODE_RANK_KEYS = ( + "speco.draft_training.node_rank", + "actor_rollout_ref.rollout.drafter.training.node_rank", +) +_MASTER_ADDR_KEYS = ( + "speco.draft_training.master_addr", + "actor_rollout_ref.rollout.drafter.training.master_addr", +) +_MASTER_PORT_KEYS = ( + "speco.draft_training.master_port", + "actor_rollout_ref.rollout.drafter.training.master_port", +) +_STANDALONE_KEYS = ( + "speco.draft_training.standalone", + "actor_rollout_ref.rollout.drafter.training.standalone", +) + + +@dataclass(frozen=True) +class DraftTrainLaunchConfig: + nproc_per_node: str + nnodes: str + node_rank: str | None + master_addr: str | None + master_port: str | None + standalone: bool + module: str + + +def _split_override(item: str) -> tuple[str, str] | None: + if "=" not in item or item.startswith("-"): + return None + key, value = item.split("=", 1) + return key, value + + +def _find_override(overrides: Iterable[str], keys: Iterable[str]) -> str | None: + wanted = set(keys) + for item in overrides: + parsed = _split_override(item) + if parsed is None: + continue + key, value = parsed + if key in wanted: + return value + return None + + +def _parse_bool(value: str | None, *, default: bool) -> bool: + if value is None: + return default + normalized = str(value).strip().lower() + if normalized in {"1", "true", "yes", "on", "y"}: + return True + if normalized in {"0", "false", "no", "off", "n"}: + return False + raise ValueError(f"Invalid boolean value for standalone: {value!r}") + + +def resolve_launch_config( + overrides: list[str], + *, + module: str = "verl_speco.draft_train", +) -> DraftTrainLaunchConfig: + """Resolve distributed-launch settings from Hydra-style CLI overrides.""" + + nproc = _find_override(overrides, _NPROC_KEYS) or "1" + nnodes = _find_override(overrides, _NNODES_KEYS) or "1" + node_rank = _find_override(overrides, _NODE_RANK_KEYS) + master_addr = _find_override(overrides, _MASTER_ADDR_KEYS) + master_port = _find_override(overrides, _MASTER_PORT_KEYS) + + standalone_default = nnodes == "1" and master_addr is None and master_port is None + standalone = _parse_bool(_find_override(overrides, _STANDALONE_KEYS), default=standalone_default) + if standalone and nnodes != "1": + raise ValueError("standalone=true requires nnodes=1") + + return DraftTrainLaunchConfig( + nproc_per_node=nproc, + nnodes=nnodes, + node_rank=node_rank, + master_addr=master_addr, + master_port=master_port, + standalone=standalone, + module=module, + ) + + +def build_torch_distributed_command( + config: DraftTrainLaunchConfig, + training_args: list[str], + *, + python_executable: str = sys.executable, +) -> list[str]: + """Build the command that starts the distributed draft trainer.""" + + command = [ + python_executable, + "-m", + "torch.distributed.run", + f"--nnodes={config.nnodes}", + f"--nproc_per_node={config.nproc_per_node}", + ] + if config.standalone: + command.append("--standalone") + else: + if config.node_rank is not None: + command.append(f"--node_rank={config.node_rank}") + if config.master_addr is not None: + command.append(f"--master_addr={config.master_addr}") + if config.master_port is not None: + command.append(f"--master_port={config.master_port}") + command.extend(["-m", config.module]) + command.extend(training_args) + return command + + +def _format_command(command: list[str]) -> str: + return " ".join(shlex.quote(part) for part in command) + + +def main(argv: list[str] | None = None) -> int: + parser = argparse.ArgumentParser(description="Launch standalone SPECO draft training.") + parser.add_argument("--dry-run", action="store_true", help="Print the resolved launch command and exit.") + parser.add_argument( + "--module", + default="verl_speco.draft_train", + help="Training module passed to the distributed launcher.", + ) + parser.add_argument( + "--python-executable", + default=sys.executable, + help="Python executable used to invoke torch.distributed.run.", + ) + args, training_args = parser.parse_known_args(argv) + + launch_config = resolve_launch_config(training_args, module=args.module) + command = build_torch_distributed_command( + launch_config, + training_args, + python_executable=args.python_executable, + ) + if args.dry_run: + print(_format_command(command)) + return 0 + return subprocess.run(command, check=False).returncode + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/verl_speco/trainer/base_trainer.py b/verl_speco/trainer/base_trainer.py index 887ddd9e..a3d6e13d 100644 --- a/verl_speco/trainer/base_trainer.py +++ b/verl_speco/trainer/base_trainer.py @@ -35,6 +35,7 @@ offload_fsdp_optimizer, ) from verl.utils.ulysses import get_ulysses_sequence_parallel_group, set_ulysses_sequence_parallel_group +from verl_speco.trainer.feature_store import DraftFeatureSample logger = logging.getLogger(__name__) logger.setLevel(os.getenv("VERL_LOGGING_LEVEL", "INFO")) @@ -2812,6 +2813,55 @@ def prepare_training_batch(self) -> Optional[dict[str, torch.Tensor]]: with self._ulysses_group_context(): return self._prepare_training_batch() + def add_feature_sample(self, sample: DraftFeatureSample | dict[str, Any]) -> None: + """Append a normalized standalone sample to the in-memory training buffer.""" + if isinstance(sample, DraftFeatureSample): + item = sample.to_training_item() + elif isinstance(sample, dict): + item = DraftFeatureSample.from_dict(sample, strict=False).to_training_item() + else: + raise TypeError(f"Unsupported draft feature sample type: {type(sample)!r}") + item["step"] = int(item.get("step", self.current_rl_step) or self.current_rl_step) + self.collected_data.append(item) + + def prepare_training_batch_from_samples( + self, + samples: list[DraftFeatureSample | dict[str, Any]], + *, + step: Optional[int] = None, + ) -> Optional[dict[str, torch.Tensor]]: + """Prepare a training batch directly from standalone feature samples.""" + current_step = int(self.current_rl_step if step is None else step) + previous_collected_data = self.collected_data + previous_current_step = self.current_rl_step + maxlen = max(len(samples), int(self.config.rollout.drafter.training.get("current_max_samples", 2000))) + self.collected_data = deque(maxlen=maxlen) + self.current_rl_step = current_step + try: + for sample in samples: + if isinstance(sample, DraftFeatureSample): + item = sample.to_training_item() + elif isinstance(sample, dict): + item = DraftFeatureSample.from_dict(sample, strict=False).to_training_item() + else: + raise TypeError(f"Unsupported draft feature sample type: {type(sample)!r}") + item["step"] = current_step + self.collected_data.append(item) + with self._ulysses_group_context(): + return self._prepare_training_batch() + finally: + self.collected_data = previous_collected_data + self.current_rl_step = previous_current_step + + async def training_step_from_batch(self, batch: dict[str, torch.Tensor], step: int) -> bool: + """Execute one optimizer step from a pre-built standalone batch.""" + try: + with torch.enable_grad(): + return await self._training_step_on_batch(batch, step) + except Exception as e: # noqa: BLE001 + logger.exception(f"Standalone training step {step} failed with error: {e}") + return False + async def training_step(self, step: int) -> bool: try: with torch.enable_grad(): @@ -2867,6 +2917,9 @@ async def _training_step_impl(self, step: int) -> bool: ) return False + return await self._training_step_on_batch(batch, step) + + async def _training_step_on_batch(self, batch: dict[str, torch.Tensor], step: int) -> bool: self.model.train() self.optimizer.zero_grad(set_to_none=True) diff --git a/verl_speco/trainer/draft_dataset.py b/verl_speco/trainer/draft_dataset.py new file mode 100644 index 00000000..01407f07 --- /dev/null +++ b/verl_speco/trainer/draft_dataset.py @@ -0,0 +1,55 @@ +"""Dataset helpers for standalone draft feature stores.""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Iterator + +from verl_speco.trainer.feature_store import DraftFeatureSample, DraftFeatureStore + + +@dataclass(frozen=True) +class DraftFeatureDataLoaderConfig: + batch_size: int + rank: int = 0 + world_size: int = 1 + shuffle: bool = True + seed: int = 0 + repeat: bool = True + + +class DraftFeatureDataLoader: + """Small iterable loader over a DraftFeatureStore. + + The first implementation deliberately keeps sharding simple: + ``rank_keys = keys[rank::world_size]``. That matches the design doc's P2 + phase-1 recommendation and works for torchrun DP/FSDP ranks. + """ + + def __init__(self, store: DraftFeatureStore, config: DraftFeatureDataLoaderConfig): + self.store = store + self.config = config + + def __iter__(self) -> Iterator[list[DraftFeatureSample]]: + epoch = 0 + while True: + keys = list( + self.store.iter_keys( + shuffle=bool(self.config.shuffle), + seed=int(self.config.seed) + epoch, + ) + ) + if not keys: + return + rank_keys = keys[int(self.config.rank) :: max(int(self.config.world_size), 1)] + batch: list[DraftFeatureSample] = [] + for key in rank_keys: + batch.append(self.store.read(key)) + if len(batch) >= int(self.config.batch_size): + yield batch + batch = [] + if batch: + yield batch + if not self.config.repeat: + return + epoch += 1 diff --git a/verl_speco/trainer/draft_training_loop.py b/verl_speco/trainer/draft_training_loop.py new file mode 100644 index 00000000..420523b4 --- /dev/null +++ b/verl_speco/trainer/draft_training_loop.py @@ -0,0 +1,153 @@ +"""Standalone torchrun training loop for SPECO draft models.""" + +from __future__ import annotations + +import asyncio +import logging +import os +from typing import Any + +import torch +import torch.distributed as dist +from omegaconf import OmegaConf +from verl.utils.device import get_device_name, get_torch_device + +from verl_speco.backends.dflash_trainer_backend import DFlashTrainerBackend +from verl_speco.backends.eagle3_trainer_backend import Eagle3TrainerBackend +from verl_speco.trainer.base_trainer import DrafterBaseTrainer +from verl_speco.trainer.draft_dataset import DraftFeatureDataLoader, DraftFeatureDataLoaderConfig +from verl_speco.trainer.feature_store import build_feature_store_from_config + +logger = logging.getLogger(__name__) + + +def run_standalone_draft_training(config) -> dict[str, Any]: + """Run independent draft training from a feature store.""" + return asyncio.run(_run_standalone_draft_training_async(config)) + + +async def _run_standalone_draft_training_async(config) -> dict[str, Any]: + rank, local_rank, world_size = _init_distributed() + draft_config = config.actor_rollout_ref + drafter_cfg = draft_config.rollout.drafter + training_cfg = drafter_cfg.training + feature_store_cfg = training_cfg.feature_store + if not feature_store_cfg.get("path"): + raise ValueError("actor_rollout_ref.rollout.drafter.training.feature_store.path is required") + + _configure_device(local_rank) + backend = _build_backend(draft_config) + trainer = DrafterBaseTrainer( + config=draft_config, + world_size=world_size, + rollout_dp_rank=rank, + training_device_mesh=None, + training_process_group=dist.group.WORLD if dist.is_initialized() and world_size > 1 else None, + data_parallel_process_group=None, + backend=backend, + ) + + activated = await trainer.activate_training_model() + if not activated: + raise RuntimeError(f"Failed to activate standalone drafter trainer on rank={rank}") + + store = build_feature_store_from_config(feature_store_cfg, read_only=True) + loader = DraftFeatureDataLoader( + store, + DraftFeatureDataLoaderConfig( + batch_size=int(training_cfg.get("batch_size_per_gpu", 4)), + rank=rank, + world_size=world_size, + shuffle=bool(feature_store_cfg.get("shuffle", True)), + repeat=bool(feature_store_cfg.get("repeat", True)), + seed=int(training_cfg.get("seed", 0) or 0), + ), + ) + + max_steps = int(training_cfg.get("max_steps", training_cfg.get("step", 1000)) or 0) + save_interval = int(training_cfg.get("save_interval_steps", 0) or 0) + successful_steps = 0 + attempted_batches = 0 + last_save_result: dict[str, Any] | None = None + try: + for samples in loader: + if max_steps > 0 and successful_steps >= max_steps: + break + attempted_batches += 1 + batch = trainer.prepare_training_batch_from_samples(samples, step=successful_steps) + has_batch = batch is not None + if not trainer._sync_batch_readiness(has_batch): + if rank == 0: + logger.warning("Stopping standalone drafter training: at least one rank has no batch") + break + if batch is None: + continue + ok = await trainer.training_step_from_batch(batch, successful_steps) + if not ok: + continue + successful_steps += 1 + if save_interval > 0 and successful_steps % save_interval == 0: + last_save_result = trainer.save_checkpoint(successful_steps, wait=True) + _barrier() + final_save = bool(training_cfg.get("save_final_checkpoint", True)) + if final_save and successful_steps > 0: + last_save_result = trainer.save_checkpoint(successful_steps, wait=True) + _barrier() + finally: + store.close() + await trainer.cleanup_training(clear_data=True) + if dist.is_initialized(): + dist.barrier() + dist.destroy_process_group() + + return { + "rank": rank, + "world_size": world_size, + "attempted_batches": attempted_batches, + "successful_steps": successful_steps, + "last_save": last_save_result, + } + + +def _build_backend(draft_config): + algo = str(draft_config.rollout.drafter.speculative_algorithm).upper() + if algo == "EAGLE3": + return Eagle3TrainerBackend(draft_config, draft_config.model) + if algo == "DFLASH": + return DFlashTrainerBackend(draft_config, draft_config.model) + if algo == "DSPARK": + from verl_speco.backends.dspark_trainer_backend import DSparkTrainerBackend + + return DSparkTrainerBackend(draft_config, draft_config.model) + raise ValueError(f"Unsupported drafter algorithm {algo!r}; expected EAGLE3, DFLASH or DSPARK") + + +def _init_distributed() -> tuple[int, int, int]: + rank = int(os.environ.get("RANK", "0")) + local_rank = int(os.environ.get("LOCAL_RANK", "0")) + world_size = int(os.environ.get("WORLD_SIZE", "1")) + if world_size > 1 and not dist.is_initialized(): + backend = "nccl" if torch.cuda.is_available() else "gloo" + dist.init_process_group(backend=backend) + return rank, local_rank, world_size + + +def _configure_device(local_rank: int) -> None: + device_name = get_device_name() + device_module = get_torch_device() + if device_name == "cpu": + return + set_device = getattr(device_module, "set_device", None) + if callable(set_device): + set_device(int(local_rank)) + + +def _barrier() -> None: + if dist.is_initialized(): + dist.barrier() + + +def log_resolved_config(config) -> None: + rank = int(os.environ.get("RANK", "0")) + if rank == 0: + logger.warning("Resolved SPECO standalone draft trainer config:\n%s", OmegaConf.to_yaml(config)) diff --git a/verl_speco/trainer/feature_store.py b/verl_speco/trainer/feature_store.py new file mode 100644 index 00000000..a9e82d9d --- /dev/null +++ b/verl_speco/trainer/feature_store.py @@ -0,0 +1,365 @@ +"""Feature store primitives for standalone SPECO draft training.""" + +from __future__ import annotations + +import json +import os +import random +import tempfile +import time +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any, Iterator, Protocol + +import torch + + +SCHEMA_VERSION = 1 +MANIFEST_NAME = "manifest.jsonl" +METADATA_NAME = "metadata.json" + + +@dataclass +class DraftFeatureSample: + """A normalized draft-training sample. + + The tensor fields intentionally mirror the format used by TorchSpec-style + independent draft training: ids, masks, hidden states and either + last-hidden supervision or sparse target logprobs. + """ + + input_ids: torch.Tensor + loss_mask: torch.Tensor + hidden_states: torch.Tensor | list[torch.Tensor] + algorithm: str = "EAGLE3" + schema_version: int = SCHEMA_VERSION + last_hidden_states: torch.Tensor | None = None + target: torch.Tensor | None = None + target_logprobs: torch.Tensor | None = None + position_ids: torch.Tensor | None = None + metadata: dict[str, Any] = field(default_factory=dict) + + @classmethod + def from_dict(cls, payload: dict[str, Any], *, strict: bool = True) -> "DraftFeatureSample": + sample = cls( + schema_version=int(payload.get("schema_version", SCHEMA_VERSION)), + algorithm=str(payload.get("algorithm", payload.get("metadata", {}).get("algorithm", "EAGLE3"))), + input_ids=payload["input_ids"], + loss_mask=payload["loss_mask"], + hidden_states=payload["hidden_states"], + last_hidden_states=payload.get("last_hidden_states", payload.get("target")), + target=payload.get("target"), + target_logprobs=payload.get("target_logprobs"), + position_ids=payload.get("position_ids"), + metadata=dict(payload.get("metadata") or {}), + ) + sample.validate(strict=strict) + return sample + + def validate(self, *, strict: bool = True) -> None: + if self.schema_version != SCHEMA_VERSION and strict: + raise ValueError(f"Unsupported DraftFeatureSample schema_version={self.schema_version}") + if not torch.is_tensor(self.input_ids): + raise TypeError("DraftFeatureSample.input_ids must be a torch.Tensor") + if not torch.is_tensor(self.loss_mask): + raise TypeError("DraftFeatureSample.loss_mask must be a torch.Tensor") + if not (torch.is_tensor(self.hidden_states) or isinstance(self.hidden_states, (list, tuple))): + raise TypeError("DraftFeatureSample.hidden_states must be a tensor or tensor list") + if self.input_ids.dim() > 1: + self.input_ids = self.input_ids.reshape(-1) + if self.loss_mask.dim() > 1: + self.loss_mask = self.loss_mask.reshape(-1) + if self.input_ids.size(0) != self.loss_mask.size(0) and strict: + raise ValueError( + "DraftFeatureSample input_ids/loss_mask length mismatch: " + f"{self.input_ids.size(0)} vs {self.loss_mask.size(0)}" + ) + if torch.is_tensor(self.hidden_states) and self.hidden_states.dim() == 3 and self.hidden_states.size(0) == 1: + self.hidden_states = self.hidden_states.squeeze(0) + if torch.is_tensor(self.last_hidden_states) and self.last_hidden_states.dim() == 3 and self.last_hidden_states.size(0) == 1: + self.last_hidden_states = self.last_hidden_states.squeeze(0) + if self.target_logprobs is not None and not torch.is_tensor(self.target_logprobs): + raise TypeError("DraftFeatureSample.target_logprobs must be a tensor when provided") + + def to_dict(self) -> dict[str, Any]: + self.validate(strict=False) + payload: dict[str, Any] = { + "schema_version": self.schema_version, + "algorithm": self.algorithm, + "input_ids": self.input_ids.detach().cpu().contiguous(), + "loss_mask": self.loss_mask.detach().cpu().float().contiguous(), + "hidden_states": _cpu_tensor_tree(self.hidden_states), + "metadata": dict(self.metadata), + } + if self.last_hidden_states is not None: + payload["last_hidden_states"] = self.last_hidden_states.detach().cpu().contiguous() + if self.target is not None: + payload["target"] = self.target.detach().cpu().contiguous() + if self.target_logprobs is not None: + payload["target_logprobs"] = self.target_logprobs.detach().cpu().contiguous() + if self.position_ids is not None: + payload["position_ids"] = self.position_ids.detach().cpu().long().contiguous() + return payload + + def to_training_item(self) -> dict[str, Any]: + payload = self.to_dict() + metadata = dict(payload.pop("metadata", {}) or {}) + item = { + "input_ids": payload.pop("input_ids"), + "loss_mask": payload.pop("loss_mask"), + "hidden_states": payload.pop("hidden_states"), + "step": int(metadata.get("global_step", metadata.get("step", 0)) or 0), + "global_step": metadata.get("global_step"), + "hidden_states_layout": metadata.get("hidden_states_layout"), + } + if "last_hidden_states" in payload: + item["last_hidden_states"] = payload["last_hidden_states"] + if "target" in payload and "last_hidden_states" not in item: + item["last_hidden_states"] = payload["target"] + if "target_logprobs" in payload: + item["target_logprobs"] = payload["target_logprobs"] + if "position_ids" in payload: + item["position_ids"] = payload["position_ids"] + for key, value in metadata.items(): + item.setdefault(key, value) + return item + + +class DraftFeatureStore(Protocol): + def write_many(self, samples: list[DraftFeatureSample | dict[str, Any]]) -> list[str]: ... + + def read(self, key: str) -> DraftFeatureSample: ... + + def iter_keys(self, *, shuffle: bool = False, seed: int = 0) -> Iterator[str]: ... + + def get_metadata(self) -> dict[str, Any]: ... + + def close(self) -> None: ... + + +class TorchShardFeatureStore: + """Local ``torch.save`` shard store. + + This is the first-stage storage backend for P2. It favors simple, + inspectable files over a long-running service. + """ + + def __init__( + self, + path: str | os.PathLike[str], + *, + max_samples_per_shard: int = 1024, + metadata: dict[str, Any] | None = None, + strict_schema: bool = True, + read_only: bool = False, + shard_prefix: str = "shard", + ): + if path is None: + raise ValueError("TorchShardFeatureStore requires a non-empty path") + self.path = Path(path) + self.max_samples_per_shard = max(int(max_samples_per_shard), 1) + self.strict_schema = bool(strict_schema) + self.read_only = bool(read_only) + self.shard_prefix = str(shard_prefix or "shard") + self.path.mkdir(parents=True, exist_ok=True) + self.manifest_path = self.path / MANIFEST_NAME + self.metadata_path = self.path / METADATA_NAME + self.metadata = { + "schema_version": SCHEMA_VERSION, + "format": "torch_shard", + "created_by": "verl_speco", + "created_at": time.time(), + } + if metadata: + self.metadata.update(metadata) + self._manifest = self._load_manifest() + self._pending: list[dict[str, Any]] = [] + self._next_shard_index = self._infer_next_shard_index() + if not self.read_only: + self._write_metadata() + + def write_many(self, samples: list[DraftFeatureSample | dict[str, Any]]) -> list[str]: + if self.read_only: + raise RuntimeError("Cannot write to a read-only TorchShardFeatureStore") + keys: list[str] = [] + for sample_like in samples: + sample = _coerce_sample(sample_like, strict=self.strict_schema) + self._pending.append(sample.to_dict()) + keys.append(f"pending:{len(self._pending) - 1}") + if len(self._pending) >= self.max_samples_per_shard: + self.flush() + return keys + + def flush(self) -> list[str]: + if not self._pending: + return [] + shard_name = f"{self.shard_prefix}_{self._next_shard_index:06d}.pt" + shard_path = self.path / shard_name + payload = { + "samples": self._pending, + "metadata": dict(self.metadata), + } + _atomic_torch_save(payload, shard_path) + entry = { + "path": shard_name, + "num_samples": len(self._pending), + "num_tokens": int(sum(_sample_token_count(sample) for sample in self._pending)), + "min_global_step": _min_metadata_int(self._pending, "global_step"), + "max_global_step": _max_metadata_int(self._pending, "global_step"), + } + with self.manifest_path.open("a", encoding="utf-8") as manifest_file: + manifest_file.write(json.dumps(entry, ensure_ascii=True, sort_keys=True) + "\n") + self._manifest.append(entry) + self._pending = [] + self._next_shard_index += 1 + return [f"{shard_name}:{idx}" for idx in range(entry["num_samples"])] + + def read(self, key: str) -> DraftFeatureSample: + shard_name, sample_index = _parse_key(key) + shard = self._load_shard(shard_name) + samples = shard.get("samples") or [] + sample = samples[int(sample_index)] + return DraftFeatureSample.from_dict(sample, strict=self.strict_schema) + + def iter_keys(self, *, shuffle: bool = False, seed: int = 0) -> Iterator[str]: + self.flush() + keys: list[str] = [] + for entry in self._load_manifest(): + shard_name = entry["path"] + for idx in range(int(entry.get("num_samples", 0))): + keys.append(f"{shard_name}:{idx}") + if shuffle: + random.Random(int(seed)).shuffle(keys) + yield from keys + + def get_metadata(self) -> dict[str, Any]: + metadata = dict(self.metadata) + if self.metadata_path.exists(): + try: + with self.metadata_path.open(encoding="utf-8") as metadata_file: + metadata.update(json.load(metadata_file)) + except (OSError, json.JSONDecodeError): + pass + metadata["num_shards"] = len(self._load_manifest()) + metadata["num_samples"] = sum(int(entry.get("num_samples", 0)) for entry in self._load_manifest()) + return metadata + + def close(self) -> None: + if not self.read_only: + self.flush() + + def _write_metadata(self) -> None: + tmp_path = self.metadata_path.with_suffix(".json.tmp") + with tmp_path.open("w", encoding="utf-8") as metadata_file: + json.dump(self.metadata, metadata_file, ensure_ascii=True, indent=2, sort_keys=True) + os.replace(tmp_path, self.metadata_path) + + def _load_manifest(self) -> list[dict[str, Any]]: + if not self.manifest_path.exists(): + return [] + entries: list[dict[str, Any]] = [] + with self.manifest_path.open(encoding="utf-8") as manifest_file: + for line in manifest_file: + line = line.strip() + if not line: + continue + entries.append(json.loads(line)) + return entries + + def _infer_next_shard_index(self) -> int: + max_index = -1 + for entry in self._manifest: + name = str(entry.get("path", "")) + stem = Path(name).stem + try: + max_index = max(max_index, int(stem.split("_")[-1])) + except (IndexError, ValueError): + continue + return max_index + 1 + + def _load_shard(self, shard_name: str) -> dict[str, Any]: + path = self.path / shard_name + try: + return torch.load(path, map_location="cpu", weights_only=False) + except TypeError: + return torch.load(path, map_location="cpu") + + +def build_feature_store_from_config(feature_store_cfg, *, read_only: bool = False) -> TorchShardFeatureStore: + store_type = str(feature_store_cfg.get("type", "torch_shard") or "torch_shard") + if store_type != "torch_shard": + raise NotImplementedError(f"Unsupported draft feature store type: {store_type}") + return TorchShardFeatureStore( + feature_store_cfg.get("path"), + max_samples_per_shard=int(feature_store_cfg.get("max_samples_per_shard", 1024)), + strict_schema=bool(feature_store_cfg.get("strict_schema", True)), + read_only=read_only, + ) + + +def _coerce_sample(sample_like: DraftFeatureSample | dict[str, Any], *, strict: bool) -> DraftFeatureSample: + if isinstance(sample_like, DraftFeatureSample): + sample_like.validate(strict=strict) + return sample_like + if isinstance(sample_like, dict): + return DraftFeatureSample.from_dict(sample_like, strict=strict) + raise TypeError(f"Unsupported draft feature sample type: {type(sample_like)!r}") + + +def _cpu_tensor_tree(value: Any) -> Any: + if torch.is_tensor(value): + return value.detach().cpu().contiguous() + if isinstance(value, (list, tuple)): + return [_cpu_tensor_tree(item) for item in value] + return value + + +def _sample_token_count(sample: dict[str, Any]) -> int: + loss_mask = sample.get("loss_mask") + if torch.is_tensor(loss_mask): + return int(loss_mask.detach().float().sum().item()) + input_ids = sample.get("input_ids") + if torch.is_tensor(input_ids): + return int(input_ids.numel()) + return 0 + + +def _metadata_int(sample: dict[str, Any], name: str) -> int | None: + metadata = sample.get("metadata") + if not isinstance(metadata, dict): + return None + value = metadata.get(name) + try: + return int(value) + except (TypeError, ValueError): + return None + + +def _min_metadata_int(samples: list[dict[str, Any]], name: str) -> int | None: + values = [_metadata_int(sample, name) for sample in samples] + values = [value for value in values if value is not None] + return min(values) if values else None + + +def _max_metadata_int(samples: list[dict[str, Any]], name: str) -> int | None: + values = [_metadata_int(sample, name) for sample in samples] + values = [value for value in values if value is not None] + return max(values) if values else None + + +def _parse_key(key: str) -> tuple[str, int]: + if ":" not in key: + raise ValueError(f"Invalid feature key {key!r}; expected 'shard.pt:index'") + shard_name, index = key.rsplit(":", 1) + return shard_name, int(index) + + +def _atomic_torch_save(payload: dict[str, Any], path: Path) -> None: + with tempfile.NamedTemporaryFile(prefix=path.name, suffix=".tmp", dir=path.parent, delete=False) as tmp_file: + tmp_name = tmp_file.name + try: + torch.save(payload, tmp_name) + os.replace(tmp_name, path) + finally: + if os.path.exists(tmp_name): + os.remove(tmp_name) diff --git a/verl_speco/trainer/speco_ray_trainer.py b/verl_speco/trainer/speco_ray_trainer.py index 31ebdac4..b0655cce 100644 --- a/verl_speco/trainer/speco_ray_trainer.py +++ b/verl_speco/trainer/speco_ray_trainer.py @@ -529,6 +529,8 @@ def _speco_prepare_drafter_checkpoint_for_worker_init(self): def _speco_should_save_drafter_checkpoint(self) -> bool: if not self.is_drafter_training_enabled(self.config): return False + if self._speco_drafter_training_mode() == "collect_only": + return False if self.drafter_wg is None: return False if not self._speco_drafter_checkpoint_save_config_enabled(): @@ -564,10 +566,16 @@ def _speco_should_train_drafter_this_step(self) -> bool: training_cfg = self._speco_drafter_training_config() return speco_step_matches_interval(self.global_steps, training_cfg.get("training_interval_steps", 1)) + def _speco_drafter_training_mode(self) -> str: + training_cfg = self._speco_drafter_training_config() + return str(training_cfg.get("mode", "online") or "online").strip().lower() + def _speco_has_collected_drafter_samples_this_step(self) -> bool: return int(getattr(self, "_speco_last_collected_samples", 0) or 0) > 0 def _speco_should_attempt_drafter_train_this_step(self) -> bool: + if self._speco_drafter_training_mode() == "collect_only": + return False if not self._speco_should_train_drafter_this_step(): return False if self._speco_has_collected_drafter_samples_this_step(): diff --git a/verl_speco/workers/speco_worker.py b/verl_speco/workers/speco_worker.py index 29290845..8d63db3a 100644 --- a/verl_speco/workers/speco_worker.py +++ b/verl_speco/workers/speco_worker.py @@ -26,6 +26,7 @@ from verl.single_controller.base.decorator import Dispatch, register from verl.utils.device import get_torch_device from verl.utils.distributed import initialize_global_process_group_ray, set_numa_affinity +from verl_speco.trainer.feature_store import DraftFeatureSample, TorchShardFeatureStore logger = logging.getLogger(__file__) logger.setLevel(os.getenv("VERL_LOGGING_LEVEL", "WARN")) @@ -34,6 +35,13 @@ DRAFTER_TARGET_SYNC_MESH = "drafter_target_sync" +def _config_str(value, default: str = "") -> str: + if value is None: + return default + text = str(value) + return default if text in {"", "None", "null"} else text + + def _is_ray_object_ref(value) -> bool: object_ref_type = getattr(ray, "ObjectRef", ()) return bool(object_ref_type) and isinstance(value, object_ref_type) @@ -264,6 +272,8 @@ def __init__(self, config: DictConfig, role: str = "speco", device_name: Optiona raise ValueError("SpecoWorker requires an explicit device_name from the trainer initialization path") self.device_name = str(device_name).lower() self.trainer = None + self.feature_writer = None + self.feature_writer_path = None self.last_global_step = None self.last_trained_step = None self.training_process_group = None @@ -417,8 +427,116 @@ def _store_rollout_sample( ): if not self.enable_drafter or not self.in_drafter_train_group or self.trainer is None: return + if self._drafter_training_mode() == "collect_only": + self._write_rollout_feature_sample(batch, hidden_states, target_logprobs) + return self.trainer.collect_online_data(batch, hidden_states, target_logprobs) + def _drafter_training_mode(self) -> str: + return str(self.config.rollout.drafter.training.get("mode", "online") or "online").strip().lower() + + def _get_feature_writer(self) -> Optional[TorchShardFeatureStore]: + feature_store_cfg = self.config.rollout.drafter.training.get("feature_store", None) + if feature_store_cfg is None: + return None + path = _config_str(feature_store_cfg.get("path", None)) + if not path: + return None + if self.feature_writer is not None and self.feature_writer_path == path: + return self.feature_writer + self.feature_writer = TorchShardFeatureStore( + path, + max_samples_per_shard=int(feature_store_cfg.get("max_samples_per_shard", 1024)), + strict_schema=bool(feature_store_cfg.get("strict_schema", True)), + metadata={ + "algorithm": str(self.config.rollout.drafter.speculative_algorithm).upper(), + "target_model_path": _config_str(getattr(self.config.model, "path", None)), + "drafter_model_path": _config_str(self.config.rollout.drafter.get("model_path", None)), + "source": "rl_collect_only", + }, + shard_prefix=f"rank{int(self.rank):05d}_pid{int(os.getpid())}", + ) + self.feature_writer_path = path + return self.feature_writer + + def _build_rollout_loss_mask(self, batch: dict, input_ids: torch.Tensor) -> torch.Tensor: + if torch.is_tensor(batch.get("loss_mask")): + return batch["loss_mask"].detach().cpu().float().reshape(-1) + ids = input_ids.detach().cpu().reshape(-1) + loss_mask = torch.zeros_like(ids, dtype=torch.float32) + prompts = batch.get("prompts") + responses = batch.get("responses") + if torch.is_tensor(prompts) and torch.is_tensor(responses): + prompt_len = int(prompts.reshape(-1).numel()) + response_ids = responses.detach().cpu().reshape(-1) + pad_token_id = int(getattr(self.config.model, "pad_token_id", 0) or 0) + max_response = max(0, min(int(response_ids.numel()), int(ids.numel()) - prompt_len)) + if max_response > 0: + loss_mask[prompt_len : prompt_len + max_response] = (response_ids[:max_response] != pad_token_id).float() + else: + loss_mask[:] = 1.0 + return loss_mask + + def _write_rollout_feature_sample( + self, + batch: dict, + hidden_states: torch.Tensor, + target_logprobs: Optional[torch.Tensor], + ) -> None: + writer = self._get_feature_writer() + if writer is None: + logger.warning( + "[SpecoWorker rank=%s] training.mode=collect_only but feature_store.path is empty; drop sample", + self.rank, + ) + return + input_ids = batch["input_ids"].detach().cpu().reshape(-1) + loss_mask = self._build_rollout_loss_mask(batch, input_ids) + metadata = { + "source": batch.get("hidden_target_logprobs_source", "rl_rollout"), + "global_step": batch.get("global_step", self.last_global_step), + "target_model_path": _config_str(getattr(self.config.model, "path", None)), + "drafter_model_path": _config_str(self.config.rollout.drafter.get("model_path", None)), + "hidden_states_layout": batch.get("hidden_states_layout") or ( + "dflash_aux" + if str(self.config.rollout.drafter.speculative_algorithm).upper() in {"DFLASH", "DSPARK"} + else "eagle3_aux_plus_last" + ), + "target_layer_ids": batch.get("target_layer_ids"), + "use_logits": bool(self.config.rollout.drafter.training.get("use_logits", False)), + "sequence_length": int(input_ids.numel()), + "loss_tokens": int(loss_mask.sum().item()), + } + for key in ( + "hidden_position_start", + "hidden_position_end", + "hidden_positions", + "hidden_prefix_cache_rows", + "hidden_window_start", + "hidden_window_end", + "hidden_lm_head_fingerprint", + "hidden_last_hidden_logprob_check", + "hidden_raw_topk_logprob_check", + "hidden_last_hidden_filter", + "hidden_last_hidden_select", + "target_logprobs_position_start", + "target_logprobs_position_end", + ): + if key in batch: + metadata[key] = batch[key] + sample = DraftFeatureSample( + algorithm=str(self.config.rollout.drafter.speculative_algorithm).upper(), + input_ids=input_ids, + loss_mask=loss_mask, + hidden_states=hidden_states.detach().cpu(), + target_logprobs=target_logprobs.detach().cpu() if torch.is_tensor(target_logprobs) else None, + metadata=metadata, + ) + writer.write_many([sample]) + flush_interval = int(self.config.rollout.drafter.training.get("feature_store", {}).get("flush_interval_steps", 1)) + if flush_interval <= 1: + writer.flush() + @register(dispatch_mode=make_nd_compute_dispatch_fn(mesh_name=DRAFTER_OWNER_ROUTE_MESH)) def collect_rollout_features(self, samples: list[dict]): if not samples: From 00fc5a8322d23bf02e909d48275c5971421b849d Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Thu, 9 Jul 2026 15:18:14 +0800 Subject: [PATCH 02/50] fix(speco): compose standalone draft training config --- examples/run_qwen3-8b_drafter_separate_training.sh | 1 + verl_speco/config/draft_trainer.yaml | 4 ++++ verl_speco/config/speco_trainer.yaml | 10 ++++++++++ 3 files changed, 15 insertions(+) diff --git a/examples/run_qwen3-8b_drafter_separate_training.sh b/examples/run_qwen3-8b_drafter_separate_training.sh index c05208be..dc05d5a1 100644 --- a/examples/run_qwen3-8b_drafter_separate_training.sh +++ b/examples/run_qwen3-8b_drafter_separate_training.sh @@ -1,3 +1,4 @@ +set -euo pipefail set -x # GPU smoke example for separate EAGLE3 draft model training. diff --git a/verl_speco/config/draft_trainer.yaml b/verl_speco/config/draft_trainer.yaml index 2d5ee2e0..505748a6 100644 --- a/verl_speco/config/draft_trainer.yaml +++ b/verl_speco/config/draft_trainer.yaml @@ -7,6 +7,10 @@ defaults: - speco_trainer - _self_ +hydra: + searchpath: + - pkg://verl.trainer.config + speco: draft_training: enable: true diff --git a/verl_speco/config/speco_trainer.yaml b/verl_speco/config/speco_trainer.yaml index b04355f5..ec74319d 100644 --- a/verl_speco/config/speco_trainer.yaml +++ b/verl_speco/config/speco_trainer.yaml @@ -53,6 +53,7 @@ actor_rollout_ref: speculative_config_overrides: {} training: + mode: online collect_hidden_states_from_sgl: false collect_hidden_states_from_old_logprob: false old_logprob_hidden_capture_impl: forward_hook @@ -136,3 +137,12 @@ actor_rollout_ref: current_max_samples: 2048 data_buffer_max_size: 1024 hidden_state_clip_value: 1.0e4 + feature_store: + type: torch_shard + path: null + max_samples_per_shard: 1024 + flush_interval_steps: 1 + shuffle: true + repeat: true + prefetch_depth: 2 + strict_schema: true From 33eb98e032557366273d80f09873d5c275f586ec Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Thu, 9 Jul 2026 17:19:16 +0800 Subject: [PATCH 03/50] fix(speco): break trainer worker import cycle Co-authored-by: Codex Signed-off-by: 755651978 <755651978@qq.com> --- tests/unit/test_package_import_boundaries.py | 34 ++++++++++++++++++++ verl_speco/trainer/__init__.py | 14 +++++++- verl_speco/workers/__init__.py | 14 +++++++- 3 files changed, 60 insertions(+), 2 deletions(-) create mode 100644 tests/unit/test_package_import_boundaries.py diff --git a/tests/unit/test_package_import_boundaries.py b/tests/unit/test_package_import_boundaries.py new file mode 100644 index 00000000..bf8fab6d --- /dev/null +++ b/tests/unit/test_package_import_boundaries.py @@ -0,0 +1,34 @@ +import subprocess +import sys + + +def _run_import_probe(source: str) -> None: + result = subprocess.run( + [sys.executable, "-c", source], + check=False, + capture_output=True, + text=True, + ) + assert result.returncode == 0, result.stderr + + +def test_trainer_package_does_not_eagerly_import_ray_trainer() -> None: + _run_import_probe( + "import sys; import verl_speco.trainer; " + "assert 'verl_speco.trainer.speco_ray_trainer' not in sys.modules" + ) + + +def test_workers_package_does_not_eagerly_import_speco_worker() -> None: + _run_import_probe( + "import sys; import verl_speco.workers; " + "assert 'verl_speco.workers.speco_worker' not in sys.modules" + ) + + +def test_resolving_feature_store_does_not_load_ray_trainer() -> None: + _run_import_probe( + "import importlib.util, sys; " + "assert importlib.util.find_spec('verl_speco.trainer.feature_store') is not None; " + "assert 'verl_speco.trainer.speco_ray_trainer' not in sys.modules" + ) diff --git a/verl_speco/trainer/__init__.py b/verl_speco/trainer/__init__.py index fd401d88..eac4a382 100644 --- a/verl_speco/trainer/__init__.py +++ b/verl_speco/trainer/__init__.py @@ -1,5 +1,17 @@ """SPECO trainer adapters.""" -from verl_speco.trainer.speco_ray_trainer import SpecoRayPPOTrainer +from typing import TYPE_CHECKING, Any + +if TYPE_CHECKING: + from verl_speco.trainer.speco_ray_trainer import SpecoRayPPOTrainer __all__ = ["SpecoRayPPOTrainer"] + + +def __getattr__(name: str) -> Any: + """Load trainer adapters without importing Ray workers during package init.""" + if name == "SpecoRayPPOTrainer": + from verl_speco.trainer.speco_ray_trainer import SpecoRayPPOTrainer + + return SpecoRayPPOTrainer + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/verl_speco/workers/__init__.py b/verl_speco/workers/__init__.py index 4121ef3f..78fb2e6b 100644 --- a/verl_speco/workers/__init__.py +++ b/verl_speco/workers/__init__.py @@ -1,5 +1,17 @@ """SPECO worker adapters.""" -from verl_speco.workers.speco_worker import SpecoWorker +from typing import TYPE_CHECKING, Any + +if TYPE_CHECKING: + from verl_speco.workers.speco_worker import SpecoWorker __all__ = ["SpecoWorker"] + + +def __getattr__(name: str) -> Any: + """Load worker adapters without importing trainer modules during package init.""" + if name == "SpecoWorker": + from verl_speco.workers.speco_worker import SpecoWorker + + return SpecoWorker + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") From 2940acd9b61de47773e2468f8deab08478b1dd2c Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Thu, 9 Jul 2026 19:28:08 +0800 Subject: [PATCH 04/50] fix(speco): flush collected features by training step Co-authored-by: Codex Signed-off-by: 755651978 <755651978@qq.com> --- tests/unit/test_draft_feature_store.py | 33 ++++++++++++++++++++++++++ verl_speco/trainer/feature_store.py | 9 +++++++ verl_speco/workers/speco_worker.py | 11 ++++++--- 3 files changed, 50 insertions(+), 3 deletions(-) diff --git a/tests/unit/test_draft_feature_store.py b/tests/unit/test_draft_feature_store.py index 576c9bc3..15cae41f 100644 --- a/tests/unit/test_draft_feature_store.py +++ b/tests/unit/test_draft_feature_store.py @@ -60,3 +60,36 @@ def test_draft_feature_dataloader_slices_keys_by_rank(tmp_path): rank1_ids = [int(sample.input_ids[0].item()) for batch in rank1 for sample in batch] assert rank0_ids == [1, 3] assert rank1_ids == [2, 4] + + +def test_flush_interval_zero_relies_on_shard_capacity(tmp_path): + store = TorchShardFeatureStore(tmp_path, max_samples_per_shard=4) + store.write_many([_sample(0), _sample(1)]) + + assert store.flush_on_step(global_step=1, interval_steps=0) == [] + assert list(tmp_path.glob("shard_*.pt")) == [] + + store.write_many([_sample(2), _sample(3)]) + assert len(list(tmp_path.glob("shard_*.pt"))) == 1 + assert store.get_metadata()["num_samples"] == 4 + + +def test_flush_interval_one_flushes_once_per_step(tmp_path): + store = TorchShardFeatureStore(tmp_path, max_samples_per_shard=32) + store.write_many([_sample(index) for index in range(16)]) + + keys = store.flush_on_step(global_step=1, interval_steps=1) + + assert len(keys) == 16 + assert len(list(tmp_path.glob("shard_*.pt"))) == 1 + assert store.get_metadata()["num_samples"] == 16 + + +def test_flush_interval_n_only_flushes_matching_steps(tmp_path): + store = TorchShardFeatureStore(tmp_path, max_samples_per_shard=32) + store.write_many([_sample(0)]) + + assert store.flush_on_step(global_step=1, interval_steps=2) == [] + assert list(tmp_path.glob("shard_*.pt")) == [] + assert len(store.flush_on_step(global_step=2, interval_steps=2)) == 1 + assert len(list(tmp_path.glob("shard_*.pt"))) == 1 diff --git a/verl_speco/trainer/feature_store.py b/verl_speco/trainer/feature_store.py index a9e82d9d..cabeead6 100644 --- a/verl_speco/trainer/feature_store.py +++ b/verl_speco/trainer/feature_store.py @@ -214,6 +214,15 @@ def flush(self) -> list[str]: self._next_shard_index += 1 return [f"{shard_name}:{idx}" for idx in range(entry["num_samples"])] + def flush_on_step(self, global_step: int | None, interval_steps: int) -> list[str]: + """Flush pending samples on configured training-step boundaries.""" + interval_steps = int(interval_steps) + if interval_steps <= 0 or global_step is None: + return [] + if int(global_step) % interval_steps != 0: + return [] + return self.flush() + def read(self, key: str) -> DraftFeatureSample: shard_name, sample_index = _parse_key(key) shard = self._load_shard(shard_name) diff --git a/verl_speco/workers/speco_worker.py b/verl_speco/workers/speco_worker.py index 8d63db3a..8d188cae 100644 --- a/verl_speco/workers/speco_worker.py +++ b/verl_speco/workers/speco_worker.py @@ -533,9 +533,13 @@ def _write_rollout_feature_sample( metadata=metadata, ) writer.write_many([sample]) - flush_interval = int(self.config.rollout.drafter.training.get("feature_store", {}).get("flush_interval_steps", 1)) - if flush_interval <= 1: - writer.flush() + + def _flush_rollout_features_for_step(self) -> None: + if self._drafter_training_mode() != "collect_only" or self.feature_writer is None: + return + feature_store_cfg = self.config.rollout.drafter.training.get("feature_store", {}) + flush_interval = int(feature_store_cfg.get("flush_interval_steps", 1)) + self.feature_writer.flush_on_step(self.last_global_step, flush_interval) @register(dispatch_mode=make_nd_compute_dispatch_fn(mesh_name=DRAFTER_OWNER_ROUTE_MESH)) def collect_rollout_features(self, samples: list[dict]): @@ -594,6 +598,7 @@ def collect_rollout_features(self, samples: list[dict]): hidden_states=hidden, target_logprobs=target_logprobs, ) + self._flush_rollout_features_for_step() @register(dispatch_mode=Dispatch.ONE_TO_ALL) def set_global_step(self, global_step: int): From 4f6ebeb64d9c76a91185ccadabf81412946fad70 Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Thu, 9 Jul 2026 19:43:17 +0800 Subject: [PATCH 05/50] fix(speco): separate primary Hydra configs Co-authored-by: Codex Signed-off-by: 755651978 <755651978@qq.com> --- tests/config/test_speco_config_overlay.py | 41 +++++-- verl_speco/config/draft_trainer.yaml | 3 +- verl_speco/config/speco_base.yaml | 136 +++++++++++++++++++++ verl_speco/config/speco_trainer.yaml | 142 +--------------------- 4 files changed, 174 insertions(+), 148 deletions(-) create mode 100644 verl_speco/config/speco_base.yaml diff --git a/tests/config/test_speco_config_overlay.py b/tests/config/test_speco_config_overlay.py index 57e0299f..37ff58c8 100644 --- a/tests/config/test_speco_config_overlay.py +++ b/tests/config/test_speco_config_overlay.py @@ -19,7 +19,7 @@ def test_overlay_has_expected_default_drafter_shape() -> None: - raw = OmegaConf.load(CONFIG_DIR / "speco_trainer.yaml") + raw = OmegaConf.load(CONFIG_DIR / "speco_base.yaml") drafter = raw.actor_rollout_ref.rollout.drafter assert raw.speco.verl_base.version == "0.8.0" @@ -42,12 +42,13 @@ def test_overlay_composes_with_pinned_upstream_verl(tmp_path: Path) -> None: # checked-out config directory so this contract remains CPU-light. composed_config_dir = tmp_path / "config" shutil.copytree(upstream_config, composed_config_dir) - overlay_source = (CONFIG_DIR / "speco_trainer.yaml").read_text(encoding="utf-8") - overlay_source = overlay_source.replace( - "pkg://verl.trainer.config", - composed_config_dir.resolve().as_uri(), - ) - (composed_config_dir / "speco_trainer.yaml").write_text(overlay_source, encoding="utf-8") + for config_name in ("speco_base.yaml", "speco_trainer.yaml"): + config_source = (CONFIG_DIR / config_name).read_text(encoding="utf-8") + config_source = config_source.replace( + "pkg://verl.trainer.config", + composed_config_dir.resolve().as_uri(), + ) + (composed_config_dir / config_name).write_text(config_source, encoding="utf-8") with initialize_config_dir(config_dir=str(composed_config_dir), version_base=None): config = compose(config_name="speco_trainer") @@ -56,3 +57,29 @@ def test_overlay_composes_with_pinned_upstream_verl(tmp_path: Path) -> None: assert config.actor_rollout_ref.rollout.drafter.enable is False assert "trainer" in config assert "algorithm" in config + + +def test_draft_trainer_composes_as_primary_config(tmp_path: Path) -> None: + upstream_root = os.getenv("VERL_SPECO_UPSTREAM_ROOT") + if not upstream_root: + pytest.skip("set VERL_SPECO_UPSTREAM_ROOT to check compose against pinned upstream verl") + upstream_config = Path(upstream_root) / "verl" / "trainer" / "config" + assert upstream_config.is_dir() + + composed_config_dir = tmp_path / "config" + shutil.copytree(upstream_config, composed_config_dir) + for config_name in ("speco_base.yaml", "draft_trainer.yaml"): + config_source = (CONFIG_DIR / config_name).read_text(encoding="utf-8") + config_source = config_source.replace( + "pkg://verl.trainer.config", + composed_config_dir.resolve().as_uri(), + ) + (composed_config_dir / config_name).write_text(config_source, encoding="utf-8") + + with initialize_config_dir(config_dir=str(composed_config_dir), version_base=None): + config = compose(config_name="draft_trainer") + + assert config.actor_rollout_ref.rollout.drafter.training.mode == "offline" + assert config.speco.draft_training.enable is True + assert "trainer" in config + assert "algorithm" in config diff --git a/verl_speco/config/draft_trainer.yaml b/verl_speco/config/draft_trainer.yaml index 505748a6..d4ce32dc 100644 --- a/verl_speco/config/draft_trainer.yaml +++ b/verl_speco/config/draft_trainer.yaml @@ -4,7 +4,8 @@ # and feature-store settings for independent draft model training. defaults: - - speco_trainer + - ppo_trainer + - speco_base - _self_ hydra: diff --git a/verl_speco/config/speco_base.yaml b/verl_speco/config/speco_base.yaml new file mode 100644 index 00000000..ec11eb5a --- /dev/null +++ b/verl_speco/config/speco_base.yaml @@ -0,0 +1,136 @@ +# Shared SPECO configuration without Hydra primary-config metadata. + +speco: + enable: true + + verl_base: + version: "0.8.0" + commit: "7aed6b230776f963fa09509c10d9c3a767d1102c" + source_modifications_allowed: false + + compatibility: + policy: warn + strict_env: VERL_SPECO_STRICT_VERL + require_verl_base_tag: v0.8.0 + require_verl_base_commit: 7aed6b230776f963fa09509c10d9c3a767d1102c + + trainer: + task_runner: verl_speco.integration.task_runner.SpecoTaskRunner + ray_trainer: verl_speco.trainer.speco_ray_trainer.SpecoRayPPOTrainer + +actor_rollout_ref: + rollout: + drafter: + enable: false + enable_drafter_training: false + speculative_algorithm: EAGLE3 + model_path: /path/to/drafter/model + checkpoint_path: null + world_size: 16 + + rollout: + spec_steps: 3 + spec_topk: 1 + spec_verify_tokens: 3 + cuda_graph_max_bs: null + + vllm: + draft_tensor_parallel_size: null + max_model_len: null + enforce_eager: null + speculative_config_overrides: {} + + training: + mode: online + collect_hidden_states_from_sgl: false + collect_hidden_states_from_old_logprob: false + old_logprob_hidden_capture_impl: forward_hook + use_data_buffer: false + collect_interval_steps: 5 + collection_sample_rate: 1.0 + max_collect_samples_per_step_per_replica: 16 + max_collect_tokens_per_step_per_replica: 2048 + hidden_state_window_mode: front + hidden_state_window_tokens_per_sample: 128 + hidden_state_window_min_rows: 512 + hidden_state_random_max_offset: null + hidden_state_random_seed_by_step: true + hidden_state_max_tokens_per_sample: null + max_seq_len: 8192 + step: 10 + batch_size_per_gpu: 4 + training_interval_steps: 5 + publish_interval_steps: 0 + publish_async: false + publish_dtype: null + publish_param_name_patterns: null + draft_update_use_shm: null + draft_update_weights_bucket_megabytes: null + draft_update_pause_generation: true + draft_update_flush_before: true + draft_update_flush_after: true + skip_heavy_cleanup_after_drafter_training: false + save_full_drafter_checkpoint: true + sample_last_n_steps: 20 + train_batches_per_cycle: 4 + lr: 1e-6 + lr_warmup_steps: 0 + min_lr_ratio: null + warmup_style: constant + use_logits: false + target_lm_head_row_restricted_sync: true + logits_topk: 128 + logits_sparse_min_intersection: 1 + logits_sparse_min_mass: null + logits_coverage_mask_min_ratio: 0.1 + logits_coverage_mask_require_top1: false + ttt_length: 1 + vocab_mapping_path: null + dflash_block_size: 16 + dflash_num_anchors: 512 + dflash_loss_decay_gamma: 7.0 + dflash_front_position_weight: 1.0 + dflash_front_position_count: 0 + dflash_hard_sample_ratio: 0.0 + dflash_hidden_size: null + dflash_num_target_layers: 5 + dflash_num_hidden_layers: 1 + dflash_mask_token_id: null + dflash_target_layer_ids: null + dflash_max_window: 512 + dflash_loss_mode: full_vocab + dflash_sampled_ce_negatives: 0 + dspark_block_size: 7 + dspark_num_anchors: 512 + dspark_loss_decay_gamma: 7.0 + dspark_hard_sample_ratio: 0.0 + dspark_hidden_size: null + dspark_num_target_layers: 5 + dspark_num_hidden_layers: 5 + dspark_mask_token_id: null + dspark_target_layer_ids: null + dspark_max_window: 512 + dspark_loss_mode: full_vocab + dspark_sampled_ce_negatives: 0 + dspark_markov_rank: 256 + dspark_markov_head_type: vanilla + dspark_confidence_head_alpha: 0.0 + dspark_confidence_loss_alpha: 0.0 + dspark_confidence_head_with_markov: true + dspark_ce_loss_alpha: 1.0 + dspark_l1_loss_alpha: 0.0 + dspark_debug_log: false + dspark_debug_log_first_n: 2 + dspark_debug_log_interval: 100 + current_max_samples: 2048 + data_buffer_max_size: 1024 + hidden_state_clip_value: 1.0e4 + feature_store: + type: torch_shard + path: null + max_samples_per_shard: 1024 + flush_interval_steps: 1 + shuffle: true + repeat: true + prefetch_depth: 2 + strict_schema: true diff --git a/verl_speco/config/speco_trainer.yaml b/verl_speco/config/speco_trainer.yaml index ec74319d..0e61ee6b 100644 --- a/verl_speco/config/speco_trainer.yaml +++ b/verl_speco/config/speco_trainer.yaml @@ -1,148 +1,10 @@ -# SPECO trainer overlay. -# -# This composes upstream verl v0.8.0 PPO defaults, then adds SPECO-only fields -# under `speco.*`. The rollout-facing `actor_rollout_ref.rollout.drafter.*` -# shape is kept as an overlay consumed by verl_speco adapters. +# SPECO online PPO trainer primary config. defaults: - ppo_trainer + - speco_base - _self_ hydra: searchpath: - pkg://verl.trainer.config - -speco: - enable: true - - verl_base: - version: "0.8.0" - commit: "7aed6b230776f963fa09509c10d9c3a767d1102c" - source_modifications_allowed: false - - compatibility: - policy: warn - strict_env: VERL_SPECO_STRICT_VERL - require_verl_base_tag: v0.8.0 - require_verl_base_commit: 7aed6b230776f963fa09509c10d9c3a767d1102c - - trainer: - task_runner: verl_speco.integration.task_runner.SpecoTaskRunner - ray_trainer: verl_speco.trainer.speco_ray_trainer.SpecoRayPPOTrainer - -actor_rollout_ref: - rollout: - drafter: - enable: false - enable_drafter_training: false - speculative_algorithm: EAGLE3 - model_path: /path/to/drafter/model - checkpoint_path: null - world_size: 16 - - rollout: - spec_steps: 3 - spec_topk: 1 - spec_verify_tokens: 3 - cuda_graph_max_bs: null - - vllm: - draft_tensor_parallel_size: null - max_model_len: null - enforce_eager: null - speculative_config_overrides: {} - - training: - mode: online - collect_hidden_states_from_sgl: false - collect_hidden_states_from_old_logprob: false - old_logprob_hidden_capture_impl: forward_hook - use_data_buffer: false - collect_interval_steps: 5 - collection_sample_rate: 1.0 - max_collect_samples_per_step_per_replica: 16 - max_collect_tokens_per_step_per_replica: 2048 - hidden_state_window_mode: front - hidden_state_window_tokens_per_sample: 128 - hidden_state_window_min_rows: 512 - hidden_state_random_max_offset: null - hidden_state_random_seed_by_step: true - hidden_state_max_tokens_per_sample: null - max_seq_len: 8192 - step: 10 - batch_size_per_gpu: 4 - training_interval_steps: 5 - publish_interval_steps: 0 - publish_async: false - publish_dtype: null - publish_param_name_patterns: null - draft_update_use_shm: null - draft_update_weights_bucket_megabytes: null - draft_update_pause_generation: true - draft_update_flush_before: true - draft_update_flush_after: true - skip_heavy_cleanup_after_drafter_training: false - save_full_drafter_checkpoint: true - sample_last_n_steps: 20 - train_batches_per_cycle: 4 - lr: 1e-6 - lr_warmup_steps: 0 - min_lr_ratio: null - warmup_style: constant - use_logits: false - target_lm_head_row_restricted_sync: true - logits_topk: 128 - logits_sparse_min_intersection: 1 - logits_sparse_min_mass: null - logits_coverage_mask_min_ratio: 0.1 - logits_coverage_mask_require_top1: false - ttt_length: 1 - vocab_mapping_path: null - dflash_block_size: 16 - dflash_num_anchors: 512 - dflash_loss_decay_gamma: 7.0 - dflash_front_position_weight: 1.0 - dflash_front_position_count: 0 - dflash_hard_sample_ratio: 0.0 - dflash_hidden_size: null - dflash_num_target_layers: 5 - dflash_num_hidden_layers: 1 - dflash_mask_token_id: null - dflash_target_layer_ids: null - dflash_max_window: 512 - dflash_loss_mode: full_vocab - dflash_sampled_ce_negatives: 0 - dspark_block_size: 7 - dspark_num_anchors: 512 - dspark_loss_decay_gamma: 7.0 - dspark_hard_sample_ratio: 0.0 - dspark_hidden_size: null - dspark_num_target_layers: 5 - dspark_num_hidden_layers: 5 - dspark_mask_token_id: null - dspark_target_layer_ids: null - dspark_max_window: 512 - dspark_loss_mode: full_vocab - dspark_sampled_ce_negatives: 0 - dspark_markov_rank: 256 - dspark_markov_head_type: vanilla - dspark_confidence_head_alpha: 0.0 - dspark_confidence_loss_alpha: 0.0 - dspark_confidence_head_with_markov: true - dspark_ce_loss_alpha: 1.0 - dspark_l1_loss_alpha: 0.0 - dspark_debug_log: false - dspark_debug_log_first_n: 2 - dspark_debug_log_interval: 100 - current_max_samples: 2048 - data_buffer_max_size: 1024 - hidden_state_clip_value: 1.0e4 - feature_store: - type: torch_shard - path: null - max_samples_per_shard: 1024 - flush_interval_steps: 1 - shuffle: true - repeat: true - prefetch_depth: 2 - strict_schema: true From 1b550ce30509738e21a7399ce6ef435f356aaac6 Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Fri, 10 Jul 2026 10:11:06 +0800 Subject: [PATCH 06/50] fix(speco): normalize offline draft training features --- tests/unit/test_draft_feature_store.py | 25 ++++ tests/unit/test_draft_train_launcher.py | 17 +++ verl_speco/draft_train_launcher.py | 36 ++++- verl_speco/inspect_feature_store.py | 152 ++++++++++++++++++++++ verl_speco/trainer/draft_training_loop.py | 77 ++++++++--- verl_speco/trainer/feature_store.py | 7 + 6 files changed, 292 insertions(+), 22 deletions(-) create mode 100644 verl_speco/inspect_feature_store.py diff --git a/tests/unit/test_draft_feature_store.py b/tests/unit/test_draft_feature_store.py index 15cae41f..6b19f8c1 100644 --- a/tests/unit/test_draft_feature_store.py +++ b/tests/unit/test_draft_feature_store.py @@ -42,6 +42,31 @@ def test_torch_shard_feature_store_roundtrip(tmp_path): assert reader.get_metadata()["num_samples"] == 2 +def test_feature_sample_normalizes_singleton_position_ids(): + sample = DraftFeatureSample( + input_ids=torch.tensor([1, 2, 3, 4], dtype=torch.long), + loss_mask=torch.tensor([0, 1, 1, 0], dtype=torch.float32), + hidden_states=torch.randn(4, 8, dtype=torch.float32), + position_ids=torch.tensor([[0, 1, 2, 3]], dtype=torch.long), + ) + + sample.validate(strict=True) + + assert sample.position_ids.shape == (4,) + + +def test_feature_sample_rejects_position_id_length_mismatch(): + sample = DraftFeatureSample( + input_ids=torch.tensor([1, 2, 3, 4], dtype=torch.long), + loss_mask=torch.tensor([0, 1, 1, 0], dtype=torch.float32), + hidden_states=torch.randn(4, 8, dtype=torch.float32), + position_ids=torch.tensor([[0, 1], [2, 3], [4, 5]], dtype=torch.long), + ) + + with pytest.raises(ValueError, match="input_ids/position_ids length mismatch"): + sample.validate(strict=True) + + def test_draft_feature_dataloader_slices_keys_by_rank(tmp_path): store = TorchShardFeatureStore(tmp_path, max_samples_per_shard=4) store.write_many([_sample(i) for i in range(4)]) diff --git a/tests/unit/test_draft_train_launcher.py b/tests/unit/test_draft_train_launcher.py index d701c9e5..829a06ec 100644 --- a/tests/unit/test_draft_train_launcher.py +++ b/tests/unit/test_draft_train_launcher.py @@ -4,6 +4,7 @@ from verl_speco.draft_train_launcher import ( build_torch_distributed_command, + normalize_training_args, resolve_launch_config, ) @@ -33,6 +34,22 @@ def test_launcher_resolves_python_friendly_gpu_count_override() -> None: assert command[-3:] == ["-m", "verl_speco.draft_train", "foo=bar"] +def test_launcher_normalizes_gpu_count_alias_for_hydra() -> None: + overrides = [ + "speco.draft_training.num_gpus_per_node=8", + "actor_rollout_ref.rollout.drafter.training.max_steps=10", + ] + config = resolve_launch_config(overrides) + + normalized = normalize_training_args(overrides, config) + + assert "speco.draft_training.num_gpus_per_node=8" not in normalized + assert "speco.draft_training.nproc_per_node=8" in normalized + assert "speco.draft_training.nnodes=1" in normalized + assert "speco.draft_training.standalone=true" in normalized + assert "actor_rollout_ref.rollout.drafter.training.max_steps=10" in normalized + + def test_launcher_resolves_multinode_settings() -> None: config = resolve_launch_config( [ diff --git a/verl_speco/draft_train_launcher.py b/verl_speco/draft_train_launcher.py index 700ebb84..08c1276e 100644 --- a/verl_speco/draft_train_launcher.py +++ b/verl_speco/draft_train_launcher.py @@ -45,6 +45,15 @@ "actor_rollout_ref.rollout.drafter.training.standalone", ) +_LAUNCH_OVERRIDE_KEYS = frozenset( + _NPROC_KEYS + + _NNODES_KEYS + + _NODE_RANK_KEYS + + _MASTER_ADDR_KEYS + + _MASTER_PORT_KEYS + + _STANDALONE_KEYS +) + @dataclass(frozen=True) class DraftTrainLaunchConfig: @@ -116,6 +125,30 @@ def resolve_launch_config( ) +def normalize_training_args(overrides: list[str], config: DraftTrainLaunchConfig) -> list[str]: + """Replace launcher aliases with canonical Hydra configuration fields.""" + + normalized = [] + for item in overrides: + parsed = _split_override(item) + if parsed is None or parsed[0] not in _LAUNCH_OVERRIDE_KEYS: + normalized.append(item) + normalized.extend( + [ + f"speco.draft_training.nproc_per_node={config.nproc_per_node}", + f"speco.draft_training.nnodes={config.nnodes}", + f"speco.draft_training.standalone={str(config.standalone).lower()}", + ] + ) + if config.node_rank is not None: + normalized.append(f"speco.draft_training.node_rank={config.node_rank}") + if config.master_addr is not None: + normalized.append(f"speco.draft_training.master_addr={config.master_addr}") + if config.master_port is not None: + normalized.append(f"speco.draft_training.master_port={config.master_port}") + return normalized + + def build_torch_distributed_command( config: DraftTrainLaunchConfig, training_args: list[str], @@ -165,9 +198,10 @@ def main(argv: list[str] | None = None) -> int: args, training_args = parser.parse_known_args(argv) launch_config = resolve_launch_config(training_args, module=args.module) + normalized_training_args = normalize_training_args(training_args, launch_config) command = build_torch_distributed_command( launch_config, - training_args, + normalized_training_args, python_executable=args.python_executable, ) if args.dry_run: diff --git a/verl_speco/inspect_feature_store.py b/verl_speco/inspect_feature_store.py new file mode 100644 index 00000000..943fb014 --- /dev/null +++ b/verl_speco/inspect_feature_store.py @@ -0,0 +1,152 @@ +"""Inspect SPECO standalone draft feature stores.""" + +from __future__ import annotations + +import argparse +import json +from collections import Counter +from pathlib import Path +from typing import Any + +import torch + +from verl_speco.trainer.feature_store import MANIFEST_NAME + + +def main() -> int: + parser = argparse.ArgumentParser(description="Inspect a torch_shard draft feature store.") + parser.add_argument("path", help="Feature store directory containing manifest.jsonl and shard .pt files.") + parser.add_argument("--max-samples", type=int, default=32, help="Maximum number of samples to inspect.") + parser.add_argument("--show-ok", action="store_true", help="Print valid samples as well as invalid samples.") + parser.add_argument("--strict-exit", action="store_true", help="Exit with code 1 when invalid samples are found.") + args = parser.parse_args() + + root = Path(args.path) + manifest_path = root / MANIFEST_NAME + if not manifest_path.exists(): + raise FileNotFoundError(f"Missing feature store manifest: {manifest_path}") + + entries = _load_manifest(manifest_path) + shape_counts: Counter[str] = Counter() + inspected = 0 + invalid = 0 + + print(f"feature_store={root}") + print(f"manifest_entries={len(entries)}") + + for entry in entries: + if inspected >= args.max_samples: + break + shard_name = str(entry.get("path")) + shard_path = root / shard_name + shard = _torch_load(shard_path) + samples = shard.get("samples") or [] + for sample_index, sample in enumerate(samples): + if inspected >= args.max_samples: + break + inspected += 1 + key = f"{shard_name}:{sample_index}" + issues = _sample_issues(sample) + summary = _sample_summary(sample) + shape_counts.update(summary.values()) + if issues: + invalid += 1 + print(f"[BAD] {key} {summary}") + for issue in issues: + print(f" - {issue}") + elif args.show_ok: + print(f"[OK] {key} {summary}") + + print(f"inspected_samples={inspected}") + print(f"invalid_samples={invalid}") + if shape_counts: + print("shape_counts:") + for shape, count in shape_counts.most_common(): + print(f" {shape}: {count}") + return 1 if invalid and args.strict_exit else 0 + + +def _load_manifest(path: Path) -> list[dict[str, Any]]: + entries: list[dict[str, Any]] = [] + with path.open(encoding="utf-8") as manifest_file: + for line in manifest_file: + line = line.strip() + if line: + entries.append(json.loads(line)) + return entries + + +def _torch_load(path: Path) -> dict[str, Any]: + try: + return torch.load(path, map_location="cpu", weights_only=False) + except TypeError: + return torch.load(path, map_location="cpu") + + +def _shape(value: Any) -> str: + if torch.is_tensor(value): + return str(tuple(value.shape)) + if isinstance(value, (list, tuple)): + return "[" + ",".join(_shape(item) for item in value) + "]" + return type(value).__name__ + + +def _tensor(value: Any, name: str, issues: list[str]) -> torch.Tensor | None: + if not torch.is_tensor(value): + issues.append(f"{name} is not a tensor: {type(value).__name__}") + return None + return value + + +def _sample_summary(sample: dict[str, Any]) -> dict[str, str]: + keys = ["input_ids", "loss_mask", "hidden_states", "last_hidden_states", "target_logprobs", "position_ids"] + return {key: _shape(sample[key]) for key in keys if key in sample} + + +def _sample_issues(sample: dict[str, Any]) -> list[str]: + issues: list[str] = [] + input_ids = _tensor(sample.get("input_ids"), "input_ids", issues) + loss_mask = _tensor(sample.get("loss_mask"), "loss_mask", issues) + hidden_states = sample.get("hidden_states") + position_ids = sample.get("position_ids") + + seq_len = None + if input_ids is not None: + if input_ids.dim() > 2 or (input_ids.dim() == 2 and 1 not in input_ids.shape): + issues.append(f"input_ids should be 1D or singleton-2D, got {_shape(input_ids)}") + seq_len = int(input_ids.numel()) + if loss_mask is not None: + if loss_mask.numel() != seq_len: + issues.append(f"loss_mask length {loss_mask.numel()} does not match input_ids length {seq_len}") + if torch.is_tensor(hidden_states): + if hidden_states.dim() == 3 and hidden_states.size(0) == 1: + hidden_len = int(hidden_states.size(1)) + elif hidden_states.dim() == 2: + hidden_len = int(hidden_states.size(0)) + else: + hidden_len = -1 + issues.append(f"hidden_states should be [seq, hidden] or [1, seq, hidden], got {_shape(hidden_states)}") + if seq_len is not None and hidden_len >= 0 and hidden_len < max(seq_len - 1, 1): + issues.append(f"hidden_states length {hidden_len} is too short for input_ids length {seq_len}") + elif isinstance(hidden_states, (list, tuple)): + for idx, tensor in enumerate(hidden_states): + if not torch.is_tensor(tensor): + issues.append(f"hidden_states[{idx}] is not a tensor: {type(tensor).__name__}") + elif tensor.dim() not in {2, 3}: + issues.append(f"hidden_states[{idx}] has unexpected shape {_shape(tensor)}") + else: + issues.append(f"hidden_states is not a tensor/list: {type(hidden_states).__name__}") + + if position_ids is not None: + if not torch.is_tensor(position_ids): + issues.append(f"position_ids is not a tensor: {type(position_ids).__name__}") + elif position_ids.numel() != seq_len: + issues.append(f"position_ids length {position_ids.numel()} does not match input_ids length {seq_len}") + elif position_ids.dim() > 1: + issues.append(f"position_ids is normalizable but not stored as 1D: {_shape(position_ids)}") + + return issues + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/verl_speco/trainer/draft_training_loop.py b/verl_speco/trainer/draft_training_loop.py index 420523b4..7334f45e 100644 --- a/verl_speco/trainer/draft_training_loop.py +++ b/verl_speco/trainer/draft_training_loop.py @@ -9,7 +9,8 @@ import torch import torch.distributed as dist -from omegaconf import OmegaConf +from torch.distributed.device_mesh import DeviceMesh +from omegaconf import OmegaConf, open_dict from verl.utils.device import get_device_name, get_torch_device from verl_speco.backends.dflash_trainer_backend import DFlashTrainerBackend @@ -34,42 +35,48 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: feature_store_cfg = training_cfg.feature_store if not feature_store_cfg.get("path"): raise ValueError("actor_rollout_ref.rollout.drafter.training.feature_store.path is required") + _disable_standalone_sequence_parallel(draft_config) _configure_device(local_rank) backend = _build_backend(draft_config) + training_device_mesh = _build_training_device_mesh(draft_config, world_size) trainer = DrafterBaseTrainer( config=draft_config, world_size=world_size, rollout_dp_rank=rank, - training_device_mesh=None, - training_process_group=dist.group.WORLD if dist.is_initialized() and world_size > 1 else None, + training_device_mesh=training_device_mesh, + training_process_group=( + None + if training_device_mesh is not None + else dist.group.WORLD if dist.is_initialized() and world_size > 1 else None + ), data_parallel_process_group=None, backend=backend, ) - activated = await trainer.activate_training_model() - if not activated: - raise RuntimeError(f"Failed to activate standalone drafter trainer on rank={rank}") - - store = build_feature_store_from_config(feature_store_cfg, read_only=True) - loader = DraftFeatureDataLoader( - store, - DraftFeatureDataLoaderConfig( - batch_size=int(training_cfg.get("batch_size_per_gpu", 4)), - rank=rank, - world_size=world_size, - shuffle=bool(feature_store_cfg.get("shuffle", True)), - repeat=bool(feature_store_cfg.get("repeat", True)), - seed=int(training_cfg.get("seed", 0) or 0), - ), - ) - max_steps = int(training_cfg.get("max_steps", training_cfg.get("step", 1000)) or 0) save_interval = int(training_cfg.get("save_interval_steps", 0) or 0) successful_steps = 0 attempted_batches = 0 last_save_result: dict[str, Any] | None = None + store = None try: + activated = await trainer.activate_training_model() + if not activated: + raise RuntimeError(f"Failed to activate standalone drafter trainer on rank={rank}") + + store = build_feature_store_from_config(feature_store_cfg, read_only=True) + loader = DraftFeatureDataLoader( + store, + DraftFeatureDataLoaderConfig( + batch_size=int(training_cfg.get("batch_size_per_gpu", 4)), + rank=rank, + world_size=world_size, + shuffle=bool(feature_store_cfg.get("shuffle", True)), + repeat=bool(feature_store_cfg.get("repeat", True)), + seed=int(training_cfg.get("seed", 0) or 0), + ), + ) for samples in loader: if max_steps > 0 and successful_steps >= max_steps: break @@ -94,7 +101,8 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: last_save_result = trainer.save_checkpoint(successful_steps, wait=True) _barrier() finally: - store.close() + if store is not None: + store.close() await trainer.cleanup_training(clear_data=True) if dist.is_initialized(): dist.barrier() @@ -122,6 +130,33 @@ def _build_backend(draft_config): raise ValueError(f"Unsupported drafter algorithm {algo!r}; expected EAGLE3, DFLASH or DSPARK") +def _disable_standalone_sequence_parallel(draft_config) -> None: + rollout_cfg = draft_config.rollout + rollout_tp_size = int(rollout_cfg.get("tensor_model_parallel_size", 1) or 1) + if rollout_tp_size <= 1: + return + logger.warning( + "Standalone draft training disables Ulysses sequence parallelism: " + "actor_rollout_ref.rollout.tensor_model_parallel_size=%s is treated as 1 for offline drafter training", + rollout_tp_size, + ) + with open_dict(rollout_cfg): + rollout_cfg.tensor_model_parallel_size = 1 + + +def _build_training_device_mesh(draft_config, world_size: int) -> DeviceMesh | None: + if world_size <= 1 or not dist.is_initialized(): + return None + strategy = str(draft_config.actor.get("strategy", "") if hasattr(draft_config, "actor") else "").lower() + if strategy != "fsdp2": + return None + return DeviceMesh( + device_type=get_device_name(), + mesh=torch.arange(world_size, dtype=torch.int64).reshape(1, world_size), + mesh_dim_names=("dp", "sp"), + ) + + def _init_distributed() -> tuple[int, int, int]: rank = int(os.environ.get("RANK", "0")) local_rank = int(os.environ.get("LOCAL_RANK", "0")) diff --git a/verl_speco/trainer/feature_store.py b/verl_speco/trainer/feature_store.py index cabeead6..695748c3 100644 --- a/verl_speco/trainer/feature_store.py +++ b/verl_speco/trainer/feature_store.py @@ -69,11 +69,18 @@ def validate(self, *, strict: bool = True) -> None: self.input_ids = self.input_ids.reshape(-1) if self.loss_mask.dim() > 1: self.loss_mask = self.loss_mask.reshape(-1) + if torch.is_tensor(self.position_ids) and self.position_ids.dim() > 1: + self.position_ids = self.position_ids.reshape(-1) if self.input_ids.size(0) != self.loss_mask.size(0) and strict: raise ValueError( "DraftFeatureSample input_ids/loss_mask length mismatch: " f"{self.input_ids.size(0)} vs {self.loss_mask.size(0)}" ) + if torch.is_tensor(self.position_ids) and self.position_ids.size(0) != self.input_ids.size(0) and strict: + raise ValueError( + "DraftFeatureSample input_ids/position_ids length mismatch: " + f"{self.input_ids.size(0)} vs {self.position_ids.size(0)}" + ) if torch.is_tensor(self.hidden_states) and self.hidden_states.dim() == 3 and self.hidden_states.size(0) == 1: self.hidden_states = self.hidden_states.squeeze(0) if torch.is_tensor(self.last_hidden_states) and self.last_hidden_states.dim() == 3 and self.last_hidden_states.size(0) == 1: From 841e32eeafefe1f1b092252f3386d11e59bef885 Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Fri, 10 Jul 2026 17:01:54 +0800 Subject: [PATCH 07/50] fix(speco): stabilize standalone draft training --- verl_speco/trainer/draft_training_loop.py | 34 ++++++++++-- verl_speco/trainer/feature_store.py | 17 ++++-- verl_speco/workers/speco_worker.py | 65 +++++++++++++++++++++-- 3 files changed, 107 insertions(+), 9 deletions(-) diff --git a/verl_speco/trainer/draft_training_loop.py b/verl_speco/trainer/draft_training_loop.py index 7334f45e..c2ba862e 100644 --- a/verl_speco/trainer/draft_training_loop.py +++ b/verl_speco/trainer/draft_training_loop.py @@ -59,6 +59,7 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: successful_steps = 0 attempted_batches = 0 last_save_result: dict[str, Any] | None = None + last_saved_step = 0 store = None try: activated = await trainer.activate_training_model() @@ -94,11 +95,12 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: continue successful_steps += 1 if save_interval > 0 and successful_steps % save_interval == 0: - last_save_result = trainer.save_checkpoint(successful_steps, wait=True) + last_save_result = _save_standalone_checkpoint(trainer, successful_steps) + last_saved_step = successful_steps _barrier() final_save = bool(training_cfg.get("save_final_checkpoint", True)) - if final_save and successful_steps > 0: - last_save_result = trainer.save_checkpoint(successful_steps, wait=True) + if final_save and successful_steps > 0 and successful_steps != last_saved_step: + last_save_result = _save_standalone_checkpoint(trainer, successful_steps) _barrier() finally: if store is not None: @@ -130,6 +132,32 @@ def _build_backend(draft_config): raise ValueError(f"Unsupported drafter algorithm {algo!r}; expected EAGLE3, DFLASH or DSPARK") +def _save_standalone_checkpoint(trainer: DrafterBaseTrainer, step: int) -> dict[str, Any]: + if not trainer.checkpoint_dir: + return {"saved": False, "reason": "missing_checkpoint_dir"} + + checkpoint_path = os.path.join(trainer.checkpoint_dir, f"draft_step_{int(step)}") + pending_full_checkpoint = getattr(trainer, "_pending_full_checkpoint_future", None) + pending_done = getattr(pending_full_checkpoint, "done", None) + if callable(pending_done) and not pending_done(): + return { + "saved": False, + "path": checkpoint_path, + "reason": "previous_save_running", + } + + future = trainer._save_checkpoint_async(int(step)) + if future is not None: + future.result() + trainer._pending_full_checkpoint_future = None + + return { + "saved": future is not None, + "path": checkpoint_path, + "reason": "saved" if future is not None else "not_checkpoint_leader", + } + + def _disable_standalone_sequence_parallel(draft_config) -> None: rollout_cfg = draft_config.rollout rollout_tp_size = int(rollout_cfg.get("tensor_model_parallel_size", 1) or 1) diff --git a/verl_speco/trainer/feature_store.py b/verl_speco/trainer/feature_store.py index 695748c3..bc2959af 100644 --- a/verl_speco/trainer/feature_store.py +++ b/verl_speco/trainer/feature_store.py @@ -265,10 +265,21 @@ def close(self) -> None: self.flush() def _write_metadata(self) -> None: - tmp_path = self.metadata_path.with_suffix(".json.tmp") - with tmp_path.open("w", encoding="utf-8") as metadata_file: + with tempfile.NamedTemporaryFile( + mode="w", + encoding="utf-8", + prefix=self.metadata_path.name, + suffix=".tmp", + dir=self.metadata_path.parent, + delete=False, + ) as metadata_file: + tmp_name = metadata_file.name json.dump(self.metadata, metadata_file, ensure_ascii=True, indent=2, sort_keys=True) - os.replace(tmp_path, self.metadata_path) + try: + os.replace(tmp_name, self.metadata_path) + finally: + if os.path.exists(tmp_name): + os.remove(tmp_name) def _load_manifest(self) -> list[dict[str, Any]]: if not self.manifest_path.exists(): diff --git a/verl_speco/workers/speco_worker.py b/verl_speco/workers/speco_worker.py index 8d188cae..997ce295 100644 --- a/verl_speco/workers/speco_worker.py +++ b/verl_speco/workers/speco_worker.py @@ -490,8 +490,24 @@ def _write_rollout_feature_sample( self.rank, ) return - input_ids = batch["input_ids"].detach().cpu().reshape(-1) - loss_mask = self._build_rollout_loss_mask(batch, input_ids) + full_input_ids = batch["input_ids"].detach().cpu().reshape(-1) + full_loss_mask = self._build_rollout_loss_mask(batch, full_input_ids) + hidden_states = hidden_states.detach().cpu() + hidden_rows = int(hidden_states.size(1) if hidden_states.dim() == 3 and hidden_states.size(0) == 1 else hidden_states.size(0)) + hidden_positions = batch.get("hidden_positions") + if torch.is_tensor(hidden_positions): + hidden_positions = hidden_positions.detach().cpu().long().reshape(-1) + else: + hidden_positions = None + feature_start, feature_end, position_ids = self._resolve_rollout_feature_window( + full_input_ids, + hidden_rows, + hidden_positions=hidden_positions, + hidden_position_start=batch.get("hidden_position_start"), + hidden_position_end=batch.get("hidden_position_end"), + ) + input_ids = full_input_ids[feature_start:feature_end] + loss_mask = full_loss_mask[feature_start:feature_end] metadata = { "source": batch.get("hidden_target_logprobs_source", "rl_rollout"), "global_step": batch.get("global_step", self.last_global_step), @@ -506,6 +522,9 @@ def _write_rollout_feature_sample( "use_logits": bool(self.config.rollout.drafter.training.get("use_logits", False)), "sequence_length": int(input_ids.numel()), "loss_tokens": int(loss_mask.sum().item()), + "full_sequence_length": int(full_input_ids.numel()), + "feature_start": int(feature_start), + "feature_end": int(feature_end), } for key in ( "hidden_position_start", @@ -528,12 +547,52 @@ def _write_rollout_feature_sample( algorithm=str(self.config.rollout.drafter.speculative_algorithm).upper(), input_ids=input_ids, loss_mask=loss_mask, - hidden_states=hidden_states.detach().cpu(), + hidden_states=hidden_states, target_logprobs=target_logprobs.detach().cpu() if torch.is_tensor(target_logprobs) else None, + position_ids=position_ids, metadata=metadata, ) writer.write_many([sample]) + @staticmethod + def _resolve_rollout_feature_window( + input_ids: torch.Tensor, + hidden_rows: int, + *, + hidden_positions: Optional[torch.Tensor], + hidden_position_start, + hidden_position_end, + ) -> tuple[int, int, torch.Tensor]: + input_len = int(input_ids.numel()) + hidden_rows = max(int(hidden_rows), 0) + if hidden_positions is not None and int(hidden_positions.numel()) > 0: + positions = hidden_positions[:hidden_rows].long() + start = int(positions[0].item()) + if int(positions.numel()) == hidden_rows and bool(torch.all(positions[1:] == positions[:-1] + 1).item()): + end = int(positions[-1].item()) + 1 + if 0 <= start < end <= input_len: + return start, end, positions + 1 + else: + positions = None + + try: + start = int(hidden_position_start) + except (TypeError, ValueError): + start = 0 + try: + end = int(hidden_position_end) + except (TypeError, ValueError): + end = start + hidden_rows + start = min(max(start, 0), input_len) + end = min(max(end, start), input_len) + if end - start != hidden_rows: + end = min(start + hidden_rows, input_len) + if end <= start: + start = 0 + end = min(hidden_rows, input_len) + position_ids = torch.arange(start + 1, end + 1, dtype=torch.long) + return start, end, position_ids + def _flush_rollout_features_for_step(self) -> None: if self._drafter_training_mode() != "collect_only" or self.feature_writer is None: return From 659949b055149c22260e134b1391c1cfb7df1622 Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Mon, 13 Jul 2026 16:01:53 +0800 Subject: [PATCH 08/50] fix(speco): restore offline draft feature alignment --- tests/unit/test_draft_feature_store.py | 34 +++++++++++++++++++++ verl_speco/inspect_feature_store.py | 16 ++++++++++ verl_speco/trainer/feature_store.py | 40 ++++++++++++++++++++++++ verl_speco/workers/speco_worker.py | 42 +++++++++++++++++++++++++- 4 files changed, 131 insertions(+), 1 deletion(-) diff --git a/tests/unit/test_draft_feature_store.py b/tests/unit/test_draft_feature_store.py index 6b19f8c1..d4e46a69 100644 --- a/tests/unit/test_draft_feature_store.py +++ b/tests/unit/test_draft_feature_store.py @@ -67,6 +67,40 @@ def test_feature_sample_rejects_position_id_length_mismatch(): sample.validate(strict=True) +def test_feature_sample_restores_online_alignment_metadata(): + sample = DraftFeatureSample( + input_ids=torch.arange(127, dtype=torch.long), + loss_mask=torch.ones(127, dtype=torch.float32), + hidden_states=torch.randn(127, 12288, dtype=torch.float32), + target_logprobs=torch.zeros(126, 128, 2, dtype=torch.float32), + position_ids=torch.arange(140, 267, dtype=torch.long), + metadata={ + "global_step": 1, + "hidden_states_layout": "eagle3_aux_plus_last", + "full_sequence_length": 1164, + "feature_start": 139, + "feature_end": 266, + "hidden_position_start": 139, + "hidden_position_end": 266, + "hidden_positions": torch.arange(139, 266, dtype=torch.long), + "target_logprobs_position_start": 140, + "target_logprobs_position_end": 266, + }, + ) + + item = sample.to_training_item() + + assert item["_verl_feature_start"] == 139 + assert item["_verl_feature_end"] == 266 + assert item["_verl_target_position_start"] == 140 + assert item["_verl_target_position_end"] == 266 + assert item["_verl_target_tensor_position_start"] == 140 + assert item["_verl_target_tensor_position_end"] == 266 + assert item["_verl_target_start"] == 0 + assert item["_verl_target_end"] == 126 + assert torch.equal(item["_verl_hidden_positions"], torch.arange(139, 266, dtype=torch.long)) + + def test_draft_feature_dataloader_slices_keys_by_rank(tmp_path): store = TorchShardFeatureStore(tmp_path, max_samples_per_shard=4) store.write_many([_sample(i) for i in range(4)]) diff --git a/verl_speco/inspect_feature_store.py b/verl_speco/inspect_feature_store.py index 943fb014..b86c8ae1 100644 --- a/verl_speco/inspect_feature_store.py +++ b/verl_speco/inspect_feature_store.py @@ -109,6 +109,7 @@ def _sample_issues(sample: dict[str, Any]) -> list[str]: loss_mask = _tensor(sample.get("loss_mask"), "loss_mask", issues) hidden_states = sample.get("hidden_states") position_ids = sample.get("position_ids") + target_logprobs = sample.get("target_logprobs") seq_len = None if input_ids is not None: @@ -144,6 +145,21 @@ def _sample_issues(sample: dict[str, Any]) -> list[str]: issues.append(f"position_ids length {position_ids.numel()} does not match input_ids length {seq_len}") elif position_ids.dim() > 1: issues.append(f"position_ids is normalizable but not stored as 1D: {_shape(position_ids)}") + if target_logprobs is not None: + if not torch.is_tensor(target_logprobs): + issues.append(f"target_logprobs is not a tensor: {type(target_logprobs).__name__}") + else: + normalized = target_logprobs + while normalized.dim() > 3 and normalized.size(0) == 1: + normalized = normalized.squeeze(0) + if normalized.dim() != 3 or normalized.size(-1) < 2: + issues.append(f"target_logprobs should be [rows, topk, 2], got {_shape(target_logprobs)}") + elif target_logprobs.dim() > 3: + issues.append(f"target_logprobs is normalizable but not stored as 3D: {_shape(target_logprobs)}") + elif seq_len is not None and int(normalized.size(0)) < max(seq_len - 2, 1): + issues.append( + f"target_logprobs rows {int(normalized.size(0))} may be too short for input_ids length {seq_len}" + ) return issues diff --git a/verl_speco/trainer/feature_store.py b/verl_speco/trainer/feature_store.py index bc2959af..4d646d8a 100644 --- a/verl_speco/trainer/feature_store.py +++ b/verl_speco/trainer/feature_store.py @@ -71,6 +71,9 @@ def validate(self, *, strict: bool = True) -> None: self.loss_mask = self.loss_mask.reshape(-1) if torch.is_tensor(self.position_ids) and self.position_ids.dim() > 1: self.position_ids = self.position_ids.reshape(-1) + if torch.is_tensor(self.target_logprobs): + while self.target_logprobs.dim() > 3 and self.target_logprobs.size(0) == 1: + self.target_logprobs = self.target_logprobs.squeeze(0) if self.input_ids.size(0) != self.loss_mask.size(0) and strict: raise ValueError( "DraftFeatureSample input_ids/loss_mask length mismatch: " @@ -87,6 +90,11 @@ def validate(self, *, strict: bool = True) -> None: self.last_hidden_states = self.last_hidden_states.squeeze(0) if self.target_logprobs is not None and not torch.is_tensor(self.target_logprobs): raise TypeError("DraftFeatureSample.target_logprobs must be a tensor when provided") + if torch.is_tensor(self.target_logprobs) and self.target_logprobs.dim() != 3 and strict: + raise ValueError( + "DraftFeatureSample.target_logprobs must have shape [rows, topk, 2], " + f"got {tuple(self.target_logprobs.shape)}" + ) def to_dict(self) -> dict[str, Any]: self.validate(strict=False) @@ -129,6 +137,7 @@ def to_training_item(self) -> dict[str, Any]: item["position_ids"] = payload["position_ids"] for key, value in metadata.items(): item.setdefault(key, value) + _populate_verl_alignment_fields(item, metadata) return item @@ -144,6 +153,37 @@ def get_metadata(self) -> dict[str, Any]: ... def close(self) -> None: ... +def _populate_verl_alignment_fields(item: dict[str, Any], metadata: dict[str, Any]) -> None: + """Restore online drafter alignment metadata for feature-store samples.""" + + direct_fields = { + "_verl_feature_start": "feature_start", + "_verl_feature_end": "feature_end", + "_verl_hidden_position_start": "hidden_position_start", + "_verl_hidden_position_end": "hidden_position_end", + "_verl_target_position_start": "target_logprobs_position_start", + "_verl_target_position_end": "target_logprobs_position_end", + "_verl_target_tensor_position_start": "target_logprobs_position_start", + "_verl_target_tensor_position_end": "target_logprobs_position_end", + "_verl_hidden_raw_target_position_start": "hidden_raw_target_logprobs_position_start", + "_verl_hidden_raw_target_position_end": "hidden_raw_target_logprobs_position_end", + "_verl_input_seq_length": "full_sequence_length", + } + for target_key, source_key in direct_fields.items(): + if target_key not in item and source_key in metadata: + item[target_key] = metadata[source_key] + + if "_verl_hidden_positions" not in item and "hidden_positions" in metadata: + item["_verl_hidden_positions"] = metadata["hidden_positions"] + if "_verl_uses_hidden_positions" not in item: + item["_verl_uses_hidden_positions"] = "hidden_positions" in metadata + + if "_verl_target_start" not in item and "target_logprobs" in item: + item["_verl_target_start"] = 0 + if "_verl_target_end" not in item and torch.is_tensor(item.get("target_logprobs")): + item["_verl_target_end"] = int(item["target_logprobs"].size(0)) + + class TorchShardFeatureStore: """Local ``torch.save`` shard store. diff --git a/verl_speco/workers/speco_worker.py b/verl_speco/workers/speco_worker.py index 997ce295..fb1a132d 100644 --- a/verl_speco/workers/speco_worker.py +++ b/verl_speco/workers/speco_worker.py @@ -508,6 +508,13 @@ def _write_rollout_feature_sample( ) input_ids = full_input_ids[feature_start:feature_end] loss_mask = full_loss_mask[feature_start:feature_end] + target_logprobs = self._align_rollout_target_logprobs( + target_logprobs, + feature_start=feature_start, + train_rows=max(int(input_ids.numel()) - 1, 0), + target_position_start=batch.get("target_logprobs_position_start"), + target_position_end=batch.get("target_logprobs_position_end"), + ) metadata = { "source": batch.get("hidden_target_logprobs_source", "rl_rollout"), "global_step": batch.get("global_step", self.last_global_step), @@ -548,12 +555,45 @@ def _write_rollout_feature_sample( input_ids=input_ids, loss_mask=loss_mask, hidden_states=hidden_states, - target_logprobs=target_logprobs.detach().cpu() if torch.is_tensor(target_logprobs) else None, + target_logprobs=target_logprobs, position_ids=position_ids, metadata=metadata, ) writer.write_many([sample]) + @staticmethod + def _align_rollout_target_logprobs( + target_logprobs: Optional[torch.Tensor], + *, + feature_start: int, + train_rows: int, + target_position_start, + target_position_end, + ) -> Optional[torch.Tensor]: + if not torch.is_tensor(target_logprobs): + return None + target = target_logprobs.detach().cpu() + while target.dim() > 3 and target.size(0) == 1: + target = target.squeeze(0) + if target.dim() != 3: + return target.contiguous() + + try: + position_start = int(target_position_start) + except (TypeError, ValueError): + position_start = int(feature_start) + 1 + try: + position_end = int(target_position_end) + except (TypeError, ValueError): + position_end = position_start + int(target.size(0)) + position_end = min(max(position_end, position_start), position_start + int(target.size(0))) + + desired_start = int(feature_start) + 1 + desired_end = desired_start + max(int(train_rows), 0) + slice_start = min(max(desired_start - position_start, 0), int(target.size(0))) + slice_end = min(max(desired_end - position_start, slice_start), int(position_end - position_start)) + return target[slice_start:slice_end].contiguous() + @staticmethod def _resolve_rollout_feature_window( input_ids: torch.Tensor, From 85fc79e3bc96b85ba4bf1c7543fb3518fe566334 Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Mon, 13 Jul 2026 16:35:59 +0800 Subject: [PATCH 09/50] fix(speco): point verl contract test at shared base config --- tests/compat/test_verl_contract.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/compat/test_verl_contract.py b/tests/compat/test_verl_contract.py index b9205de7..0957f024 100644 --- a/tests/compat/test_verl_contract.py +++ b/tests/compat/test_verl_contract.py @@ -7,7 +7,7 @@ ROOT = Path(__file__).resolve().parents[2] REQUIRED_VERL = ROOT / "REQUIRED_VERL.txt" -OVERLAY_CONFIG = ROOT / "verl_speco" / "config" / "speco_trainer.yaml" +OVERLAY_CONFIG = ROOT / "verl_speco" / "config" / "speco_base.yaml" def _required_verl_values() -> dict[str, str]: From 05552bdaad24c5f2b9d3c080c16d63102fb42aeb Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Tue, 14 Jul 2026 10:48:05 +0800 Subject: [PATCH 10/50] fix(speco): address standalone draft review feedback --- tests/unit/test_draft_feature_store.py | 20 +++++++++ tests/unit/test_draft_training_loop.py | 55 +++++++++++++++++++++++ verl_speco/trainer/draft_dataset.py | 10 ++++- verl_speco/trainer/draft_training_loop.py | 11 ++--- verl_speco/workers/speco_worker.py | 11 +++-- 5 files changed, 98 insertions(+), 9 deletions(-) create mode 100644 tests/unit/test_draft_training_loop.py diff --git a/tests/unit/test_draft_feature_store.py b/tests/unit/test_draft_feature_store.py index d4e46a69..b198165e 100644 --- a/tests/unit/test_draft_feature_store.py +++ b/tests/unit/test_draft_feature_store.py @@ -121,6 +121,26 @@ def test_draft_feature_dataloader_slices_keys_by_rank(tmp_path): assert rank1_ids == [2, 4] +def test_draft_feature_dataloader_rejects_rank_out_of_range(tmp_path): + store = TorchShardFeatureStore(tmp_path, max_samples_per_shard=4) + + with pytest.raises(ValueError, match="Invalid rank/world_size configuration"): + DraftFeatureDataLoader( + store, + DraftFeatureDataLoaderConfig(batch_size=1, rank=2, world_size=2), + ) + + +def test_draft_feature_dataloader_rejects_non_positive_world_size(tmp_path): + store = TorchShardFeatureStore(tmp_path, max_samples_per_shard=4) + + with pytest.raises(ValueError, match="Invalid world_size"): + DraftFeatureDataLoader( + store, + DraftFeatureDataLoaderConfig(batch_size=1, rank=0, world_size=0), + ) + + def test_flush_interval_zero_relies_on_shard_capacity(tmp_path): store = TorchShardFeatureStore(tmp_path, max_samples_per_shard=4) store.write_many([_sample(0), _sample(1)]) diff --git a/tests/unit/test_draft_training_loop.py b/tests/unit/test_draft_training_loop.py new file mode 100644 index 00000000..6f9e57c4 --- /dev/null +++ b/tests/unit/test_draft_training_loop.py @@ -0,0 +1,55 @@ +from __future__ import annotations + +from concurrent.futures import Future +from types import SimpleNamespace + +import pytest + +pytest.importorskip("torch") + +from verl_speco.trainer.draft_training_loop import _save_standalone_checkpoint + + +class _FakeTrainer: + def __init__(self): + self.checkpoint_dir = "/tmp/draft" + self._pending_full_checkpoint_future = None + self.future = Future() + self.calls = 0 + + def _save_checkpoint_async(self, step: int): + self.calls += 1 + self.step = step + self._pending_full_checkpoint_future = self.future + return self.future + + +def test_standalone_checkpoint_schedules_without_waiting(): + trainer = _FakeTrainer() + + result = _save_standalone_checkpoint(trainer, 5) + + assert result["saved"] is True + assert result["reason"] == "scheduled" + assert trainer.calls == 1 + assert trainer._pending_full_checkpoint_future is trainer.future + + +def test_standalone_checkpoint_waits_when_requested(): + trainer = _FakeTrainer() + trainer.future.set_result(None) + + result = _save_standalone_checkpoint(trainer, 5, wait=True) + + assert result["saved"] is True + assert result["reason"] == "saved" + assert trainer._pending_full_checkpoint_future is None + + +def test_standalone_checkpoint_skips_when_previous_save_is_running(): + trainer = SimpleNamespace(checkpoint_dir="/tmp/draft", _pending_full_checkpoint_future=Future()) + + result = _save_standalone_checkpoint(trainer, 5) + + assert result["saved"] is False + assert result["reason"] == "previous_save_running" diff --git a/verl_speco/trainer/draft_dataset.py b/verl_speco/trainer/draft_dataset.py index 01407f07..e83a808d 100644 --- a/verl_speco/trainer/draft_dataset.py +++ b/verl_speco/trainer/draft_dataset.py @@ -29,6 +29,14 @@ class DraftFeatureDataLoader: def __init__(self, store: DraftFeatureStore, config: DraftFeatureDataLoaderConfig): self.store = store self.config = config + rank = int(config.rank) + world_size = int(config.world_size) + if world_size <= 0: + raise ValueError(f"Invalid world_size: {world_size}") + if not (0 <= rank < world_size): + raise ValueError( + f"Invalid rank/world_size configuration: rank={rank}, world_size={world_size}" + ) def __iter__(self) -> Iterator[list[DraftFeatureSample]]: epoch = 0 @@ -41,7 +49,7 @@ def __iter__(self) -> Iterator[list[DraftFeatureSample]]: ) if not keys: return - rank_keys = keys[int(self.config.rank) :: max(int(self.config.world_size), 1)] + rank_keys = keys[int(self.config.rank) :: int(self.config.world_size)] batch: list[DraftFeatureSample] = [] for key in rank_keys: batch.append(self.store.read(key)) diff --git a/verl_speco/trainer/draft_training_loop.py b/verl_speco/trainer/draft_training_loop.py index c2ba862e..695272ab 100644 --- a/verl_speco/trainer/draft_training_loop.py +++ b/verl_speco/trainer/draft_training_loop.py @@ -96,11 +96,12 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: successful_steps += 1 if save_interval > 0 and successful_steps % save_interval == 0: last_save_result = _save_standalone_checkpoint(trainer, successful_steps) - last_saved_step = successful_steps + if last_save_result.get("saved"): + last_saved_step = successful_steps _barrier() final_save = bool(training_cfg.get("save_final_checkpoint", True)) if final_save and successful_steps > 0 and successful_steps != last_saved_step: - last_save_result = _save_standalone_checkpoint(trainer, successful_steps) + last_save_result = _save_standalone_checkpoint(trainer, successful_steps, wait=True) _barrier() finally: if store is not None: @@ -132,7 +133,7 @@ def _build_backend(draft_config): raise ValueError(f"Unsupported drafter algorithm {algo!r}; expected EAGLE3, DFLASH or DSPARK") -def _save_standalone_checkpoint(trainer: DrafterBaseTrainer, step: int) -> dict[str, Any]: +def _save_standalone_checkpoint(trainer: DrafterBaseTrainer, step: int, *, wait: bool = False) -> dict[str, Any]: if not trainer.checkpoint_dir: return {"saved": False, "reason": "missing_checkpoint_dir"} @@ -147,14 +148,14 @@ def _save_standalone_checkpoint(trainer: DrafterBaseTrainer, step: int) -> dict[ } future = trainer._save_checkpoint_async(int(step)) - if future is not None: + if future is not None and wait: future.result() trainer._pending_full_checkpoint_future = None return { "saved": future is not None, "path": checkpoint_path, - "reason": "saved" if future is not None else "not_checkpoint_leader", + "reason": "saved" if future is not None and wait else "scheduled" if future is not None else "not_checkpoint_leader", } diff --git a/verl_speco/workers/speco_worker.py b/verl_speco/workers/speco_worker.py index fb1a132d..8d7b5bbc 100644 --- a/verl_speco/workers/speco_worker.py +++ b/verl_speco/workers/speco_worker.py @@ -444,13 +444,15 @@ def _get_feature_writer(self) -> Optional[TorchShardFeatureStore]: return None if self.feature_writer is not None and self.feature_writer_path == path: return self.feature_writer + model_cfg = self.config.get("model", None) + target_model_path = _config_str(model_cfg.get("path", None)) if model_cfg is not None else "" self.feature_writer = TorchShardFeatureStore( path, max_samples_per_shard=int(feature_store_cfg.get("max_samples_per_shard", 1024)), strict_schema=bool(feature_store_cfg.get("strict_schema", True)), metadata={ "algorithm": str(self.config.rollout.drafter.speculative_algorithm).upper(), - "target_model_path": _config_str(getattr(self.config.model, "path", None)), + "target_model_path": target_model_path, "drafter_model_path": _config_str(self.config.rollout.drafter.get("model_path", None)), "source": "rl_collect_only", }, @@ -469,7 +471,8 @@ def _build_rollout_loss_mask(self, batch: dict, input_ids: torch.Tensor) -> torc if torch.is_tensor(prompts) and torch.is_tensor(responses): prompt_len = int(prompts.reshape(-1).numel()) response_ids = responses.detach().cpu().reshape(-1) - pad_token_id = int(getattr(self.config.model, "pad_token_id", 0) or 0) + model_cfg = self.config.get("model", None) + pad_token_id = int(model_cfg.get("pad_token_id", 0) or 0) if model_cfg is not None else 0 max_response = max(0, min(int(response_ids.numel()), int(ids.numel()) - prompt_len)) if max_response > 0: loss_mask[prompt_len : prompt_len + max_response] = (response_ids[:max_response] != pad_token_id).float() @@ -515,10 +518,12 @@ def _write_rollout_feature_sample( target_position_start=batch.get("target_logprobs_position_start"), target_position_end=batch.get("target_logprobs_position_end"), ) + model_cfg = self.config.get("model", None) + target_model_path = _config_str(model_cfg.get("path", None)) if model_cfg is not None else "" metadata = { "source": batch.get("hidden_target_logprobs_source", "rl_rollout"), "global_step": batch.get("global_step", self.last_global_step), - "target_model_path": _config_str(getattr(self.config.model, "path", None)), + "target_model_path": target_model_path, "drafter_model_path": _config_str(self.config.rollout.drafter.get("model_path", None)), "hidden_states_layout": batch.get("hidden_states_layout") or ( "dflash_aux" From b2ba8f1d6b31fe83f4199dfae539cb6f16942b1c Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Tue, 14 Jul 2026 14:39:13 +0800 Subject: [PATCH 11/50] update readme --- README.md | 33 +++++++++++++++++++++++++++++++++ 1 file changed, 33 insertions(+) diff --git a/README.md b/README.md index 1aab1e38..426d02be 100644 --- a/README.md +++ b/README.md @@ -158,6 +158,39 @@ actor_rollout_ref.rollout.drafter.model_path=/path/to/drafter actor_rollout_ref.rollout.drafter.speculative_algorithm=EAGLE3 ``` +## Separate Draft Model Training + +verl-SpeCo also supports a separate draft model training workflow. In this +mode, rollout workers collect drafter training features into a feature store, +and the draft model can be trained separately after feature collection. + +Quickstart: + +```bash +bash examples/run_qwen3-8b_drafter_separate_training.sh +``` + +Replace the model, drafter, dataset, feature-store, and checkpoint paths in +the script before running it. The script uses `collect_only` mode for rollout +feature collection and `offline` mode for standalone drafter training. + +The main mode values are: + +| Mode | Meaning | +| --- | --- | +| `online` | Default. Collects rollout features, trains the drafter inside the online PPO/Ray workflow, and can publish updated drafter weights back to the rollout engine. | +| `collect_only` | Collects rollout features into `feature_store.path` without running drafter training in the PPO/Ray workflow. | +| `offline` | Reads collected features from `feature_store.path` and trains the drafter with the standalone multi-GPU workflow. | + +Collected feature stores can be inspected before offline training: + +```bash +python -m verl_speco.inspect_feature_store /path/to/features \ + --max-samples 200 \ + --show-ok \ + --strict-exit +``` + ## Configuration SPECO-specific options live under: From 782de130100f8ad55d7d69f1914a236b274978bb Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Tue, 14 Jul 2026 19:37:28 +0800 Subject: [PATCH 12/50] docs: add vLLM and SGLang GPU Dockerfiles --- Dockerfile.sglang | 17 +++++++++++++ Dockerfile.vllm | 17 +++++++++++++ README.md | 63 +++++++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 97 insertions(+) create mode 100644 Dockerfile.sglang create mode 100644 Dockerfile.vllm diff --git a/Dockerfile.sglang b/Dockerfile.sglang new file mode 100644 index 00000000..5df5fbdb --- /dev/null +++ b/Dockerfile.sglang @@ -0,0 +1,17 @@ +# GPU SGLang runtime image. +FROM verlai/verl:sgl0512.dev1 + +ARG VERL_COMMIT=7aed6b230776f963fa09509c10d9c3a767d1102c +ARG VERL_REPO=https://github.com/verl-project/verl.git + +WORKDIR /workspace + +RUN git clone ${VERL_REPO} /workspace/verl \ + && cd /workspace/verl \ + && git checkout ${VERL_COMMIT} \ + && pip install -e . + +COPY . /workspace/verl-SpeCo + +ENV PYTHONPATH=/workspace/verl-SpeCo:${PYTHONPATH} +WORKDIR /workspace/verl-SpeCo diff --git a/Dockerfile.vllm b/Dockerfile.vllm new file mode 100644 index 00000000..621b2ce5 --- /dev/null +++ b/Dockerfile.vllm @@ -0,0 +1,17 @@ +# GPU vLLM runtime image. +FROM verlai/verl:vllm023.dev1 + +ARG VERL_COMMIT=7aed6b230776f963fa09509c10d9c3a767d1102c +ARG VERL_REPO=https://github.com/verl-project/verl.git + +WORKDIR /workspace + +RUN git clone ${VERL_REPO} /workspace/verl \ + && cd /workspace/verl \ + && git checkout ${VERL_COMMIT} \ + && pip install -e . + +COPY . /workspace/verl-SpeCo + +ENV PYTHONPATH=/workspace/verl-SpeCo:${PYTHONPATH} +WORKDIR /workspace/verl-SpeCo diff --git a/README.md b/README.md index 426d02be..4e82718f 100644 --- a/README.md +++ b/README.md @@ -121,6 +121,69 @@ cd /path/to/verl-SpeCo export PYTHONPATH="$PWD:$PYTHONPATH" ``` +### Docker Images + +You can also build GPU runtime images from the official `verlai/verl` +development images and then pin the importable upstream `verl` checkout to the +required v0.8.0 commit. The Dockerfiles below target GPU deployments; use the +matching accelerator image for NPU or other accelerator runtimes. + +For GPU vLLM-based examples, use this Dockerfile: + +```dockerfile +# GPU vLLM runtime image. +FROM verlai/verl:vllm023.dev1 + +ARG VERL_COMMIT=7aed6b230776f963fa09509c10d9c3a767d1102c +ARG VERL_REPO=https://github.com/verl-project/verl.git + +WORKDIR /workspace + +RUN git clone ${VERL_REPO} /workspace/verl \ + && cd /workspace/verl \ + && git checkout ${VERL_COMMIT} \ + && pip install -e . + +COPY . /workspace/verl-SpeCo + +ENV PYTHONPATH=/workspace/verl-SpeCo:${PYTHONPATH} +WORKDIR /workspace/verl-SpeCo +``` + +Build it from the `verl-SpeCo` repository root: + +```bash +docker build -f Dockerfile.vllm -t verl-speco:vllm023-verl080 . +``` + +For GPU SGLang-based examples, use the same layout with the SGLang base image: + +```dockerfile +# GPU SGLang runtime image. +FROM verlai/verl:sgl0512.dev1 + +ARG VERL_COMMIT=7aed6b230776f963fa09509c10d9c3a767d1102c +ARG VERL_REPO=https://github.com/verl-project/verl.git + +WORKDIR /workspace + +RUN git clone ${VERL_REPO} /workspace/verl \ + && cd /workspace/verl \ + && git checkout ${VERL_COMMIT} \ + && pip install -e . + +COPY . /workspace/verl-SpeCo + +ENV PYTHONPATH=/workspace/verl-SpeCo:${PYTHONPATH} +WORKDIR /workspace/verl-SpeCo +``` + +Build it from the `verl-SpeCo` repository root: + +```bash +docker build -f Dockerfile.sglang -t verl-speco:sgl0512-verl080 . +``` + Install the rollout engine and accelerator runtime that match the script you intend to run, for example vLLM on GPU, SGLang on GPU, or vLLM-Ascend on NPU. Those runtime packages are intentionally not pinned by this repository. From 54991543e63d24da205a3a1ab22d47e127edaedd Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Wed, 15 Jul 2026 17:35:12 +0800 Subject: [PATCH 13/50] fix: export standalone drafter runtime config --- tests/unit/test_draft_training_loop.py | 76 ++++++++++- verl_speco/trainer/draft_training_loop.py | 157 +++++++++++++++++++++- 2 files changed, 230 insertions(+), 3 deletions(-) diff --git a/tests/unit/test_draft_training_loop.py b/tests/unit/test_draft_training_loop.py index 6f9e57c4..a741f908 100644 --- a/tests/unit/test_draft_training_loop.py +++ b/tests/unit/test_draft_training_loop.py @@ -1,5 +1,6 @@ from __future__ import annotations +import json from concurrent.futures import Future from types import SimpleNamespace @@ -7,7 +8,7 @@ pytest.importorskip("torch") -from verl_speco.trainer.draft_training_loop import _save_standalone_checkpoint +from verl_speco.trainer.draft_training_loop import _rewrite_standalone_block_runtime_config, _save_standalone_checkpoint class _FakeTrainer: @@ -53,3 +54,76 @@ def test_standalone_checkpoint_skips_when_previous_save_is_running(): assert result["saved"] is False assert result["reason"] == "previous_save_running" + + +def test_standalone_dspark_checkpoint_preserves_source_runtime_config(tmp_path): + checkpoint_dir = tmp_path / "draft_step_5" + checkpoint_dir.mkdir() + source_dir = tmp_path / "source_dspark" + source_dir.mkdir() + source_config = { + "model_type": "deepseek_v3", + "architectures": ["DeepSeekDSparkModel"], + "target_layer_ids": [1, 9, 17], + } + (source_dir / "config.json").write_text(json.dumps(source_config), encoding="utf-8") + training_config = { + "model_type": "dspark", + "architectures": ["DSparkDraftModel"], + "target_layer_ids": [1, 9, 17], + "mask_token_id": 151669, + "markov_head_type": "vanilla", + "markov_rank": 256, + "block_size": 7, + "num_context_layers": 3, + } + (checkpoint_dir / "config.json").write_text(json.dumps(training_config), encoding="utf-8") + trainer = SimpleNamespace( + backend=SimpleNamespace(model_type="dspark"), + config=SimpleNamespace(rollout=SimpleNamespace(drafter=SimpleNamespace(model_path=str(source_dir)))), + ) + + _rewrite_standalone_block_runtime_config(trainer, str(checkpoint_dir)) + + runtime_config = json.loads((checkpoint_dir / "config.json").read_text(encoding="utf-8")) + saved_training_config = json.loads((checkpoint_dir / "speco_training_config.json").read_text(encoding="utf-8")) + assert runtime_config["model_type"] == "deepseek_v3" + assert runtime_config["architectures"] == ["DeepSeekDSparkModel"] + assert runtime_config["dspark_config"]["markov_head_type"] == "vanilla" + assert runtime_config["dflash_config"]["target_layer_ids"] == [1, 9, 17] + assert runtime_config["eagle_aux_hidden_state_layer_ids"] == [2, 10, 18] + assert saved_training_config == training_config + + +def test_standalone_dflash_checkpoint_preserves_source_runtime_config(tmp_path): + checkpoint_dir = tmp_path / "draft_step_5" + checkpoint_dir.mkdir() + source_dir = tmp_path / "source_dflash" + source_dir.mkdir() + source_config = { + "model_type": "qwen3", + "architectures": ["DFlashForCausalLM"], + } + (source_dir / "config.json").write_text(json.dumps(source_config), encoding="utf-8") + training_config = { + "model_type": "dflash", + "architectures": ["DFlashDraftModel"], + "target_layer_ids": [2, 10, 18], + "mask_token_id": 151669, + "num_context_layers": 3, + } + (checkpoint_dir / "config.json").write_text(json.dumps(training_config), encoding="utf-8") + trainer = SimpleNamespace( + backend=SimpleNamespace(model_type="dflash"), + config=SimpleNamespace(rollout=SimpleNamespace(drafter=SimpleNamespace(model_path=str(source_dir)))), + ) + + _rewrite_standalone_block_runtime_config(trainer, str(checkpoint_dir)) + + runtime_config = json.loads((checkpoint_dir / "config.json").read_text(encoding="utf-8")) + saved_training_config = json.loads((checkpoint_dir / "speco_training_config.json").read_text(encoding="utf-8")) + assert runtime_config["model_type"] == "qwen3" + assert runtime_config["architectures"] == ["DFlashForCausalLM"] + assert runtime_config["dflash_config"]["target_layer_ids"] == [2, 10, 18] + assert runtime_config["eagle_aux_hidden_state_layer_ids"] == [3, 11, 19] + assert saved_training_config == training_config diff --git a/verl_speco/trainer/draft_training_loop.py b/verl_speco/trainer/draft_training_loop.py index 695272ab..3c575357 100644 --- a/verl_speco/trainer/draft_training_loop.py +++ b/verl_speco/trainer/draft_training_loop.py @@ -3,6 +3,8 @@ from __future__ import annotations import asyncio +from copy import deepcopy +import json import logging import os from typing import Any @@ -96,7 +98,7 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: successful_steps += 1 if save_interval > 0 and successful_steps % save_interval == 0: last_save_result = _save_standalone_checkpoint(trainer, successful_steps) - if last_save_result.get("saved"): + if _sync_any_rank_saved_checkpoint(last_save_result.get("saved")): last_saved_step = successful_steps _barrier() final_save = bool(training_cfg.get("save_final_checkpoint", True)) @@ -108,7 +110,6 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: store.close() await trainer.cleanup_training(clear_data=True) if dist.is_initialized(): - dist.barrier() dist.destroy_process_group() return { @@ -151,6 +152,9 @@ def _save_standalone_checkpoint(trainer: DrafterBaseTrainer, step: int, *, wait: if future is not None and wait: future.result() trainer._pending_full_checkpoint_future = None + _rewrite_standalone_block_runtime_config(trainer, checkpoint_path) + elif future is not None: + future.add_done_callback(lambda completed: _rewrite_standalone_block_runtime_config(trainer, checkpoint_path, completed)) return { "saved": future is not None, @@ -159,6 +163,140 @@ def _save_standalone_checkpoint(trainer: DrafterBaseTrainer, step: int, *, wait: } +def _ensure_dict_child(config: dict[str, Any], key: str) -> dict[str, Any]: + value = config.get(key) + if isinstance(value, dict): + return value + value = {} + config[key] = value + return value + + +def _load_source_drafter_config(trainer: DrafterBaseTrainer) -> dict[str, Any] | None: + model_path = getattr(getattr(getattr(trainer, "config", None), "rollout", None), "drafter", None) + model_path = getattr(model_path, "model_path", None) + if not model_path: + return None + config_path = os.path.join(os.fspath(model_path), "config.json") + if not os.path.exists(config_path): + return None + try: + with open(config_path, "r", encoding="utf-8") as f: + loaded = json.load(f) + except (OSError, json.JSONDecodeError) as exc: + logger.warning("Failed to load source drafter config %s: %s", config_path, exc) + return None + return loaded if isinstance(loaded, dict) else None + + +def _fill_if_missing(dst: dict[str, Any], src: dict[str, Any], keys: tuple[str, ...]) -> None: + for key in keys: + if key in src and key not in dst: + dst[key] = deepcopy(src[key]) + + +def _rewrite_standalone_block_runtime_config( + trainer: DrafterBaseTrainer, + checkpoint_path: str, + completed_future=None, +) -> None: + """Export standalone DFlash/DSpark checkpoints with runtime-facing config. + + The training wrapper saves an internal SpeCo config. For standalone + checkpoints we keep the original drafter ``config.json`` as the runtime + contract and only merge the alias fields needed by vLLM/SGLang. + """ + backend_type = getattr(getattr(trainer, "backend", None), "model_type", None) + if backend_type not in {"dflash", "dspark"}: + return + + if completed_future is not None: + try: + completed_future.result() + except Exception as exc: # noqa: BLE001 + logger.warning("Skip standalone runtime config rewrite because checkpoint save failed: %s", exc) + return + + config_path = os.path.join(checkpoint_path, "config.json") + if not os.path.exists(config_path): + logger.warning("Cannot rewrite standalone runtime config: missing %s", config_path) + return + + try: + with open(config_path, "r", encoding="utf-8") as f: + training_config = json.load(f) + except (OSError, json.JSONDecodeError) as exc: + logger.warning("Cannot rewrite standalone runtime config %s: %s", config_path, exc) + return + if not isinstance(training_config, dict): + logger.warning("Cannot rewrite standalone runtime config %s: expected object", config_path) + return + + training_config_path = os.path.join(checkpoint_path, "speco_training_config.json") + try: + with open(training_config_path, "w", encoding="utf-8") as f: + json.dump(training_config, f, indent=2, sort_keys=True) + except OSError as exc: + logger.warning("Failed to write standalone training config copy %s: %s", training_config_path, exc) + + runtime_config = _load_source_drafter_config(trainer) + if runtime_config is None: + runtime_config = deepcopy(training_config) + logger.warning( + "Source drafter config is unavailable; standalone checkpoint keeps SpeCo training config as runtime config" + ) + + runtime_config["speco_training_model_type"] = backend_type + common_alias_keys = ("target_layer_ids", "mask_token_id", "num_context_layers") + _fill_if_missing(runtime_config, training_config, common_alias_keys) + + dflash_config = _ensure_dict_child(runtime_config, "dflash_config") + _fill_if_missing(dflash_config, training_config, common_alias_keys) + + if backend_type == "dspark": + dspark_config = _ensure_dict_child(runtime_config, "dspark_config") + _fill_if_missing( + dspark_config, + training_config, + ( + "block_size", + "num_anchors", + "markov_rank", + "markov_head_type", + "confidence_head_alpha", + "confidence_head_with_markov", + "ce_loss_alpha", + "l1_loss_alpha", + "loss_decay_gamma", + "target_layer_ids", + "num_context_layers", + "num_target_layers", + "target_num_hidden_layers", + "mask_token_id", + ), + ) + else: + dspark_config = {} + + target_layer_ids = ( + runtime_config.get("target_layer_ids") + or dflash_config.get("target_layer_ids") + or dspark_config.get("target_layer_ids") + ) + if target_layer_ids is not None and "eagle_aux_hidden_state_layer_ids" not in runtime_config: + try: + runtime_config["eagle_aux_hidden_state_layer_ids"] = [int(layer_id) + 1 for layer_id in target_layer_ids] + except (TypeError, ValueError): + logger.warning("Invalid target_layer_ids in standalone exported config: %r", target_layer_ids) + + try: + with open(config_path, "w", encoding="utf-8") as f: + json.dump(runtime_config, f, indent=2, sort_keys=True) + f.write("\n") + except OSError as exc: + logger.warning("Failed to write standalone runtime config %s: %s", config_path, exc) + + def _disable_standalone_sequence_parallel(draft_config) -> None: rollout_cfg = draft_config.rollout rollout_tp_size = int(rollout_cfg.get("tensor_model_parallel_size", 1) or 1) @@ -211,6 +349,21 @@ def _barrier() -> None: dist.barrier() +def _sync_any_rank_saved_checkpoint(saved: Any) -> bool: + if not dist.is_initialized(): + return bool(saved) + device_name = get_device_name() + if device_name == "cpu": + device = torch.device("cpu") + else: + current_device = getattr(get_torch_device(), "current_device", None) + device_index = current_device() if callable(current_device) else 0 + device = torch.device(f"{device_name}:{int(device_index)}") + flag = torch.tensor([1 if saved else 0], dtype=torch.int32, device=device) + dist.all_reduce(flag, op=dist.ReduceOp.MAX) + return bool(flag.item()) + + def log_resolved_config(config) -> None: rank = int(os.environ.get("RANK", "0")) if rank == 0: From a1089fc0c002747767a5b16bde2e7ce9d4553fc1 Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Thu, 16 Jul 2026 17:14:27 +0800 Subject: [PATCH 14/50] Add trusted NPU DSpark example CI --- .github/workflows/npu_example_tests.yml | 38 ++++++---- ci/README.md | 53 ++++++++++++- ci/run_example_test.sh | 65 +++++++++++++++- tests/examples/test_ci_example_runner.py | 95 ++++++++++++++++++++---- 4 files changed, 221 insertions(+), 30 deletions(-) diff --git a/.github/workflows/npu_example_tests.yml b/.github/workflows/npu_example_tests.yml index b194e76c..a050b136 100644 --- a/.github/workflows/npu_example_tests.yml +++ b/.github/workflows/npu_example_tests.yml @@ -3,6 +3,9 @@ name: npu_example_tests run-name: NPU example tests (${{ github.ref_name }}) on: + pull_request: + branches: + - main workflow_dispatch: inputs: run_training: @@ -22,6 +25,10 @@ on: description: DFlash drafter path; falls back to SPECO_DFLASH_DRAFT_MODEL required: false type: string + dspark_draft_model: + description: DSpark drafter path; falls back to SPECO_DSPARK_DRAFT_MODEL + required: false + type: string train_file: description: Training dataset path; falls back to SPECO_TRAIN_FILE required: false @@ -57,7 +64,7 @@ permissions: concurrency: group: npu-example-tests-${{ github.ref }} - cancel-in-progress: false + cancel-in-progress: true env: SPECO_DEFAULT_MODEL_ROOT: ${{ vars.SPECO_MODEL_ROOT || '/home/runner/models' }} @@ -65,6 +72,7 @@ env: SPECO_DEFAULT_TARGET_MODEL: ${{ vars.SPECO_TARGET_MODEL || '/home/runner/models/Qwen/Qwen2.5-0.5B-Instruct' }} SPECO_DEFAULT_EAGLE3_DRAFT_MODEL: ${{ vars.SPECO_EAGLE3_DRAFT_MODEL || '/home/runner/models/speco/eagle3-drafter' }} SPECO_DEFAULT_DFLASH_DRAFT_MODEL: ${{ vars.SPECO_DFLASH_DRAFT_MODEL || '/home/runner/models/speco/dflash-drafter' }} + SPECO_DEFAULT_DSPARK_DRAFT_MODEL: ${{ vars.SPECO_DSPARK_DRAFT_MODEL || '/home/runner/models/speco/dspark-drafter' }} SPECO_DEFAULT_TRAIN_FILE: ${{ vars.SPECO_TRAIN_FILE || '/home/runner/models/hf_data/gsm8k/train.parquet' }} SPECO_DEFAULT_TEST_FILE: ${{ vars.SPECO_TEST_FILE || '/home/runner/models/hf_data/gsm8k/test.parquet' }} SPECO_DEFAULT_ACCELERATOR_COUNT: ${{ vars.SPECO_ACCELERATOR_COUNT || '1' }} @@ -74,21 +82,13 @@ env: jobs: example: name: NPU ${{ matrix.backend }} ${{ matrix.drafter }} example - runs-on: [self-hosted, linux, x64, npu] + if: ${{ github.event_name != 'pull_request' || github.event.pull_request.head.repo.full_name == github.repository }} + runs-on: [self-hosted, linux-aarch64-a2-8] timeout-minutes: 180 environment: speco-npu-ci strategy: fail-fast: false - matrix: - include: - - backend: vllm - drafter: eagle3 - - backend: vllm - drafter: dflash - - backend: sglang - drafter: eagle3 - - backend: sglang - drafter: dflash + matrix: ${{ fromJSON(github.event_name == 'pull_request' && '{"include":[{"backend":"vllm","drafter":"dspark"},{"backend":"sglang","drafter":"eagle3"},{"backend":"sglang","drafter":"dflash"}]}' || '{"include":[{"backend":"vllm","drafter":"eagle3"},{"backend":"vllm","drafter":"dflash"},{"backend":"vllm","drafter":"dspark"},{"backend":"sglang","drafter":"eagle3"},{"backend":"sglang","drafter":"dflash"}]}') }} steps: - name: Check out verl-SpeCo @@ -97,7 +97,12 @@ jobs: - name: Verify NPU runtime run: | npu-smi info - python -c "import torch, torch_npu; assert torch.npu.is_available()" + python - <<'PY' + import torch + import torch_npu + assert torch.npu.is_available() + print("torch.npu.device_count =", torch.npu.device_count()) + PY if [[ "${{ matrix.backend }}" == "vllm" ]]; then python -c "import vllm; print(vllm.__version__)" else @@ -110,12 +115,17 @@ jobs: SPECO_TARGET_MODEL: ${{ inputs.target_model || env.SPECO_DEFAULT_TARGET_MODEL }} SPECO_EAGLE3_DRAFT_MODEL: ${{ inputs.eagle3_draft_model || env.SPECO_DEFAULT_EAGLE3_DRAFT_MODEL }} SPECO_DFLASH_DRAFT_MODEL: ${{ inputs.dflash_draft_model || env.SPECO_DEFAULT_DFLASH_DRAFT_MODEL }} + SPECO_DSPARK_DRAFT_MODEL: ${{ inputs.dspark_draft_model || env.SPECO_DEFAULT_DSPARK_DRAFT_MODEL }} SPECO_TRAIN_FILE: ${{ inputs.train_file || env.SPECO_DEFAULT_TRAIN_FILE }} SPECO_TEST_FILE: ${{ inputs.test_file || env.SPECO_DEFAULT_TEST_FILE }} SPECO_CKPT_DIR: ${{ vars.SPECO_CKPT_DIR || runner.temp }} SPECO_ACCELERATOR_COUNT: ${{ inputs.accelerator_count || env.SPECO_DEFAULT_ACCELERATOR_COUNT }} SPECO_TENSOR_PARALLEL_SIZE: ${{ inputs.tensor_parallel_size || env.SPECO_DEFAULT_TENSOR_PARALLEL_SIZE }} SPECO_SEQUENCE_PARALLEL_SIZE: ${{ inputs.sequence_parallel_size || env.SPECO_DEFAULT_SEQUENCE_PARALLEL_SIZE }} - SPECO_ENABLE_TRAINING: ${{ github.event_name == 'schedule' && 'true' || inputs.run_training }} + SPECO_ENABLE_TRAINING: ${{ (github.event_name == 'schedule' || github.event_name == 'pull_request') && 'true' || inputs.run_training }} + SPECO_TOTAL_TRAINING_STEPS: ${{ github.event_name == 'pull_request' && '1' || vars.SPECO_TOTAL_TRAINING_STEPS }} + SPECO_TRAIN_MAX_SAMPLES: ${{ github.event_name == 'pull_request' && '1' || vars.SPECO_TRAIN_MAX_SAMPLES }} + SPECO_VAL_MAX_SAMPLES: ${{ github.event_name == 'pull_request' && '1' || vars.SPECO_VAL_MAX_SAMPLES }} + SPECO_DATALOADER_NUM_WORKERS: ${{ github.event_name == 'pull_request' && '0' || vars.SPECO_DATALOADER_NUM_WORKERS }} SPECO_EXTRA_HYDRA_ARGS: ${{ inputs.extra_hydra_args || vars.SPECO_EXTRA_HYDRA_ARGS }} run: bash ci/run_example_test.sh npu "${{ matrix.backend }}" "${{ matrix.drafter }}" diff --git a/ci/README.md b/ci/README.md index 0562c3c6..7915544d 100644 --- a/ci/README.md +++ b/ci/README.md @@ -4,12 +4,27 @@ The repository uses three workflow layers: - `cpu_unit_tests.yml`: required PR checks without installing this repository or using accelerator runtimes. - `gpu_example_tests.yml`: scheduled/manual vLLM and SGLang example-script runs on GPU. -- `npu_example_tests.yml`: scheduled/manual vLLM and SGLang example-script runs on NPU. +- `npu_example_tests.yml`: trusted PR, scheduled, and manual vLLM/SGLang example-script runs on NPU. The hardware workflows require self-hosted runner labels: - GPU: `self-hosted`, `linux`, `x64`, `gpu` -- NPU: `self-hosted`, `linux`, `x64`, `npu` +- NPU: `self-hosted`, `linux-aarch64-a2-8` + +The NPU workflow intentionally does not run self-hosted jobs for forked pull +requests. Pull requests from the same repository run a one-step smoke matrix: + +- vLLM + DSpark +- SGLang + EAGLE3 +- SGLang + DFlash + +Scheduled and manual NPU runs use the broader matrix: + +- vLLM + EAGLE3 +- vLLM + DFlash +- vLLM + DSpark +- SGLang + EAGLE3 +- SGLang + DFlash Like verl's CI, the hardware workflows assume the runner image has the runtime stack and a small default model/data cache. By default they look under: @@ -39,6 +54,7 @@ GitHub environments, or pass them as manual workflow inputs where available: - `SPECO_TARGET_MODEL` - `SPECO_EAGLE3_DRAFT_MODEL` - `SPECO_DFLASH_DRAFT_MODEL` +- `SPECO_DSPARK_DRAFT_MODEL` - `SPECO_TRAIN_FILE` - `SPECO_TEST_FILE` - `SPECO_CKPT_DIR` @@ -51,8 +67,28 @@ GitHub environments, or pass them as manual workflow inputs where available: - `SPECO_SPEC_VERIFY_TOKENS` - `SPECO_DFLASH_NUM_ANCHORS` - `SPECO_DFLASH_MAX_WINDOW` +- `SPECO_DSPARK_BLOCK_SIZE` +- `SPECO_DSPARK_NUM_ANCHORS` +- `SPECO_DSPARK_MAX_WINDOW` +- `SPECO_TOTAL_TRAINING_STEPS` +- `SPECO_TRAIN_MAX_SAMPLES` +- `SPECO_VAL_MAX_SAMPLES` +- `SPECO_DATALOADER_NUM_WORKERS` - `SPECO_EXTRA_HYDRA_ARGS` +For NPU runs, `ci/run_example_test.sh` generates +`ASCEND_RT_VISIBLE_DEVICES=0,...,N-1` from `SPECO_ACCELERATOR_COUNT` when the +caller has not already set `ASCEND_RT_VISIBLE_DEVICES`. If the caller provides +`ASCEND_RT_VISIBLE_DEVICES`, the script preserves it and checks that +`SPECO_ACCELERATOR_COUNT` does not exceed the visible device count. + +PR smoke jobs force lightweight settings through environment variables: + +- `SPECO_TOTAL_TRAINING_STEPS=1` +- `SPECO_TRAIN_MAX_SAMPLES=1` +- `SPECO_VAL_MAX_SAMPLES=1` +- `SPECO_DATALOADER_NUM_WORKERS=0` + The runner image is responsible for providing the matching verl, vLLM/SGLang, PyTorch accelerator runtime, and model files. Hardware workflows deliberately fail closed when required models or datasets are absent. @@ -76,6 +112,19 @@ D:\git\bin\bash.exe -n ci/run_example_test.sh python -m pytest tests/compat tests/config tests/examples tests/integration -q ``` +You can inspect the selected script and Hydra overrides without launching a +model by setting `SPECO_DRY_RUN=true`: + +```bash +SPECO_DRY_RUN=true \ +SPECO_TARGET_MODEL=/models/target \ +SPECO_DSPARK_DRAFT_MODEL=/models/dspark \ +SPECO_TRAIN_FILE=/data/train.parquet \ +SPECO_TEST_FILE=/data/test.parquet \ +SPECO_CKPT_DIR=/tmp/speco \ +bash ci/run_example_test.sh npu vllm dspark +``` + To test the hardware workflows on GitHub, open Actions, choose `gpu_example_tests` or `npu_example_tests`, and run the workflow without inputs after preparing the default paths above. Fill the manual inputs only when you diff --git a/ci/run_example_test.sh b/ci/run_example_test.sh index c9e81e17..18104cd0 100644 --- a/ci/run_example_test.sh +++ b/ci/run_example_test.sh @@ -12,6 +12,9 @@ case "${platform}/${backend}/${drafter}" in gpu/vllm/dflash) example="examples/run_qwen3-8b_drafter_dflash_vllm.sh" ;; + gpu/vllm/dspark) + example="examples/run_qwen3-8b_drafter_dspark_vllm.sh" + ;; gpu/sglang/eagle3) example="examples/run_qwen3-8b_drafter_eagle3_sglang.sh" ;; @@ -24,6 +27,9 @@ case "${platform}/${backend}/${drafter}" in npu/vllm/dflash) example="examples/run_qwen3-8b_drafter_dflash_vllm_npu.sh" ;; + npu/vllm/dspark) + example="examples/run_qwen3-8b_drafter_dspark_vllm_npu.sh" + ;; npu/sglang/eagle3) example="examples/run_qwen3-8b_drafter_eagle3_sglang_npu.sh" ;; @@ -31,7 +37,7 @@ case "${platform}/${backend}/${drafter}" in example="examples/run_qwen3-8b_drafter_dflash_sglang.sh" ;; *) - echo "usage: $0 {gpu|npu} {vllm|sglang} {eagle3|dflash}" >&2 + echo "usage: $0 {gpu|npu} {vllm|sglang} {eagle3|dflash|dspark}" >&2 exit 2 ;; esac @@ -58,6 +64,10 @@ case "${drafter}" in draft_model="${SPECO_DFLASH_DRAFT_MODEL:-}" draft_algorithm="DFLASH" ;; + dspark) + draft_model="${SPECO_DSPARK_DRAFT_MODEL:-}" + draft_algorithm="DSPARK" + ;; esac if [[ -z "${draft_model}" ]]; then echo "required ${drafter} draft model environment variable is not set" >&2 @@ -69,6 +79,22 @@ tensor_parallel_size="${SPECO_TENSOR_PARALLEL_SIZE:-1}" sequence_parallel_size="${SPECO_SEQUENCE_PARALLEL_SIZE:-1}" if [[ "${platform}" == "npu" ]]; then + if [[ "${SPECO_DRY_RUN:-false}" != "true" ]]; then + physical_npu_count="$(python - <<'PY' +import torch +import torch_npu +print(torch.npu.device_count()) +PY +)" + if (( accelerator_count > physical_npu_count )); then + echo "SPECO_ACCELERATOR_COUNT=${accelerator_count} exceeds physical NPU count ${physical_npu_count}" >&2 + exit 2 + fi + fi + if (( accelerator_count < 1 )); then + echo "SPECO_ACCELERATOR_COUNT must be >= 1, got ${accelerator_count}" >&2 + exit 2 + fi if [[ -z "${ASCEND_RT_VISIBLE_DEVICES:-}" ]]; then visible_devices="" for ((device_index = 0; device_index < accelerator_count; device_index++)); do @@ -78,6 +104,15 @@ if [[ "${platform}" == "npu" ]]; then visible_devices+="${device_index}" done export ASCEND_RT_VISIBLE_DEVICES="${visible_devices}" + else + visible_count=1 + if [[ -n "${ASCEND_RT_VISIBLE_DEVICES}" ]]; then + visible_count="$(awk -F, '{print NF}' <<< "${ASCEND_RT_VISIBLE_DEVICES}")" + fi + if (( accelerator_count > visible_count )); then + echo "SPECO_ACCELERATOR_COUNT=${accelerator_count} exceeds visible NPU count ${visible_count} from ASCEND_RT_VISIBLE_DEVICES=${ASCEND_RT_VISIBLE_DEVICES}" >&2 + exit 2 + fi fi export HCCL_HOST_SOCKET_PORT_RANGE="${HCCL_HOST_SOCKET_PORT_RANGE:-60000-60050}" export HCCL_NPU_SOCKET_PORT_RANGE="${HCCL_NPU_SOCKET_PORT_RANGE:-61000-61050}" @@ -139,6 +174,10 @@ overrides=( "trainer.save_freq=${SPECO_SAVE_FREQ:--1}" "trainer.test_freq=${SPECO_TEST_FREQ:--1}" "trainer.total_epochs=${total_epochs}" + "trainer.total_training_steps=${SPECO_TOTAL_TRAINING_STEPS:-2}" + "data.train_max_samples=${SPECO_TRAIN_MAX_SAMPLES:-1}" + "data.val_max_samples=${SPECO_VAL_MAX_SAMPLES:-1}" + "data.dataloader_num_workers=${SPECO_DATALOADER_NUM_WORKERS:-0}" ) if [[ "${drafter}" == "dflash" ]]; then @@ -154,6 +193,18 @@ if [[ "${drafter}" == "dflash" ]]; then ) fi +if [[ "${drafter}" == "dspark" ]]; then + overrides+=( + "actor_rollout_ref.rollout.drafter.rollout.spec_steps=${SPECO_DSPARK_SPEC_STEPS:-1}" + "actor_rollout_ref.rollout.drafter.rollout.spec_verify_tokens=${SPECO_DSPARK_SPEC_VERIFY_TOKENS:-7}" + "actor_rollout_ref.rollout.drafter.training.hidden_state_window_min_rows=${SPECO_HIDDEN_STATE_WINDOW_MIN_ROWS:-1}" + "actor_rollout_ref.rollout.drafter.training.hidden_state_window_tokens_per_sample=${SPECO_HIDDEN_STATE_WINDOW_TOKENS_PER_SAMPLE:-64}" + "actor_rollout_ref.rollout.drafter.training.dspark_block_size=${SPECO_DSPARK_BLOCK_SIZE:-7}" + "actor_rollout_ref.rollout.drafter.training.dspark_num_anchors=${SPECO_DSPARK_NUM_ANCHORS:-8}" + "actor_rollout_ref.rollout.drafter.training.dspark_max_window=${SPECO_DSPARK_MAX_WINDOW:-64}" + ) +fi + if [[ -n "${SPECO_EXTRA_HYDRA_ARGS:-}" ]]; then while IFS= read -r extra_arg; do [[ -z "${extra_arg}" ]] && continue @@ -161,4 +212,16 @@ if [[ -n "${SPECO_EXTRA_HYDRA_ARGS:-}" ]]; then done <<< "${SPECO_EXTRA_HYDRA_ARGS}" fi +if [[ "${SPECO_DRY_RUN:-false}" == "true" ]]; then + echo "platform=${platform}" + echo "backend=${backend}" + echo "drafter=${drafter}" + echo "example=${example}" + echo "draft_algorithm=${draft_algorithm}" + echo "ASCEND_RT_VISIBLE_DEVICES=${ASCEND_RT_VISIBLE_DEVICES:-}" + printf 'Hydra overrides:\n' + printf ' %q\n' "${overrides[@]}" + exit 0 +fi + bash "${example}" "${overrides[@]}" diff --git a/tests/examples/test_ci_example_runner.py b/tests/examples/test_ci_example_runner.py index e613c39f..70d78efb 100644 --- a/tests/examples/test_ci_example_runner.py +++ b/tests/examples/test_ci_example_runner.py @@ -1,5 +1,7 @@ from __future__ import annotations +import os +import shlex import subprocess import shutil from pathlib import Path @@ -33,6 +35,20 @@ def _require_working_bash() -> str: return bash +def _bash_path(path: Path, bash: str) -> str: + if os.name != "nt": + return str(path) + if "system32" in bash.lower(): + drive = path.drive.rstrip(":").lower() + rest = path.relative_to(path.anchor).as_posix() + return f"/mnt/{drive}/{rest}" + return path.as_posix() + + +def _runner_script() -> str: + return "\n".join(RUNNER.read_text(encoding="utf-8").splitlines()) + "\n" + + def test_ci_layers_match_required_shape() -> None: expected = { "cpu_unit_tests.yml", @@ -43,7 +59,7 @@ def test_ci_layers_match_required_shape() -> None: assert expected <= {path.name for path in WORKFLOWS.glob("*.yml")} assert "pull_request" in _workflow("cpu_unit_tests.yml")["on"] assert "pull_request" not in _workflow("gpu_example_tests.yml")["on"] - assert "pull_request" not in _workflow("npu_example_tests.yml")["on"] + assert "pull_request" in _workflow("npu_example_tests.yml")["on"] def test_cpu_unit_workflow_is_lightweight_pr_gate() -> None: @@ -78,28 +94,41 @@ def test_gpu_and_npu_workflows_run_examples_on_self_hosted_runners() -> None: assert "SPECO_TARGET_MODEL" in source assert "SPECO_EAGLE3_DRAFT_MODEL" in source assert "SPECO_DFLASH_DRAFT_MODEL" in source + if workflow_name == "npu_example_tests.yml": + assert "SPECO_DSPARK_DRAFT_MODEL" in source assert "SPECO_ACCELERATOR_COUNT" in source assert "SPECO_TENSOR_PARALLEL_SIZE" in source assert "SPECO_SEQUENCE_PARALLEL_SIZE" in source assert "SPECO_ENABLE_TRAINING" in source assert "SPECO_EXTRA_HYDRA_ARGS" in source - matrix_entries = { - (entry["backend"], entry["drafter"]) - for entry in workflow["jobs"]["example"]["strategy"]["matrix"]["include"] - } - assert { - ("vllm", "eagle3"), - ("vllm", "dflash"), - ("sglang", "eagle3"), - ("sglang", "dflash"), - } <= matrix_entries + if workflow_name == "npu_example_tests.yml": + assert "github.event.pull_request.head.repo.full_name == github.repository" in source + assert "linux-aarch64-a2-8" in source + assert "linux-aarch64-a2-4" not in source + assert '"backend":"vllm","drafter":"dspark"' in source + assert '"backend":"sglang","drafter":"eagle3"' in source + assert '"backend":"sglang","drafter":"dflash"' in source + else: + matrix_entries = { + (entry["backend"], entry["drafter"]) + for entry in workflow["jobs"]["example"]["strategy"]["matrix"]["include"] + } + assert { + ("vllm", "eagle3"), + ("vllm", "dflash"), + ("sglang", "eagle3"), + ("sglang", "dflash"), + } <= matrix_entries for job in workflow["jobs"].values(): - assert {"self-hosted", label} <= set(job["runs-on"]) + if workflow_name == "npu_example_tests.yml": + assert set(job["runs-on"]) == {"self-hosted", "linux-aarch64-a2-8"} + else: + assert {"self-hosted", label} <= set(job["runs-on"]) def test_example_runner_shell_syntax_is_valid() -> None: bash = _require_working_bash() - subprocess.run([bash, "-n", str(RUNNER)], check=True) + subprocess.run([bash, "-n", "-s"], input=_runner_script().encode("utf-8"), check=True) def test_example_runner_covers_gpu_and_npu_backend_matrix() -> None: @@ -107,10 +136,12 @@ def test_example_runner_covers_gpu_and_npu_backend_matrix() -> None: assert "gpu/vllm/eagle3" in source assert "gpu/vllm/dflash" in source + assert "gpu/vllm/dspark" in source assert "gpu/sglang/eagle3" in source assert "gpu/sglang/dflash" in source assert "npu/vllm/eagle3" in source assert "npu/vllm/dflash" in source + assert "npu/vllm/dspark" in source assert "npu/sglang/eagle3" in source assert "npu/sglang/dflash" in source assert "examples/run_qwen3-8b_drafter_eagle3_vllm.sh" in source @@ -120,6 +151,8 @@ def test_example_runner_covers_gpu_and_npu_backend_matrix() -> None: assert "examples/run_qwen3-8b_drafter_dflash_vllm.sh" in source assert "examples/run_qwen3-8b_drafter_dflash_vllm_npu.sh" in source assert "examples/run_qwen3-8b_drafter_dflash_sglang.sh" in source + assert "examples/run_qwen3-8b_drafter_dspark_vllm.sh" in source + assert "examples/run_qwen3-8b_drafter_dspark_vllm_npu.sh" in source def test_example_runner_exposes_required_hydra_overrides() -> None: @@ -137,4 +170,40 @@ def test_example_runner_exposes_required_hydra_overrides() -> None: assert "SPECO_SPEC_VERIFY_TOKENS" in source assert "SPECO_DFLASH_NUM_ANCHORS" in source assert "SPECO_DFLASH_MAX_WINDOW" in source + assert "SPECO_DSPARK_DRAFT_MODEL" in source + assert "SPECO_DSPARK_NUM_ANCHORS" in source + assert "SPECO_DSPARK_MAX_WINDOW" in source + assert "SPECO_TOTAL_TRAINING_STEPS" in source + assert "SPECO_TRAIN_MAX_SAMPLES" in source + assert "SPECO_VAL_MAX_SAMPLES" in source + assert "SPECO_DATALOADER_NUM_WORKERS" in source assert "SPECO_EXTRA_HYDRA_ARGS" in source + + +def test_example_runner_dry_run_covers_npu_dspark() -> None: + bash = _require_working_bash() + env = { + "SPECO_DRY_RUN": "true", + "SPECO_TARGET_MODEL": "/models/target", + "SPECO_DSPARK_DRAFT_MODEL": "/models/dspark", + "SPECO_TRAIN_FILE": "/data/train.parquet", + "SPECO_TEST_FILE": "/data/test.parquet", + "SPECO_CKPT_DIR": "/tmp/speco", + "SPECO_ACCELERATOR_COUNT": "1", + } + script = "".join(f"export {name}={shlex.quote(value)}\n" for name, value in env.items()) + script += _runner_script() + result = subprocess.run( + [bash, "-s", "--", "npu", "vllm", "dspark"], + env=os.environ.copy(), + input=script.encode("utf-8"), + capture_output=True, + check=True, + ) + stdout = result.stdout.decode("utf-8", errors="replace") + + assert "example=examples/run_qwen3-8b_drafter_dspark_vllm_npu.sh" in stdout + assert "draft_algorithm=DSPARK" in stdout + assert "actor_rollout_ref.rollout.drafter.speculative_algorithm=DSPARK" in stdout + assert "actor_rollout_ref.rollout.drafter.training.dspark_block_size=7" in stdout + assert "actor_rollout_ref.rollout.drafter.rollout.spec_verify_tokens=7" in stdout From 6f783adfa469d704de898959462ec5907a00c153 Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Thu, 16 Jul 2026 19:06:59 +0800 Subject: [PATCH 15/50] Address DSpark NPU CI review feedback --- ci/README.md | 2 ++ ci/run_example_test.sh | 3 ++- verl_speco/trainer/draft_training_loop.py | 8 +------- 3 files changed, 5 insertions(+), 8 deletions(-) diff --git a/ci/README.md b/ci/README.md index c1767059..fad0ef02 100644 --- a/ci/README.md +++ b/ci/README.md @@ -68,6 +68,8 @@ GitHub environments, or pass them as manual workflow inputs where available: - `SPECO_DFLASH_NUM_ANCHORS` - `SPECO_DFLASH_MAX_WINDOW` - `SPECO_DSPARK_BLOCK_SIZE` +- `SPECO_DSPARK_SPEC_STEPS` +- `SPECO_DSPARK_SPEC_VERIFY_TOKENS` - `SPECO_DSPARK_NUM_ANCHORS` - `SPECO_DSPARK_MAX_WINDOW` - `SPECO_TOTAL_TRAINING_STEPS` diff --git a/ci/run_example_test.sh b/ci/run_example_test.sh index 18104cd0..3edf8b69 100644 --- a/ci/run_example_test.sh +++ b/ci/run_example_test.sh @@ -107,7 +107,8 @@ PY else visible_count=1 if [[ -n "${ASCEND_RT_VISIBLE_DEVICES}" ]]; then - visible_count="$(awk -F, '{print NF}' <<< "${ASCEND_RT_VISIBLE_DEVICES}")" + visible_commas="${ASCEND_RT_VISIBLE_DEVICES//[^,]/}" + visible_count=$(( ${#visible_commas} + 1 )) fi if (( accelerator_count > visible_count )); then echo "SPECO_ACCELERATOR_COUNT=${accelerator_count} exceeds visible NPU count ${visible_count} from ASCEND_RT_VISIBLE_DEVICES=${ASCEND_RT_VISIBLE_DEVICES}" >&2 diff --git a/verl_speco/trainer/draft_training_loop.py b/verl_speco/trainer/draft_training_loop.py index 64226e16..26314a19 100644 --- a/verl_speco/trainer/draft_training_loop.py +++ b/verl_speco/trainer/draft_training_loop.py @@ -353,13 +353,7 @@ def _barrier() -> None: def _sync_any_rank_saved_checkpoint(saved: Any) -> bool: if not dist.is_initialized(): return bool(saved) - device_name = get_device_name() - if device_name == "cpu": - device = torch.device("cpu") - else: - current_device = getattr(get_torch_device(), "current_device", None) - device_index = current_device() if callable(current_device) else 0 - device = torch.device(f"{device_name}:{int(device_index)}") + device = torch.device(get_device_name()) flag = torch.tensor([1 if saved else 0], dtype=torch.int32, device=device) dist.all_reduce(flag, op=dist.ReduceOp.MAX) return bool(flag.item()) From b664a91e6b9fb2de424fa78436d0de7a052cd694 Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Thu, 16 Jul 2026 19:28:05 +0800 Subject: [PATCH 16/50] Align NPU examples with 8-card CI runners --- examples/run_qwen3-8b_drafter_dflash_vllm_npu.sh | 5 +++-- examples/run_qwen3-8b_drafter_dspark_vllm_npu.sh | 5 +++-- examples/run_qwen3-8b_drafter_eagle3_sglang_npu.sh | 5 +++-- examples/run_qwen3-8b_drafter_eagle3_vllm_npu.sh | 5 +++-- 4 files changed, 12 insertions(+), 8 deletions(-) diff --git a/examples/run_qwen3-8b_drafter_dflash_vllm_npu.sh b/examples/run_qwen3-8b_drafter_dflash_vllm_npu.sh index 8df0abc8..d1b2bbcc 100644 --- a/examples/run_qwen3-8b_drafter_dflash_vllm_npu.sh +++ b/examples/run_qwen3-8b_drafter_dflash_vllm_npu.sh @@ -1,5 +1,5 @@ set -x -export ASCEND_RT_VISIBLE_DEVICES=0,1,2,3,4,5,6,7,8,9,10,11,12,13,14,15 +export ASCEND_RT_VISIBLE_DEVICES="${ASCEND_RT_VISIBLE_DEVICES:-0,1,2,3,4,5,6,7}" # NPU example for the native DFlash proposer in vLLM. project_name='verl_grpo_example_dflash_drafter' @@ -7,6 +7,7 @@ exp_name='qwen3_8b_dflash_drafter_vllm_npu' gen_tp=2 train_sp=4 +ppo_gpus_per_node=${SPECO_ACCELERATOR_COUNT:-8} MODEL_PATH=/path/to/model CKPTS_DIR=/path/to/checkpoint @@ -95,7 +96,7 @@ PYTHONUNBUFFERED=1 python3 -m verl_speco.main \ trainer.logger='["console"]' \ trainer.project_name=${project_name} \ trainer.experiment_name=${exp_name} \ - trainer.n_gpus_per_node=16 \ + trainer.n_gpus_per_node=${ppo_gpus_per_node} \ trainer.nnodes=1 \ trainer.resume_mode=disable \ trainer.default_local_dir=${CKPTS_DIR} \ diff --git a/examples/run_qwen3-8b_drafter_dspark_vllm_npu.sh b/examples/run_qwen3-8b_drafter_dspark_vllm_npu.sh index 267d465d..344eeabe 100644 --- a/examples/run_qwen3-8b_drafter_dspark_vllm_npu.sh +++ b/examples/run_qwen3-8b_drafter_dspark_vllm_npu.sh @@ -1,5 +1,5 @@ set -x -export ASCEND_RT_VISIBLE_DEVICES=0,1,2,3,4,5,6,7,8,9,10,11,12,13,14,15 +export ASCEND_RT_VISIBLE_DEVICES="${ASCEND_RT_VISIBLE_DEVICES:-0,1,2,3,4,5,6,7}" # NPU example for DSpark on vLLM-Ascend. SPECO keeps the user-facing # algorithm as DSPARK and maps it to vLLM's dflash speculative method. @@ -8,6 +8,7 @@ exp_name='qwen3_8b_dspark_drafter_vllm_npu' gen_tp=2 train_sp=4 +ppo_gpus_per_node=${SPECO_ACCELERATOR_COUNT:-8} MODEL_PATH=/path/to/model CKPTS_DIR=/path/to/checkpoint @@ -104,7 +105,7 @@ PYTHONUNBUFFERED=1 python3 -m verl_speco.main \ trainer.logger='["console"]' \ trainer.project_name=${project_name} \ trainer.experiment_name=${exp_name} \ - trainer.n_gpus_per_node=16 \ + trainer.n_gpus_per_node=${ppo_gpus_per_node} \ trainer.nnodes=1 \ trainer.resume_mode=disable \ trainer.default_local_dir=${CKPTS_DIR} \ diff --git a/examples/run_qwen3-8b_drafter_eagle3_sglang_npu.sh b/examples/run_qwen3-8b_drafter_eagle3_sglang_npu.sh index 15524c8a..b29e05ea 100644 --- a/examples/run_qwen3-8b_drafter_eagle3_sglang_npu.sh +++ b/examples/run_qwen3-8b_drafter_eagle3_sglang_npu.sh @@ -2,7 +2,7 @@ set -x export HCCL_HOST_SOCKET_PORT_RANGE=60000-60050 export HCCL_NPU_SOCKET_PORT_RANGE=61000-61050 export RAY_EXPERIMENTAL_NOSET_ASCEND_RT_VISIBLE_DEVICES=1 -export ASCEND_RT_VISIBLE_DEVICES=0,1,2,3,4,5,6,7,8,9,10,11,12,13,14,15 +export ASCEND_RT_VISIBLE_DEVICES="${ASCEND_RT_VISIBLE_DEVICES:-0,1,2,3,4,5,6,7}" export SGLANG_DEEPEP_BF16_DISPATCH=1 export SGLANG_SET_CPU_AFFINITY=1 @@ -16,6 +16,7 @@ exp_name='qwen3_8b_function_rm_drafter' gen_tp=2 train_sp=4 +ppo_gpus_per_node=${SPECO_ACCELERATOR_COUNT:-8} MODEL_PATH=/path/to/model CKPTS_DIR=/path/to/checkpoint @@ -85,7 +86,7 @@ PYTHONUNBUFFERED=1 python3 -m verl_speco.main \ trainer.logger='["console"]' \ trainer.project_name=project_name \ trainer.experiment_name=exp_name \ - trainer.n_gpus_per_node=16 \ + trainer.n_gpus_per_node=${ppo_gpus_per_node} \ trainer.nnodes=1 \ trainer.default_local_dir=${CKPTS_DIR} \ trainer.save_freq=20 \ diff --git a/examples/run_qwen3-8b_drafter_eagle3_vllm_npu.sh b/examples/run_qwen3-8b_drafter_eagle3_vllm_npu.sh index 5fe6af55..d0a63f13 100644 --- a/examples/run_qwen3-8b_drafter_eagle3_vllm_npu.sh +++ b/examples/run_qwen3-8b_drafter_eagle3_vllm_npu.sh @@ -2,7 +2,7 @@ set -x export HCCL_HOST_SOCKET_PORT_RANGE=60000-60050 export HCCL_NPU_SOCKET_PORT_RANGE=61000-61050 export RAY_EXPERIMENTAL_NOSET_ASCEND_RT_VISIBLE_DEVICES=1 -export ASCEND_RT_VISIBLE_DEVICES=0,1,2,3,4,5,6,7,8,9,10,11,12,13,14,15 +export ASCEND_RT_VISIBLE_DEVICES="${ASCEND_RT_VISIBLE_DEVICES:-0,1,2,3,4,5,6,7}" export PYTORCH_NPU_ALLOC_CONF=expandable_segments:True export STREAMS_PER_DEVICE=32 @@ -14,6 +14,7 @@ exp_name='qwen3_8b_function_rm_drafter_vllm_npu' gen_tp=2 train_sp=4 +ppo_gpus_per_node=${SPECO_ACCELERATOR_COUNT:-8} MODEL_PATH=/path/to/model CKPTS_DIR=/path/to/checkpoint @@ -89,7 +90,7 @@ PYTHONUNBUFFERED=1 python3 -m verl_speco.main \ trainer.logger='["console"]' \ trainer.project_name=${project_name} \ trainer.experiment_name=${exp_name} \ - trainer.n_gpus_per_node=16 \ + trainer.n_gpus_per_node=${ppo_gpus_per_node} \ trainer.nnodes=1 \ trainer.default_local_dir=${CKPTS_DIR} \ trainer.save_freq=20 \ From eef7dff31c63e412b7c9d2374ce6346212fe8b70 Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Mon, 20 Jul 2026 14:47:06 +0800 Subject: [PATCH 17/50] Fix standalone draft checkpoint lm_head export --- tests/unit/test_draft_training_loop.py | 36 ++++++- verl_speco/trainer/draft_training_loop.py | 113 +++++++++++++++++++++- 2 files changed, 145 insertions(+), 4 deletions(-) diff --git a/tests/unit/test_draft_training_loop.py b/tests/unit/test_draft_training_loop.py index 6b921066..7a36be45 100644 --- a/tests/unit/test_draft_training_loop.py +++ b/tests/unit/test_draft_training_loop.py @@ -6,9 +6,12 @@ import pytest -pytest.importorskip("torch") +torch = pytest.importorskip("torch") -from verl_speco.trainer.draft_training_loop import _rewrite_standalone_block_runtime_config, _save_standalone_checkpoint +from verl_speco.trainer.draft_training_loop import ( # noqa: E402 + _rewrite_standalone_block_runtime_config, + _save_standalone_checkpoint, +) class _FakeTrainer: @@ -176,3 +179,32 @@ def test_standalone_dflash_checkpoint_preserves_source_runtime_config(tmp_path): assert runtime_config["dflash_config"]["target_layer_ids"] == [2, 10, 18] assert runtime_config["eagle_aux_hidden_state_layer_ids"] == [3, 11, 19] assert saved_training_config == training_config + + +def test_standalone_block_checkpoint_appends_source_lm_head_weight(tmp_path): + safetensors_torch = pytest.importorskip("safetensors.torch") + checkpoint_dir = tmp_path / "draft_step_5" + checkpoint_dir.mkdir() + source_dir = tmp_path / "source_dspark" + source_dir.mkdir() + (source_dir / "config.json").write_text( + json.dumps({"model_type": "qwen3", "architectures": ["DSparkForCausalLM"]}), + encoding="utf-8", + ) + (checkpoint_dir / "config.json").write_text( + json.dumps({"model_type": "dspark", "architectures": ["DSparkDraftModel"]}), + encoding="utf-8", + ) + lm_head = torch.arange(12, dtype=torch.float32).reshape(3, 4) + safetensors_torch.save_file({"lm_head.weight": lm_head}, str(source_dir / "model.safetensors")) + safetensors_torch.save_file({"fc.weight": torch.ones(2, 2)}, str(checkpoint_dir / "model.safetensors")) + trainer = SimpleNamespace( + backend=SimpleNamespace(model_type="dspark"), + config=SimpleNamespace(rollout=SimpleNamespace(drafter=SimpleNamespace(model_path=str(source_dir)))), + ) + + _rewrite_standalone_block_runtime_config(trainer, str(checkpoint_dir)) + + exported_state = safetensors_torch.load_file(str(checkpoint_dir / "model.safetensors"), device="cpu") + assert torch.equal(exported_state["lm_head.weight"], lm_head) + assert torch.equal(exported_state["fc.weight"], torch.ones(2, 2)) diff --git a/verl_speco/trainer/draft_training_loop.py b/verl_speco/trainer/draft_training_loop.py index 3cd85391..f94b4024 100644 --- a/verl_speco/trainer/draft_training_loop.py +++ b/verl_speco/trainer/draft_training_loop.py @@ -211,12 +211,19 @@ def _ensure_dict_child(config: dict[str, Any], key: str) -> dict[str, Any]: return value -def _load_source_drafter_config(trainer: DrafterBaseTrainer) -> dict[str, Any] | None: +def _source_drafter_model_path(trainer: DrafterBaseTrainer) -> str | None: model_path = getattr(getattr(getattr(trainer, "config", None), "rollout", None), "drafter", None) model_path = getattr(model_path, "model_path", None) if not model_path: return None - config_path = os.path.join(os.fspath(model_path), "config.json") + return os.fspath(model_path) + + +def _load_source_drafter_config(trainer: DrafterBaseTrainer) -> dict[str, Any] | None: + model_path = _source_drafter_model_path(trainer) + if not model_path: + return None + config_path = os.path.join(model_path, "config.json") if not os.path.exists(config_path): return None try: @@ -228,6 +235,105 @@ def _load_source_drafter_config(trainer: DrafterBaseTrainer) -> dict[str, Any] | return loaded if isinstance(loaded, dict) else None +def _load_tensor_from_safetensors(path: str, key: str) -> torch.Tensor | None: + try: + from safetensors import safe_open + + with safe_open(path, framework="pt", device="cpu") as f: + if key in f.keys(): + return f.get_tensor(key) + except Exception as exc: # noqa: BLE001 + logger.warning("Failed to load %s from %s: %s", key, path, exc) + return None + + +def _load_tensor_from_torch(path: str, key: str) -> torch.Tensor | None: + try: + state = torch.load(path, map_location="cpu", weights_only=True) + except Exception as exc: # noqa: BLE001 + logger.warning("Failed to load %s from %s: %s", key, path, exc) + return None + if isinstance(state, dict): + value = state.get(key) + if isinstance(value, torch.Tensor): + return value + return None + + +def _load_source_lm_head_weight(model_path: str | None) -> torch.Tensor | None: + if not model_path: + return None + + for index_name in ("model.safetensors.index.json", "pytorch_model.bin.index.json"): + index_path = os.path.join(model_path, index_name) + if not os.path.exists(index_path): + continue + try: + with open(index_path, "r", encoding="utf-8") as f: + weight_map = json.load(f).get("weight_map", {}) + except (OSError, json.JSONDecodeError) as exc: + logger.warning("Failed to read source weight index %s: %s", index_path, exc) + continue + shard_name = weight_map.get("lm_head.weight") if isinstance(weight_map, dict) else None + if not shard_name: + continue + shard_path = os.path.join(model_path, os.fspath(shard_name)) + if index_name.endswith(".safetensors.index.json"): + return _load_tensor_from_safetensors(shard_path, "lm_head.weight") + return _load_tensor_from_torch(shard_path, "lm_head.weight") + + safetensors_path = os.path.join(model_path, "model.safetensors") + if os.path.exists(safetensors_path): + return _load_tensor_from_safetensors(safetensors_path, "lm_head.weight") + + torch_path = os.path.join(model_path, "pytorch_model.bin") + if os.path.exists(torch_path): + return _load_tensor_from_torch(torch_path, "lm_head.weight") + + return None + + +def _append_lm_head_weight_if_missing(checkpoint_path: str, source_model_path: str | None) -> None: + lm_head_weight = _load_source_lm_head_weight(source_model_path) + if lm_head_weight is None: + logger.warning("Standalone checkpoint export could not find source lm_head.weight in %s", source_model_path) + return + lm_head_weight = lm_head_weight.detach().cpu() + + safetensors_path = os.path.join(checkpoint_path, "model.safetensors") + if os.path.exists(safetensors_path): + try: + from safetensors.torch import load_file, save_file + + state = load_file(safetensors_path, device="cpu") + if "lm_head.weight" in state: + return + state["lm_head.weight"] = lm_head_weight + save_file(state, safetensors_path) + logger.info("Added lm_head.weight to standalone checkpoint %s", safetensors_path) + except Exception as exc: # noqa: BLE001 + logger.warning("Failed to append lm_head.weight to %s: %s", safetensors_path, exc) + return + + torch_path = os.path.join(checkpoint_path, "pytorch_model.bin") + if os.path.exists(torch_path): + try: + state = torch.load(torch_path, map_location="cpu", weights_only=True) + if not isinstance(state, dict): + logger.warning("Cannot append lm_head.weight to %s: expected dict state", torch_path) + return + if "lm_head.weight" in state: + return + state["lm_head.weight"] = lm_head_weight + torch.save(state, torch_path) + logger.info("Added lm_head.weight to standalone checkpoint %s", torch_path) + except Exception as exc: # noqa: BLE001 + logger.warning("Failed to append lm_head.weight to %s: %s", torch_path, exc) + return + + logger.warning("Standalone checkpoint export found no model.safetensors or pytorch_model.bin under %s", checkpoint_path) + + def _fill_if_missing(dst: dict[str, Any], src: dict[str, Any], keys: tuple[str, ...]) -> None: for key in keys: if key in src and key not in dst: @@ -278,6 +384,7 @@ def _rewrite_standalone_block_runtime_config( except OSError as exc: logger.warning("Failed to write standalone training config copy %s: %s", training_config_path, exc) + source_model_path = _source_drafter_model_path(trainer) runtime_config = _load_source_drafter_config(trainer) if runtime_config is None: runtime_config = deepcopy(training_config) @@ -335,6 +442,8 @@ def _rewrite_standalone_block_runtime_config( except OSError as exc: logger.warning("Failed to write standalone runtime config %s: %s", config_path, exc) + _append_lm_head_weight_if_missing(checkpoint_path, source_model_path) + def _disable_standalone_sequence_parallel(draft_config) -> None: rollout_cfg = draft_config.rollout From 5b768545c2d6c79956ecfe7e88345fb194926b34 Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Tue, 21 Jul 2026 18:46:54 +0800 Subject: [PATCH 18/50] Fix: Support sharded standalone lm_head export --- tests/unit/test_draft_training_loop.py | 65 +++++++++++++++++++ verl_speco/trainer/draft_training_loop.py | 76 ++++++++++++++++++++++- 2 files changed, 139 insertions(+), 2 deletions(-) diff --git a/tests/unit/test_draft_training_loop.py b/tests/unit/test_draft_training_loop.py index 7a36be45..fbd1e92f 100644 --- a/tests/unit/test_draft_training_loop.py +++ b/tests/unit/test_draft_training_loop.py @@ -11,6 +11,7 @@ from verl_speco.trainer.draft_training_loop import ( # noqa: E402 _rewrite_standalone_block_runtime_config, _save_standalone_checkpoint, + _torch_load_cpu, ) @@ -208,3 +209,67 @@ def test_standalone_block_checkpoint_appends_source_lm_head_weight(tmp_path): exported_state = safetensors_torch.load_file(str(checkpoint_dir / "model.safetensors"), device="cpu") assert torch.equal(exported_state["lm_head.weight"], lm_head) assert torch.equal(exported_state["fc.weight"], torch.ones(2, 2)) + + +def test_standalone_block_checkpoint_appends_lm_head_to_sharded_safetensors_index(tmp_path): + safetensors_torch = pytest.importorskip("safetensors.torch") + checkpoint_dir = tmp_path / "draft_step_5" + checkpoint_dir.mkdir() + source_dir = tmp_path / "source_dspark" + source_dir.mkdir() + (source_dir / "config.json").write_text( + json.dumps({"model_type": "qwen3", "architectures": ["DSparkForCausalLM"]}), + encoding="utf-8", + ) + (checkpoint_dir / "config.json").write_text( + json.dumps({"model_type": "dspark", "architectures": ["DSparkDraftModel"]}), + encoding="utf-8", + ) + lm_head = torch.arange(12, dtype=torch.float32).reshape(3, 4) + fc_weight = torch.ones(2, 2) + safetensors_torch.save_file({"lm_head.weight": lm_head}, str(source_dir / "model.safetensors")) + safetensors_torch.save_file({"fc.weight": fc_weight}, str(checkpoint_dir / "model-00001-of-00001.safetensors")) + (checkpoint_dir / "model.safetensors.index.json").write_text( + json.dumps( + { + "metadata": {"total_size": fc_weight.numel() * fc_weight.element_size()}, + "weight_map": {"fc.weight": "model-00001-of-00001.safetensors"}, + } + ), + encoding="utf-8", + ) + trainer = SimpleNamespace( + backend=SimpleNamespace(model_type="dspark"), + config=SimpleNamespace(rollout=SimpleNamespace(drafter=SimpleNamespace(model_path=str(source_dir)))), + ) + + _rewrite_standalone_block_runtime_config(trainer, str(checkpoint_dir)) + + index_data = json.loads((checkpoint_dir / "model.safetensors.index.json").read_text(encoding="utf-8")) + assert index_data["weight_map"]["lm_head.weight"] == "model-lm-head.safetensors" + assert index_data["metadata"]["total_size"] == ( + fc_weight.numel() * fc_weight.element_size() + lm_head.numel() * lm_head.element_size() + ) + added_state = safetensors_torch.load_file(str(checkpoint_dir / "model-lm-head.safetensors"), device="cpu") + assert torch.equal(added_state["lm_head.weight"], lm_head) + + +def test_torch_load_cpu_falls_back_without_weights_only(monkeypatch, tmp_path): + checkpoint_path = tmp_path / "pytorch_model.bin" + expected = {"lm_head.weight": torch.ones(2, 2)} + calls = [] + + def fake_load(path, **kwargs): + calls.append(kwargs) + if "weights_only" in kwargs: + raise TypeError("weights_only is unsupported") + assert path == str(checkpoint_path) + return expected + + monkeypatch.setattr(torch, "load", fake_load) + + assert _torch_load_cpu(str(checkpoint_path)) is expected + assert calls == [ + {"map_location": "cpu", "weights_only": True}, + {"map_location": "cpu"}, + ] diff --git a/verl_speco/trainer/draft_training_loop.py b/verl_speco/trainer/draft_training_loop.py index f94b4024..c9de8fe7 100644 --- a/verl_speco/trainer/draft_training_loop.py +++ b/verl_speco/trainer/draft_training_loop.py @@ -247,9 +247,16 @@ def _load_tensor_from_safetensors(path: str, key: str) -> torch.Tensor | None: return None +def _torch_load_cpu(path: str) -> Any: + try: + return torch.load(path, map_location="cpu", weights_only=True) + except TypeError: + return torch.load(path, map_location="cpu") + + def _load_tensor_from_torch(path: str, key: str) -> torch.Tensor | None: try: - state = torch.load(path, map_location="cpu", weights_only=True) + state = _torch_load_cpu(path) except Exception as exc: # noqa: BLE001 logger.warning("Failed to load %s from %s: %s", key, path, exc) return None @@ -293,6 +300,66 @@ def _load_source_lm_head_weight(model_path: str | None) -> torch.Tensor | None: return None +def _append_lm_head_to_safetensors_index(checkpoint_path: str, lm_head_weight: torch.Tensor) -> bool: + index_path = os.path.join(checkpoint_path, "model.safetensors.index.json") + if not os.path.exists(index_path): + return False + try: + from safetensors.torch import save_file + + with open(index_path, "r", encoding="utf-8") as f: + index_data = json.load(f) + weight_map = index_data.setdefault("weight_map", {}) + if not isinstance(weight_map, dict): + logger.warning("Cannot append lm_head.weight to %s: expected weight_map object", index_path) + return True + if "lm_head.weight" in weight_map: + return True + + shard_name = "model-lm-head.safetensors" + save_file({"lm_head.weight": lm_head_weight}, os.path.join(checkpoint_path, shard_name)) + weight_map["lm_head.weight"] = shard_name + metadata = index_data.setdefault("metadata", {}) + if isinstance(metadata, dict) and "total_size" in metadata: + metadata["total_size"] = int(metadata["total_size"]) + lm_head_weight.numel() * lm_head_weight.element_size() + with open(index_path, "w", encoding="utf-8") as f: + json.dump(index_data, f, indent=2, sort_keys=True) + f.write("\n") + logger.info("Added lm_head.weight to standalone sharded checkpoint %s", index_path) + except Exception as exc: # noqa: BLE001 + logger.warning("Failed to append lm_head.weight to sharded checkpoint %s: %s", index_path, exc) + return True + + +def _append_lm_head_to_torch_index(checkpoint_path: str, lm_head_weight: torch.Tensor) -> bool: + index_path = os.path.join(checkpoint_path, "pytorch_model.bin.index.json") + if not os.path.exists(index_path): + return False + try: + with open(index_path, "r", encoding="utf-8") as f: + index_data = json.load(f) + weight_map = index_data.setdefault("weight_map", {}) + if not isinstance(weight_map, dict): + logger.warning("Cannot append lm_head.weight to %s: expected weight_map object", index_path) + return True + if "lm_head.weight" in weight_map: + return True + + shard_name = "pytorch_model-lm-head.bin" + torch.save({"lm_head.weight": lm_head_weight}, os.path.join(checkpoint_path, shard_name)) + weight_map["lm_head.weight"] = shard_name + metadata = index_data.setdefault("metadata", {}) + if isinstance(metadata, dict) and "total_size" in metadata: + metadata["total_size"] = int(metadata["total_size"]) + lm_head_weight.numel() * lm_head_weight.element_size() + with open(index_path, "w", encoding="utf-8") as f: + json.dump(index_data, f, indent=2, sort_keys=True) + f.write("\n") + logger.info("Added lm_head.weight to standalone sharded checkpoint %s", index_path) + except Exception as exc: # noqa: BLE001 + logger.warning("Failed to append lm_head.weight to sharded checkpoint %s: %s", index_path, exc) + return True + + def _append_lm_head_weight_if_missing(checkpoint_path: str, source_model_path: str | None) -> None: lm_head_weight = _load_source_lm_head_weight(source_model_path) if lm_head_weight is None: @@ -300,6 +367,11 @@ def _append_lm_head_weight_if_missing(checkpoint_path: str, source_model_path: s return lm_head_weight = lm_head_weight.detach().cpu() + if _append_lm_head_to_safetensors_index(checkpoint_path, lm_head_weight): + return + if _append_lm_head_to_torch_index(checkpoint_path, lm_head_weight): + return + safetensors_path = os.path.join(checkpoint_path, "model.safetensors") if os.path.exists(safetensors_path): try: @@ -318,7 +390,7 @@ def _append_lm_head_weight_if_missing(checkpoint_path: str, source_model_path: s torch_path = os.path.join(checkpoint_path, "pytorch_model.bin") if os.path.exists(torch_path): try: - state = torch.load(torch_path, map_location="cpu", weights_only=True) + state = _torch_load_cpu(torch_path) if not isinstance(state, dict): logger.warning("Cannot append lm_head.weight to %s: expected dict state", torch_path) return From 13f79c6d1a0179d1417f670aab88d9496baf18bf Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Fri, 24 Jul 2026 14:42:52 +0800 Subject: [PATCH 19/50] Add standalone draft training workflow and checkpoint fixes --- tests/unit/test_draft_training_loop.py | 227 ++++++++++ verl_speco/backends/dflash_trainer_backend.py | 14 +- verl_speco/backends/dspark_trainer_backend.py | 4 +- verl_speco/backends/eagle3_trainer_backend.py | 31 +- .../models/dflash/configuration_dflash.py | 3 + verl_speco/trainer/base_trainer.py | 12 +- verl_speco/trainer/draft_training_loop.py | 399 +++++++++++------- verl_speco/trainer/standalone_checkpoint.py | 282 +++++++++++++ 8 files changed, 803 insertions(+), 169 deletions(-) create mode 100644 verl_speco/trainer/standalone_checkpoint.py diff --git a/tests/unit/test_draft_training_loop.py b/tests/unit/test_draft_training_loop.py index fbd1e92f..20f3459e 100644 --- a/tests/unit/test_draft_training_loop.py +++ b/tests/unit/test_draft_training_loop.py @@ -8,6 +8,7 @@ torch = pytest.importorskip("torch") +from verl_speco.trainer.standalone_checkpoint import rewrite_standalone_runtime_config # noqa: E402 from verl_speco.trainer.draft_training_loop import ( # noqa: E402 _rewrite_standalone_block_runtime_config, _save_standalone_checkpoint, @@ -148,6 +149,50 @@ def test_standalone_dspark_checkpoint_preserves_source_runtime_config(tmp_path): assert saved_training_config == training_config +def test_standalone_dspark_checkpoint_rewrites_generic_qwen3_architecture(tmp_path): + checkpoint_dir = tmp_path / "draft_step_5" + checkpoint_dir.mkdir() + source_dir = tmp_path / "source_dspark" + source_dir.mkdir() + target_dir = tmp_path / "target_qwen3" + target_dir.mkdir() + (source_dir / "config.json").write_text( + json.dumps( + { + "model_type": "qwen3", + "architectures": ["DSparkDraftModel"], + "markov_head_type": "vanilla", + } + ), + encoding="utf-8", + ) + (target_dir / "config.json").write_text(json.dumps({"model_type": "qwen3"}), encoding="utf-8") + (checkpoint_dir / "config.json").write_text( + json.dumps( + { + "model_type": "dspark", + "architectures": ["DSparkDraftModel"], + "markov_head_type": "vanilla", + } + ), + encoding="utf-8", + ) + trainer = SimpleNamespace( + backend=SimpleNamespace(model_type="dspark"), + config=SimpleNamespace( + model=SimpleNamespace(path=str(target_dir)), + rollout=SimpleNamespace(drafter=SimpleNamespace(model_path=str(source_dir))), + ), + ) + + rewrite_standalone_runtime_config(trainer, str(checkpoint_dir)) + + runtime_config = json.loads((checkpoint_dir / "config.json").read_text(encoding="utf-8")) + assert runtime_config["model_type"] == "qwen3" + assert runtime_config["architectures"] == ["Qwen3DSparkModel"] + assert runtime_config["speco_training_model_type"] == "dspark" + + def test_standalone_dflash_checkpoint_preserves_source_runtime_config(tmp_path): checkpoint_dir = tmp_path / "draft_step_5" checkpoint_dir.mkdir() @@ -182,6 +227,188 @@ def test_standalone_dflash_checkpoint_preserves_source_runtime_config(tmp_path): assert saved_training_config == training_config +def test_standalone_block_checkpoint_uses_target_model_type_without_source_config(tmp_path): + checkpoint_dir = tmp_path / "draft_step_5" + checkpoint_dir.mkdir() + target_dir = tmp_path / "target_qwen3" + target_dir.mkdir() + missing_source_dir = tmp_path / "missing_source_dspark" + (target_dir / "config.json").write_text( + json.dumps( + { + "model_type": "qwen3", + "head_dim": 128, + "rope_theta": 1000000.0, + "max_position_embeddings": 40960, + } + ), + encoding="utf-8", + ) + training_config = { + "model_type": "dspark", + "architectures": ["DSparkDraftModel"], + "target_layer_ids": [1, 9, 17], + "markov_head_type": "vanilla", + "head_dim": 80, + "rope_theta": 10000.0, + } + (checkpoint_dir / "config.json").write_text(json.dumps(training_config), encoding="utf-8") + trainer = SimpleNamespace( + backend=SimpleNamespace(model_type="dspark"), + config=SimpleNamespace( + model=SimpleNamespace(path=str(target_dir)), + rollout=SimpleNamespace(drafter=SimpleNamespace(model_path=str(missing_source_dir))), + ), + ) + + source_model_path = rewrite_standalone_runtime_config(trainer, str(checkpoint_dir)) + + runtime_config = json.loads((checkpoint_dir / "config.json").read_text(encoding="utf-8")) + saved_training_config = json.loads((checkpoint_dir / "speco_training_config.json").read_text(encoding="utf-8")) + assert source_model_path == str(missing_source_dir) + assert runtime_config["model_type"] == "qwen3" + assert runtime_config["architectures"] == ["Qwen3DSparkModel"] + assert runtime_config["draft_model_type"] == "dspark" + assert runtime_config["speculative_algorithm"] == "DSPARK" + assert runtime_config["speco_training_model_type"] == "dspark" + assert runtime_config["head_dim"] == 128 + assert runtime_config["max_position_embeddings"] == 40960 + assert runtime_config["dflash_config"]["head_dim"] == 128 + assert runtime_config["dspark_config"]["head_dim"] == 128 + assert runtime_config["dspark_config"]["markov_head_type"] == "vanilla" + assert runtime_config["rope_parameters"] == {"rope_theta": 1000000.0, "rope_type": "default"} + assert "rope_theta" not in runtime_config + assert "rope_theta" not in runtime_config["dflash_config"] + assert "rope_theta" not in runtime_config["dspark_config"] + assert saved_training_config == training_config + + +def test_standalone_dflash_checkpoint_preserves_source_lm_head(tmp_path): + safetensors_torch = pytest.importorskip("safetensors.torch") + checkpoint_dir = tmp_path / "draft_step_5" + checkpoint_dir.mkdir() + source_dir = tmp_path / "source_dflash" + source_dir.mkdir() + (source_dir / "config.json").write_text( + json.dumps({"model_type": "qwen3", "architectures": ["DFlashForCausalLM"]}), + encoding="utf-8", + ) + (checkpoint_dir / "config.json").write_text( + json.dumps({"model_type": "dflash", "architectures": ["DFlashDraftModel"]}), + encoding="utf-8", + ) + safetensors_torch.save_file({"lm_head.weight": torch.ones(3, 4)}, str(source_dir / "model.safetensors")) + safetensors_torch.save_file({"fc.weight": torch.ones(2, 2)}, str(checkpoint_dir / "model.safetensors")) + trainer = SimpleNamespace( + backend=SimpleNamespace(model_type="dflash"), + config=SimpleNamespace(rollout=SimpleNamespace(drafter=SimpleNamespace(model_path=str(source_dir)))), + ) + + _rewrite_standalone_block_runtime_config(trainer, str(checkpoint_dir)) + + exported_state = safetensors_torch.load_file(str(checkpoint_dir / "model.safetensors"), device="cpu") + assert torch.equal(exported_state["lm_head.weight"], torch.ones(3, 4)) + assert torch.equal(exported_state["fc.weight"], torch.ones(2, 2)) + + +def test_standalone_dflash_checkpoint_does_not_create_lm_head_without_source(tmp_path): + safetensors_torch = pytest.importorskip("safetensors.torch") + checkpoint_dir = tmp_path / "draft_step_5" + checkpoint_dir.mkdir() + target_dir = tmp_path / "target_qwen3" + target_dir.mkdir() + missing_source_dir = tmp_path / "missing_source_dflash" + (target_dir / "config.json").write_text(json.dumps({"model_type": "qwen3"}), encoding="utf-8") + (checkpoint_dir / "config.json").write_text( + json.dumps({"model_type": "dflash", "architectures": ["DFlashDraftModel"]}), + encoding="utf-8", + ) + safetensors_torch.save_file({"model.embed_tokens.weight": torch.ones(3, 4)}, str(target_dir / "model.safetensors")) + safetensors_torch.save_file({"fc.weight": torch.ones(2, 2)}, str(checkpoint_dir / "model.safetensors")) + trainer = SimpleNamespace( + backend=SimpleNamespace(model_type="dflash"), + config=SimpleNamespace( + model=SimpleNamespace(path=str(target_dir)), + rollout=SimpleNamespace(drafter=SimpleNamespace(model_path=str(missing_source_dir))), + ), + ) + + _rewrite_standalone_block_runtime_config(trainer, str(checkpoint_dir)) + + exported_state = safetensors_torch.load_file(str(checkpoint_dir / "model.safetensors"), device="cpu") + runtime_config = json.loads((checkpoint_dir / "config.json").read_text(encoding="utf-8")) + assert runtime_config["model_type"] == "qwen3" + assert "lm_head.weight" not in exported_state + assert torch.equal(exported_state["fc.weight"], torch.ones(2, 2)) + + +def test_standalone_dspark_checkpoint_appends_target_tied_embedding(tmp_path): + safetensors_torch = pytest.importorskip("safetensors.torch") + checkpoint_dir = tmp_path / "draft_step_5" + checkpoint_dir.mkdir() + target_dir = tmp_path / "target_qwen3" + target_dir.mkdir() + missing_source_dir = tmp_path / "missing_source_dspark" + (target_dir / "config.json").write_text( + json.dumps({"model_type": "qwen3", "tie_word_embeddings": True}), + encoding="utf-8", + ) + (checkpoint_dir / "config.json").write_text( + json.dumps({"model_type": "dspark", "architectures": ["DSparkDraftModel"]}), + encoding="utf-8", + ) + embedding = torch.arange(12, dtype=torch.float32).reshape(3, 4) + safetensors_torch.save_file({"model.embed_tokens.weight": embedding}, str(target_dir / "model.safetensors")) + safetensors_torch.save_file({"fc.weight": torch.ones(2, 2)}, str(checkpoint_dir / "model.safetensors")) + trainer = SimpleNamespace( + backend=SimpleNamespace(model_type="dspark"), + config=SimpleNamespace( + model=SimpleNamespace(path=str(target_dir)), + rollout=SimpleNamespace(drafter=SimpleNamespace(model_path=str(missing_source_dir))), + ), + ) + + _rewrite_standalone_block_runtime_config(trainer, str(checkpoint_dir)) + + exported_state = safetensors_torch.load_file(str(checkpoint_dir / "model.safetensors"), device="cpu") + runtime_config = json.loads((checkpoint_dir / "config.json").read_text(encoding="utf-8")) + assert runtime_config["model_type"] == "qwen3" + assert torch.equal(exported_state["lm_head.weight"], embedding) + assert torch.equal(exported_state["fc.weight"], torch.ones(2, 2)) + + +def test_standalone_dspark_checkpoint_skips_untied_target_embedding(tmp_path): + safetensors_torch = pytest.importorskip("safetensors.torch") + checkpoint_dir = tmp_path / "draft_step_5" + checkpoint_dir.mkdir() + target_dir = tmp_path / "target_qwen3" + target_dir.mkdir() + missing_source_dir = tmp_path / "missing_source_dspark" + (target_dir / "config.json").write_text( + json.dumps({"model_type": "qwen3", "tie_word_embeddings": False}), + encoding="utf-8", + ) + (checkpoint_dir / "config.json").write_text( + json.dumps({"model_type": "dspark", "architectures": ["DSparkDraftModel"]}), + encoding="utf-8", + ) + safetensors_torch.save_file({"model.embed_tokens.weight": torch.ones(3, 4)}, str(target_dir / "model.safetensors")) + safetensors_torch.save_file({"fc.weight": torch.ones(2, 2)}, str(checkpoint_dir / "model.safetensors")) + trainer = SimpleNamespace( + backend=SimpleNamespace(model_type="dspark"), + config=SimpleNamespace( + model=SimpleNamespace(path=str(target_dir)), + rollout=SimpleNamespace(drafter=SimpleNamespace(model_path=str(missing_source_dir))), + ), + ) + + _rewrite_standalone_block_runtime_config(trainer, str(checkpoint_dir)) + + exported_state = safetensors_torch.load_file(str(checkpoint_dir / "model.safetensors"), device="cpu") + assert "lm_head.weight" not in exported_state + assert torch.equal(exported_state["fc.weight"], torch.ones(2, 2)) + + def test_standalone_block_checkpoint_appends_source_lm_head_weight(tmp_path): safetensors_torch = pytest.importorskip("safetensors.torch") checkpoint_dir = tmp_path / "draft_step_5" diff --git a/verl_speco/backends/dflash_trainer_backend.py b/verl_speco/backends/dflash_trainer_backend.py index bdbaa3b9..8bc0e043 100644 --- a/verl_speco/backends/dflash_trainer_backend.py +++ b/verl_speco/backends/dflash_trainer_backend.py @@ -492,6 +492,7 @@ def _build_fallback_config(self, target_hf_config): target_num_hidden_layers = int(getattr(target_text_config, "num_hidden_layers", 36)) mask_token_id_cfg = training_cfg.get("dflash_mask_token_id", None) mask_token_id = int(mask_token_id_cfg if mask_token_id_cfg is not None else target_text_config.vocab_size - 1) + target_head_dim = getattr(target_text_config, "head_dim", None) target_layer_ids = training_cfg.get("dflash_target_layer_ids", None) if target_layer_ids is None: target_layer_ids = build_target_layer_ids(num_context_layers, target_num_hidden_layers) @@ -501,10 +502,11 @@ def _build_fallback_config(self, target_hf_config): num_hidden_layers=int(training_cfg.get("dflash_num_hidden_layers", 1)), num_attention_heads=int(getattr(target_text_config, "num_attention_heads")), num_key_value_heads=int(getattr(target_text_config, "num_key_value_heads", getattr(target_text_config, "num_attention_heads"))), + head_dim=int(target_head_dim) if target_head_dim is not None else None, vocab_size=int(target_text_config.vocab_size), rms_norm_eps=float(getattr(target_text_config, "rms_norm_eps", 1e-6)), max_position_embeddings=int(getattr(target_text_config, "max_position_embeddings", 32768)), - rope_theta=float(getattr(target_text_config, "rope_theta", 10000.0)), + rope_theta=self._target_rope_theta(target_text_config), num_target_layers=target_num_hidden_layers, num_context_layers=num_context_layers, target_hidden_size=int(target_text_config.hidden_size), @@ -514,6 +516,16 @@ def _build_fallback_config(self, target_hf_config): architectures=["DFlashDraftModel"], ) + @staticmethod + def _target_rope_theta(target_text_config) -> float: + rope_theta = getattr(target_text_config, "rope_theta", None) + if rope_theta is not None: + return float(rope_theta) + rope_parameters = getattr(target_text_config, "rope_parameters", None) + if isinstance(rope_parameters, dict) and rope_parameters.get("rope_theta") is not None: + return float(rope_parameters["rope_theta"]) + return 10000.0 + def _load_state_file(self, path: str) -> dict: if path.endswith(".safetensors"): with safe_open(path, framework="pt", device="cpu") as f: diff --git a/verl_speco/backends/dspark_trainer_backend.py b/verl_speco/backends/dspark_trainer_backend.py index 81da6a26..a0c193d0 100644 --- a/verl_speco/backends/dspark_trainer_backend.py +++ b/verl_speco/backends/dspark_trainer_backend.py @@ -552,6 +552,7 @@ def _build_fallback_config(self, target_hf_config): target_num_hidden_layers = int(getattr(target_text_config, "num_hidden_layers", 36)) mask_token_id_cfg = self._training_value(training_cfg, "dspark_mask_token_id", "dflash_mask_token_id", None) mask_token_id = int(mask_token_id_cfg if mask_token_id_cfg is not None else target_text_config.vocab_size - 1) + target_head_dim = getattr(target_text_config, "head_dim", None) target_layer_ids = self._training_value(training_cfg, "dspark_target_layer_ids", "dflash_target_layer_ids", None) if target_layer_ids is None: from verl_speco.models.dflash import build_target_layer_ids @@ -563,10 +564,11 @@ def _build_fallback_config(self, target_hf_config): num_hidden_layers=int(self._training_value(training_cfg, "dspark_num_hidden_layers", "dflash_num_hidden_layers", 1)), num_attention_heads=int(getattr(target_text_config, "num_attention_heads")), num_key_value_heads=int(getattr(target_text_config, "num_key_value_heads", getattr(target_text_config, "num_attention_heads"))), + head_dim=int(target_head_dim) if target_head_dim is not None else None, vocab_size=int(target_text_config.vocab_size), rms_norm_eps=float(getattr(target_text_config, "rms_norm_eps", 1e-6)), max_position_embeddings=int(getattr(target_text_config, "max_position_embeddings", 32768)), - rope_theta=float(getattr(target_text_config, "rope_theta", 10000.0)), + rope_theta=self._target_rope_theta(target_text_config), num_target_layers=target_num_hidden_layers, num_context_layers=num_context_layers, target_hidden_size=int(target_text_config.hidden_size), diff --git a/verl_speco/backends/eagle3_trainer_backend.py b/verl_speco/backends/eagle3_trainer_backend.py index 6178aa96..fd3e970a 100644 --- a/verl_speco/backends/eagle3_trainer_backend.py +++ b/verl_speco/backends/eagle3_trainer_backend.py @@ -1019,6 +1019,14 @@ def compute_loss(self, model, batch, _current_pad_size): quality_tokens = torch.tensor(0.0, device=input_ids.device, dtype=torch.float32) quality_topk = min(5, int(all_step_logits[0].size(-1))) quality_step_stats = [] + collect_diagnostics = bool(getattr(self, "enable_standalone_training_metrics", False)) + loss_sum_per_position = None + correct_per_position = None + count_per_position = None + if collect_diagnostics: + loss_sum_per_position = torch.zeros(length, device=input_ids.device, dtype=torch.float32) + correct_per_position = torch.zeros(length, device=input_ids.device, dtype=torch.float32) + count_per_position = torch.zeros(length, device=input_ids.device, dtype=torch.float32) sparse_base_tokens = torch.tensor(0.0, device=input_ids.device, dtype=torch.float32) sparse_valid_tokens = torch.tensor(0.0, device=input_ids.device, dtype=torch.float32) sparse_intersection_sum = torch.tensor(0.0, device=input_ids.device, dtype=torch.float32) @@ -1088,6 +1096,9 @@ def compute_loss(self, model, batch, _current_pad_size): quality_topk_correct += step_topk_correct step_tokens = valid_position.float().sum() quality_tokens += step_tokens + if collect_diagnostics: + correct_per_position[idx] = step_top1_correct + count_per_position[idx] = step_tokens quality_step_stats.append( { "step": idx, @@ -1103,6 +1114,8 @@ def compute_loss(self, model, batch, _current_pad_size): } ) step_loss_sum = per_token_ploss.sum() + if collect_diagnostics: + loss_sum_per_position[idx] = step_loss_sum # Apply EAGLE3 step-wise temporal decay total_local_ploss += (gamma ** idx) * step_loss_sum @@ -1134,13 +1147,27 @@ def compute_loss(self, model, batch, _current_pad_size): quality_step_stats, ) - return { + result = { "total_local_vloss": torch.tensor(0.0, device=input_ids.device), "total_local_ploss": total_local_ploss, "local_num_tokens": total_local_tokens, "v_weight": 0.0, - "p_weight": 1.0 + "p_weight": 1.0, } + if collect_diagnostics: + result["diagnostics"] = { + "correct_count": quality_top1_correct.detach(), + "eval_token_count": quality_tokens.detach(), + "top1_correct_count": quality_top1_correct.detach(), + "top5_correct_count": quality_topk_correct.detach(), + "quality_token_count": quality_tokens.detach(), + "valid_token_count": quality_tokens.detach(), + "weighted_token_count": total_local_tokens.detach(), + "loss_sum_per_position": loss_sum_per_position.detach(), + "correct_per_position": correct_per_position.detach(), + "count_per_position": count_per_position.detach(), + } + return result def _compute_target_p_padded(self, target_scores, t2d, loss_mask, length): with torch.no_grad(): diff --git a/verl_speco/models/dflash/configuration_dflash.py b/verl_speco/models/dflash/configuration_dflash.py index b9defde2..9dd162ab 100644 --- a/verl_speco/models/dflash/configuration_dflash.py +++ b/verl_speco/models/dflash/configuration_dflash.py @@ -23,6 +23,7 @@ def __init__( num_hidden_layers: int = 1, num_attention_heads: int = 32, num_key_value_heads: int = 8, + head_dim: Optional[int] = None, vocab_size: int = 152064, rms_norm_eps: float = 1e-6, max_position_embeddings: int = 32768, @@ -42,6 +43,8 @@ def __init__( self.num_hidden_layers = num_hidden_layers self.num_attention_heads = num_attention_heads self.num_key_value_heads = num_key_value_heads + if head_dim is not None: + self.head_dim = int(head_dim) self.vocab_size = vocab_size self.rms_norm_eps = rms_norm_eps self.max_position_embeddings = max_position_embeddings diff --git a/verl_speco/trainer/base_trainer.py b/verl_speco/trainer/base_trainer.py index 286b684e..39c308cd 100644 --- a/verl_speco/trainer/base_trainer.py +++ b/verl_speco/trainer/base_trainer.py @@ -563,7 +563,7 @@ def _is_block_drafter_backend(self) -> bool: def _block_drafter_metric_prefix(self) -> str: model_type = str(getattr(self.backend, "model_type", "dflash") or "dflash") - if model_type in {"dspark", "domino"}: + if model_type in {"dspark", "domino", "eagle3"}: return model_type return "dflash" @@ -619,7 +619,15 @@ def get_training_metrics(self) -> dict[str, float]: if self.optimizer is not None and self.optimizer.param_groups: metrics["drafter/current_lr"] = float(self.optimizer.param_groups[0]["lr"]) - for pos in range(int(self._block_drafter_config_value("block_size", 16))): + count_prefix = f"{prefix}/count_per_position/" + positions = sorted( + int(key.removeprefix(count_prefix)) + for key in sums + if key.startswith(count_prefix) and key.removeprefix(count_prefix).isdigit() + ) + if not positions: + positions = list(range(int(self._block_drafter_config_value("block_size", 16)))) + for pos in positions: count_key = f"{prefix}/count_per_position/{pos}" count = sums.get(count_key, 0.0) if count <= 0: diff --git a/verl_speco/trainer/draft_training_loop.py b/verl_speco/trainer/draft_training_loop.py index c9de8fe7..16b29671 100644 --- a/verl_speco/trainer/draft_training_loop.py +++ b/verl_speco/trainer/draft_training_loop.py @@ -3,10 +3,10 @@ from __future__ import annotations import asyncio -from copy import deepcopy import json import logging import os +import time from typing import Any import torch @@ -20,6 +20,7 @@ from verl_speco.trainer.base_trainer import DrafterBaseTrainer from verl_speco.trainer.draft_dataset import DraftFeatureDataLoader, DraftFeatureDataLoaderConfig from verl_speco.trainer.feature_store import build_feature_store_from_config +from verl_speco.trainer.standalone_checkpoint import rewrite_standalone_runtime_config logger = logging.getLogger(__name__) @@ -41,6 +42,7 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: _configure_device(local_rank) backend = _build_backend(draft_config) + setattr(backend, "enable_standalone_training_metrics", True) training_device_mesh = _build_training_device_mesh(draft_config, world_size) trainer = DrafterBaseTrainer( config=draft_config, @@ -58,7 +60,6 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: data_parallel_process_group=None, backend=backend, ) - max_steps = int(training_cfg.get("max_steps", training_cfg.get("step", 1000)) or 0) save_interval = int(training_cfg.get("save_interval_steps", 0) or 0) successful_steps = 0 @@ -91,6 +92,7 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: for samples in loader: if max_steps > 0 and successful_steps >= max_steps: break + step_started = time.perf_counter() attempted_batches += 1 batch = trainer.prepare_training_batch_from_samples(samples, step=optimizer_step) has_batch = batch is not None @@ -100,6 +102,7 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: continue if batch is None: continue + trainer.reset_training_metrics() ok = await trainer.training_step_from_batch(batch, optimizer_step) if not _all_ranks_true(ok, trainer.runtime_device): continue @@ -107,6 +110,13 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: optimizer_step = int(trainer.optimizer_steps_total) if optimizer_step <= initial_optimizer_step: optimizer_step = initial_optimizer_step + successful_steps + step_metrics = _standalone_step_metrics( + trainer, + successful_steps=successful_steps, + attempted_batches=attempted_batches, + step_elapsed_sec=time.perf_counter() - step_started, + ) + _log_standalone_step_metrics(step_metrics, rank=rank) if save_interval > 0 and optimizer_step % save_interval == 0: last_save_result = _save_standalone_checkpoint(trainer, optimizer_step) if _sync_any_rank_saved_checkpoint(last_save_result.get("saved")): @@ -198,52 +208,25 @@ def _save_standalone_checkpoint(trainer: DrafterBaseTrainer, step: int, *, wait: return { "saved": future is not None, "path": checkpoint_path, - "reason": "saved" if future is not None and wait else "scheduled" if future is not None else "not_checkpoint_leader", + "reason": ( + "saved" + if future is not None and wait + else "scheduled" if future is not None else "not_checkpoint_leader" + ), } -def _ensure_dict_child(config: dict[str, Any], key: str) -> dict[str, Any]: - value = config.get(key) - if isinstance(value, dict): - return value - value = {} - config[key] = value - return value - - -def _source_drafter_model_path(trainer: DrafterBaseTrainer) -> str | None: - model_path = getattr(getattr(getattr(trainer, "config", None), "rollout", None), "drafter", None) - model_path = getattr(model_path, "model_path", None) - if not model_path: - return None - return os.fspath(model_path) - - -def _load_source_drafter_config(trainer: DrafterBaseTrainer) -> dict[str, Any] | None: - model_path = _source_drafter_model_path(trainer) - if not model_path: - return None - config_path = os.path.join(model_path, "config.json") - if not os.path.exists(config_path): - return None - try: - with open(config_path, "r", encoding="utf-8") as f: - loaded = json.load(f) - except (OSError, json.JSONDecodeError) as exc: - logger.warning("Failed to load source drafter config %s: %s", config_path, exc) - return None - return loaded if isinstance(loaded, dict) else None - - -def _load_tensor_from_safetensors(path: str, key: str) -> torch.Tensor | None: +def _load_tensor_from_safetensors(path: str, keys: tuple[str, ...]) -> tuple[str, torch.Tensor] | None: try: from safetensors import safe_open with safe_open(path, framework="pt", device="cpu") as f: - if key in f.keys(): - return f.get_tensor(key) + available_keys = set(f.keys()) + for key in keys: + if key in available_keys: + return key, f.get_tensor(key) except Exception as exc: # noqa: BLE001 - logger.warning("Failed to load %s from %s: %s", key, path, exc) + logger.warning("Failed to load any of %s from %s: %s", keys, path, exc) return None @@ -254,23 +237,52 @@ def _torch_load_cpu(path: str) -> Any: return torch.load(path, map_location="cpu") -def _load_tensor_from_torch(path: str, key: str) -> torch.Tensor | None: +def _load_tensor_from_torch(path: str, keys: tuple[str, ...]) -> tuple[str, torch.Tensor] | None: try: state = _torch_load_cpu(path) except Exception as exc: # noqa: BLE001 - logger.warning("Failed to load %s from %s: %s", key, path, exc) + logger.warning("Failed to load any of %s from %s: %s", keys, path, exc) return None if isinstance(state, dict): - value = state.get(key) - if isinstance(value, torch.Tensor): - return value + for key in keys: + value = state.get(key) + if isinstance(value, torch.Tensor): + return key, value return None -def _load_source_lm_head_weight(model_path: str | None) -> torch.Tensor | None: +def _target_model_path_for_lm_head(trainer: DrafterBaseTrainer) -> str | None: + model_path = getattr(getattr(trainer, "config", None), "model", None) + model_path = getattr(model_path, "path", None) if not model_path: return None + return os.fspath(model_path) + +def _model_ties_word_embeddings(model_path: str | None) -> bool: + if not model_path: + return False + config_path = os.path.join(model_path, "config.json") + if not os.path.exists(config_path): + return False + try: + with open(config_path, "r", encoding="utf-8") as f: + config = json.load(f) + except (OSError, json.JSONDecodeError) as exc: + logger.warning("Failed to read model config %s for tied embedding check: %s", config_path, exc) + return False + return bool(isinstance(config, dict) and config.get("tie_word_embeddings") is True) + + +def _load_lm_head_weight( + model_path: str | None, + *, + allow_tied_embedding: bool = False, +) -> tuple[str, torch.Tensor] | None: + if not model_path: + return None + + keys = ("lm_head.weight", "model.embed_tokens.weight") if allow_tied_embedding else ("lm_head.weight",) for index_name in ("model.safetensors.index.json", "pytorch_model.bin.index.json"): index_path = os.path.join(model_path, index_name) if not os.path.exists(index_path): @@ -279,23 +291,31 @@ def _load_source_lm_head_weight(model_path: str | None) -> torch.Tensor | None: with open(index_path, "r", encoding="utf-8") as f: weight_map = json.load(f).get("weight_map", {}) except (OSError, json.JSONDecodeError) as exc: - logger.warning("Failed to read source weight index %s: %s", index_path, exc) + logger.warning("Failed to read weight index %s: %s", index_path, exc) continue - shard_name = weight_map.get("lm_head.weight") if isinstance(weight_map, dict) else None - if not shard_name: + if not isinstance(weight_map, dict): continue - shard_path = os.path.join(model_path, os.fspath(shard_name)) - if index_name.endswith(".safetensors.index.json"): - return _load_tensor_from_safetensors(shard_path, "lm_head.weight") - return _load_tensor_from_torch(shard_path, "lm_head.weight") + for key in keys: + shard_name = weight_map.get(key) + if not shard_name: + continue + shard_path = os.path.join(model_path, os.fspath(shard_name)) + if index_name.endswith(".safetensors.index.json"): + loaded = _load_tensor_from_safetensors(shard_path, (key,)) + else: + loaded = _load_tensor_from_torch(shard_path, (key,)) + if loaded is not None: + return loaded safetensors_path = os.path.join(model_path, "model.safetensors") if os.path.exists(safetensors_path): - return _load_tensor_from_safetensors(safetensors_path, "lm_head.weight") + loaded = _load_tensor_from_safetensors(safetensors_path, keys) + if loaded is not None: + return loaded torch_path = os.path.join(model_path, "pytorch_model.bin") if os.path.exists(torch_path): - return _load_tensor_from_torch(torch_path, "lm_head.weight") + return _load_tensor_from_torch(torch_path, keys) return None @@ -321,7 +341,9 @@ def _append_lm_head_to_safetensors_index(checkpoint_path: str, lm_head_weight: t weight_map["lm_head.weight"] = shard_name metadata = index_data.setdefault("metadata", {}) if isinstance(metadata, dict) and "total_size" in metadata: - metadata["total_size"] = int(metadata["total_size"]) + lm_head_weight.numel() * lm_head_weight.element_size() + metadata["total_size"] = int(metadata["total_size"]) + ( + lm_head_weight.numel() * lm_head_weight.element_size() + ) with open(index_path, "w", encoding="utf-8") as f: json.dump(index_data, f, indent=2, sort_keys=True) f.write("\n") @@ -350,7 +372,9 @@ def _append_lm_head_to_torch_index(checkpoint_path: str, lm_head_weight: torch.T weight_map["lm_head.weight"] = shard_name metadata = index_data.setdefault("metadata", {}) if isinstance(metadata, dict) and "total_size" in metadata: - metadata["total_size"] = int(metadata["total_size"]) + lm_head_weight.numel() * lm_head_weight.element_size() + metadata["total_size"] = int(metadata["total_size"]) + ( + lm_head_weight.numel() * lm_head_weight.element_size() + ) with open(index_path, "w", encoding="utf-8") as f: json.dump(index_data, f, indent=2, sort_keys=True) f.write("\n") @@ -360,12 +384,32 @@ def _append_lm_head_to_torch_index(checkpoint_path: str, lm_head_weight: torch.T return True -def _append_lm_head_weight_if_missing(checkpoint_path: str, source_model_path: str | None) -> None: - lm_head_weight = _load_source_lm_head_weight(source_model_path) - if lm_head_weight is None: - logger.warning("Standalone checkpoint export could not find source lm_head.weight in %s", source_model_path) +def _append_lm_head_weight_if_missing( + checkpoint_path: str, + source_model_path: str | None, + target_model_path: str | None, + *, + allow_target_fallback: bool, +) -> None: + loaded = _load_lm_head_weight(source_model_path) + source_label = source_model_path + if loaded is None and allow_target_fallback: + loaded = _load_lm_head_weight(target_model_path) + source_label = target_model_path + if loaded is None and _model_ties_word_embeddings(target_model_path): + loaded = _load_lm_head_weight(target_model_path, allow_tied_embedding=True) + if loaded is None: + if allow_target_fallback: + logger.warning( + "Standalone DSpark checkpoint export could not find lm_head.weight in source %s, " + "or target lm_head.weight / tied model.embed_tokens.weight in target %s", + source_model_path, + target_model_path, + ) return + loaded_key, lm_head_weight = loaded lm_head_weight = lm_head_weight.detach().cpu() + logger.info("Using %s from %s as standalone drafter lm_head.weight", loaded_key, source_label) if _append_lm_head_to_safetensors_index(checkpoint_path, lm_head_weight): return @@ -403,13 +447,10 @@ def _append_lm_head_weight_if_missing(checkpoint_path: str, source_model_path: s logger.warning("Failed to append lm_head.weight to %s: %s", torch_path, exc) return - logger.warning("Standalone checkpoint export found no model.safetensors or pytorch_model.bin under %s", checkpoint_path) - - -def _fill_if_missing(dst: dict[str, Any], src: dict[str, Any], keys: tuple[str, ...]) -> None: - for key in keys: - if key in src and key not in dst: - dst[key] = deepcopy(src[key]) + logger.warning( + "Standalone checkpoint export found no model.safetensors or pytorch_model.bin under %s", + checkpoint_path, + ) def _rewrite_standalone_block_runtime_config( @@ -417,104 +458,22 @@ def _rewrite_standalone_block_runtime_config( checkpoint_path: str, completed_future=None, ) -> None: - """Export standalone DFlash/DSpark checkpoints with runtime-facing config. - - The training wrapper saves an internal SpeCo config. For standalone - checkpoints we keep the original drafter ``config.json`` as the runtime - contract and only merge the alias fields needed by vLLM/SGLang. - """ + source_model_path = rewrite_standalone_runtime_config(trainer, checkpoint_path, completed_future) backend_type = getattr(getattr(trainer, "backend", None), "model_type", None) - if backend_type not in {"dflash", "dspark"}: - return - - if completed_future is not None: - try: - completed_future.result() - except Exception as exc: # noqa: BLE001 - logger.warning("Skip standalone runtime config rewrite because checkpoint save failed: %s", exc) - return - - config_path = os.path.join(checkpoint_path, "config.json") - if not os.path.exists(config_path): - logger.warning("Cannot rewrite standalone runtime config: missing %s", config_path) - return - - try: - with open(config_path, "r", encoding="utf-8") as f: - training_config = json.load(f) - except (OSError, json.JSONDecodeError) as exc: - logger.warning("Cannot rewrite standalone runtime config %s: %s", config_path, exc) - return - if not isinstance(training_config, dict): - logger.warning("Cannot rewrite standalone runtime config %s: expected object", config_path) - return - - training_config_path = os.path.join(checkpoint_path, "speco_training_config.json") - try: - with open(training_config_path, "w", encoding="utf-8") as f: - json.dump(training_config, f, indent=2, sort_keys=True) - except OSError as exc: - logger.warning("Failed to write standalone training config copy %s: %s", training_config_path, exc) - - source_model_path = _source_drafter_model_path(trainer) - runtime_config = _load_source_drafter_config(trainer) - if runtime_config is None: - runtime_config = deepcopy(training_config) - logger.warning( - "Source drafter config is unavailable; standalone checkpoint keeps SpeCo training config as runtime config" - ) - - runtime_config["speco_training_model_type"] = backend_type - common_alias_keys = ("target_layer_ids", "mask_token_id", "num_context_layers") - _fill_if_missing(runtime_config, training_config, common_alias_keys) - - dflash_config = _ensure_dict_child(runtime_config, "dflash_config") - _fill_if_missing(dflash_config, training_config, common_alias_keys) - if backend_type == "dspark": - dspark_config = _ensure_dict_child(runtime_config, "dspark_config") - _fill_if_missing( - dspark_config, - training_config, - ( - "block_size", - "num_anchors", - "markov_rank", - "markov_head_type", - "confidence_head_alpha", - "confidence_head_with_markov", - "ce_loss_alpha", - "l1_loss_alpha", - "loss_decay_gamma", - "target_layer_ids", - "num_context_layers", - "num_target_layers", - "target_num_hidden_layers", - "mask_token_id", - ), + _append_lm_head_weight_if_missing( + checkpoint_path, + source_model_path, + _target_model_path_for_lm_head(trainer), + allow_target_fallback=True, + ) + elif backend_type == "dflash": + _append_lm_head_weight_if_missing( + checkpoint_path, + source_model_path, + None, + allow_target_fallback=False, ) - else: - dspark_config = {} - - target_layer_ids = ( - runtime_config.get("target_layer_ids") - or dflash_config.get("target_layer_ids") - or dspark_config.get("target_layer_ids") - ) - if target_layer_ids is not None and "eagle_aux_hidden_state_layer_ids" not in runtime_config: - try: - runtime_config["eagle_aux_hidden_state_layer_ids"] = [int(layer_id) + 1 for layer_id in target_layer_ids] - except (TypeError, ValueError): - logger.warning("Invalid target_layer_ids in standalone exported config: %r", target_layer_ids) - - try: - with open(config_path, "w", encoding="utf-8") as f: - json.dump(runtime_config, f, indent=2, sort_keys=True) - f.write("\n") - except OSError as exc: - logger.warning("Failed to write standalone runtime config %s: %s", config_path, exc) - - _append_lm_head_weight_if_missing(checkpoint_path, source_model_path) def _disable_standalone_sequence_parallel(draft_config) -> None: @@ -544,6 +503,120 @@ def _build_training_device_mesh(draft_config, world_size: int) -> DeviceMesh | N ) + +def _block_metric_prefix(trainer: DrafterBaseTrainer) -> str | None: + model_type = str(getattr(getattr(trainer, "backend", None), "model_type", "") or "") + if model_type in {"dflash", "dspark", "eagle3"}: + return model_type + return None + + +def _current_learning_rate(trainer: DrafterBaseTrainer) -> float: + optimizer = getattr(trainer, "optimizer", None) + param_groups = getattr(optimizer, "param_groups", None) + if not param_groups: + return 0.0 + return float(param_groups[0].get("lr", 0.0)) + + +def _position_metric_series(metrics: dict[str, float], prefix: str, name: str) -> list[float]: + values: list[float] = [] + pos = 0 + while True: + key = f"{prefix}/{name}/{pos}" + if key not in metrics: + break + values.append(float(metrics[key])) + pos += 1 + return values + + +def _weighted_average(values: list[float], counts: list[float]) -> float | None: + if not values or not counts: + return None + total_count = sum(counts[: len(values)]) + if total_count <= 0: + return None + return sum(value * count for value, count in zip(values, counts, strict=False)) / total_count + + +def _simulated_accept_length(accuracies: list[float]) -> float: + cumulative = 1.0 + simulated = 0.0 + for accuracy in accuracies: + cumulative *= max(0.0, min(1.0, float(accuracy))) + simulated += cumulative + return simulated + + +def _standalone_step_metrics( + trainer: DrafterBaseTrainer, + *, + successful_steps: int, + attempted_batches: int, + step_elapsed_sec: float, +) -> dict[str, float]: + raw_metrics = trainer.get_training_metrics() + metrics: dict[str, float] = {key: float(value) for key, value in raw_metrics.items()} + prefix = _block_metric_prefix(trainer) + if prefix is not None: + anchor_offset = 1 if prefix == "dflash" else 0 + losses = _position_metric_series(raw_metrics, prefix, "loss_per_position") + accuracies = _position_metric_series(raw_metrics, prefix, "accuracy_per_position") + counts = _position_metric_series(raw_metrics, prefix, "count_per_position") + pred_losses = losses[anchor_offset:] + pred_accuracies = accuracies[anchor_offset:] + pred_counts = counts[anchor_offset:] + + avg_loss = _weighted_average(pred_losses, pred_counts) + avg_acc = _weighted_average(pred_accuracies, pred_counts) + if avg_loss is not None: + metrics["train/avg_loss"] = avg_loss + if avg_acc is not None: + metrics["train/avg_acc"] = avg_acc + if pred_accuracies: + metrics["train/simulated_acc_len"] = _simulated_accept_length(pred_accuracies) + if f"{prefix}/top1_acc" in raw_metrics: + metrics["train/top1_acc"] = float(raw_metrics[f"{prefix}/top1_acc"]) + if f"{prefix}/top5_acc" in raw_metrics: + metrics["train/top5_acc"] = float(raw_metrics[f"{prefix}/top5_acc"]) + for idx, value in enumerate(pred_losses): + metrics[f"train/ploss_{idx}"] = float(value) + for idx, value in enumerate(pred_accuracies): + metrics[f"train/acc_{idx}"] = float(value) + metrics["train/step"] = float(successful_steps) + metrics["train/global_step"] = float(getattr(trainer, "training_steps", successful_steps)) + metrics["train/lr"] = _current_learning_rate(trainer) + metrics["drafter/train_successful_steps"] = float(successful_steps) + metrics["drafter/train_attempted_batches"] = float(attempted_batches) + metrics["perf/step_time"] = float(step_elapsed_sec) + return metrics + + +def _log_standalone_step_metrics(metrics: dict[str, float], *, rank: int) -> None: + if rank != 0: + return + fields = [f"step={int(metrics.get('train/step', 0.0))}"] + for key, label in ( + ("train/avg_loss", "avg_loss"), + ("train/avg_acc", "avg_acc"), + ("train/top1_acc", "top1"), + ("train/top5_acc", "top5"), + ("train/simulated_acc_len", "sim_acc_len"), + ("train/lr", "lr"), + ("perf/step_time", "step_time"), + ): + if key not in metrics: + continue + value = float(metrics[key]) + if key == "train/lr": + fields.append(f"{label}={value:.3e}") + elif key == "perf/step_time": + fields.append(f"{label}={value:.3f}s") + else: + fields.append(f"{label}={value:.4f}") + logger.warning("[standalone drafter metrics] %s", " ".join(fields)) + def _init_distributed() -> tuple[int, int, int]: rank = int(os.environ.get("RANK", "0")) local_rank = int(os.environ.get("LOCAL_RANK", "0")) diff --git a/verl_speco/trainer/standalone_checkpoint.py b/verl_speco/trainer/standalone_checkpoint.py new file mode 100644 index 00000000..d5e50b19 --- /dev/null +++ b/verl_speco/trainer/standalone_checkpoint.py @@ -0,0 +1,282 @@ +"""Standalone drafter checkpoint runtime config helpers.""" + +from __future__ import annotations + +from copy import deepcopy +import json +import logging +import os +from typing import Any + +logger = logging.getLogger(__name__) + + +def _ensure_dict_child(config: dict[str, Any], key: str) -> dict[str, Any]: + value = config.get(key) + if isinstance(value, dict): + return value + value = {} + config[key] = value + return value + + +def _source_drafter_model_path(trainer: Any) -> str | None: + model_path = getattr(getattr(getattr(trainer, "config", None), "rollout", None), "drafter", None) + model_path = getattr(model_path, "model_path", None) + if not model_path: + return None + return os.fspath(model_path) + + +def _target_model_path(trainer: Any) -> str | None: + model_path = getattr(getattr(trainer, "config", None), "model", None) + model_path = getattr(model_path, "path", None) + if not model_path: + return None + return os.fspath(model_path) + + +def _load_source_drafter_config(trainer: Any) -> dict[str, Any] | None: + model_path = _source_drafter_model_path(trainer) + if not model_path: + return None + config_path = os.path.join(model_path, "config.json") + if not os.path.exists(config_path): + return None + try: + with open(config_path, "r", encoding="utf-8") as f: + loaded = json.load(f) + except (OSError, json.JSONDecodeError) as exc: + logger.warning("Failed to load source drafter config %s: %s", config_path, exc) + return None + return loaded if isinstance(loaded, dict) else None + + +def _load_target_runtime_config(trainer: Any) -> dict[str, Any] | None: + model_path = _target_model_path(trainer) + if not model_path: + return None + config_path = os.path.join(model_path, "config.json") + if not os.path.exists(config_path): + return None + try: + with open(config_path, "r", encoding="utf-8") as f: + loaded = json.load(f) + except (OSError, json.JSONDecodeError) as exc: + logger.warning("Failed to load target model config %s: %s", config_path, exc) + return None + return loaded if isinstance(loaded, dict) else None + + +def _target_runtime_model_type(target_config: dict[str, Any] | None) -> str | None: + if not isinstance(target_config, dict) or target_config.get("model_type") is None: + return None + return str(target_config["model_type"]) + + +def _fill_if_missing(dst: dict[str, Any], src: dict[str, Any], keys: tuple[str, ...]) -> None: + for key in keys: + if key in src and key not in dst: + dst[key] = deepcopy(src[key]) + + +def _copy_if_present(dst: dict[str, Any], src: dict[str, Any] | None, keys: tuple[str, ...]) -> None: + if src is None: + return + for key in keys: + if key in src: + dst[key] = deepcopy(src[key]) + + +def _drop_keys(config: dict[str, Any], keys: tuple[str, ...]) -> None: + for key in keys: + config.pop(key, None) + + +def _target_rope_parameters(target_config: dict[str, Any] | None) -> dict[str, Any] | None: + if not isinstance(target_config, dict): + return None + rope_parameters = target_config.get("rope_parameters") + if isinstance(rope_parameters, dict): + return deepcopy(rope_parameters) + rope_theta = target_config.get("rope_theta") + if rope_theta is not None: + return {"rope_theta": rope_theta, "rope_type": "default"} + return None + + +def _dspark_runtime_architecture(model_type: Any) -> str | None: + normalized = str(model_type or "").lower() + if normalized.startswith("qwen3"): + return "Qwen3DSparkModel" + if normalized.startswith("deepseek"): + return "DeepSeekDSparkModel" + return None + + +def _normalize_block_runtime_model_type( + runtime_config: dict[str, Any], + target_model_type: str | None, + backend_type: str, +) -> None: + training_model_types = {"dflash", "dspark", "qwen3_dspark"} + model_type = str(runtime_config.get("model_type") or "") + if model_type in training_model_types and target_model_type is not None: + runtime_config["model_type"] = target_model_type + runtime_config.setdefault("draft_model_type", backend_type) + runtime_config.setdefault("speculative_algorithm", backend_type.upper()) + + +def _normalize_dspark_runtime_architecture( + runtime_config: dict[str, Any], + target_model_type: str | None, +) -> None: + model_type = runtime_config.get("model_type") or target_model_type + architecture = _dspark_runtime_architecture(model_type) + if architecture is None: + architectures = runtime_config.get("architectures") or [] + if architectures == ["DSparkDraftModel"]: + logger.warning( + "Standalone DSpark runtime config keeps generic architecture for unsupported model_type=%r", + model_type, + ) + return + + architectures = runtime_config.get("architectures") or [] + if architectures != [architecture]: + runtime_config["architectures"] = [architecture] + + +def rewrite_standalone_runtime_config( + trainer: Any, + checkpoint_path: str, + completed_future: Any = None, +) -> str | None: + """Rewrite standalone DFlash/DSpark checkpoints with runtime-facing config. + + Returns the source drafter model path so callers can perform additional + checkpoint post-processing such as appending lm_head.weight. + """ + + backend_type = getattr(getattr(trainer, "backend", None), "model_type", None) + if backend_type not in {"dflash", "dspark"}: + return None + + if completed_future is not None: + try: + completed_future.result() + except Exception as exc: # noqa: BLE001 + logger.warning("Skip standalone runtime config rewrite because checkpoint save failed: %s", exc) + return None + + config_path = os.path.join(checkpoint_path, "config.json") + if not os.path.exists(config_path): + logger.warning("Cannot rewrite standalone runtime config: missing %s", config_path) + return None + + try: + with open(config_path, "r", encoding="utf-8") as f: + training_config = json.load(f) + except (OSError, json.JSONDecodeError) as exc: + logger.warning("Cannot rewrite standalone runtime config %s: %s", config_path, exc) + return None + if not isinstance(training_config, dict): + logger.warning("Cannot rewrite standalone runtime config %s: expected object", config_path) + return None + + training_config_path = os.path.join(checkpoint_path, "speco_training_config.json") + try: + with open(training_config_path, "w", encoding="utf-8") as f: + json.dump(training_config, f, indent=2, sort_keys=True) + except OSError as exc: + logger.warning("Failed to write standalone training config copy %s: %s", training_config_path, exc) + + source_model_path = _source_drafter_model_path(trainer) + source_runtime_config = _load_source_drafter_config(trainer) + runtime_config = source_runtime_config + target_runtime_config = _load_target_runtime_config(trainer) + target_model_type = _target_runtime_model_type(target_runtime_config) + if runtime_config is None: + runtime_config = deepcopy(training_config) + if target_model_type is not None: + runtime_config["model_type"] = target_model_type + runtime_config.setdefault("draft_model_type", backend_type) + runtime_config.setdefault("speculative_algorithm", backend_type.upper()) + else: + logger.warning( + "Source drafter config is unavailable and target model_type could not be inferred; " + "standalone checkpoint keeps SpeCo training config as runtime config" + ) + + _normalize_block_runtime_model_type(runtime_config, target_model_type, backend_type) + if backend_type == "dspark": + _normalize_dspark_runtime_architecture(runtime_config, target_model_type) + + runtime_config["speco_training_model_type"] = backend_type + target_runtime_keys = ("head_dim", "rope_theta", "max_position_embeddings") + common_alias_keys = ("head_dim", "rope_theta", "target_layer_ids", "mask_token_id", "num_context_layers") + _fill_if_missing(runtime_config, training_config, common_alias_keys) + _copy_if_present(runtime_config, target_runtime_config, target_runtime_keys) + + dflash_config = _ensure_dict_child(runtime_config, "dflash_config") + _fill_if_missing(dflash_config, training_config, common_alias_keys) + _copy_if_present(dflash_config, target_runtime_config, ("head_dim", "rope_theta")) + + if backend_type == "dspark": + dspark_config = _ensure_dict_child(runtime_config, "dspark_config") + _fill_if_missing( + dspark_config, + training_config, + ( + "block_size", + "head_dim", + "rope_theta", + "num_anchors", + "markov_rank", + "markov_head_type", + "confidence_head_alpha", + "confidence_head_with_markov", + "ce_loss_alpha", + "l1_loss_alpha", + "loss_decay_gamma", + "target_layer_ids", + "num_context_layers", + "num_target_layers", + "target_num_hidden_layers", + "mask_token_id", + ), + ) + _copy_if_present(dspark_config, target_runtime_config, ("head_dim", "rope_theta")) + else: + dspark_config = {} + + source_has_rope_theta = isinstance(source_runtime_config, dict) and "rope_theta" in source_runtime_config + if backend_type == "dspark" and str(target_model_type or "").lower().startswith("qwen3") and not source_has_rope_theta: + _drop_keys(runtime_config, ("rope_theta",)) + _drop_keys(dflash_config, ("rope_theta",)) + _drop_keys(dspark_config, ("rope_theta",)) + if "rope_parameters" not in runtime_config: + rope_parameters = _target_rope_parameters(target_runtime_config) + if rope_parameters is not None: + runtime_config["rope_parameters"] = rope_parameters + + target_layer_ids = ( + runtime_config.get("target_layer_ids") + or dflash_config.get("target_layer_ids") + or dspark_config.get("target_layer_ids") + ) + if target_layer_ids is not None and "eagle_aux_hidden_state_layer_ids" not in runtime_config: + try: + runtime_config["eagle_aux_hidden_state_layer_ids"] = [int(layer_id) + 1 for layer_id in target_layer_ids] + except (TypeError, ValueError): + logger.warning("Invalid target_layer_ids in standalone exported config: %r", target_layer_ids) + + try: + with open(config_path, "w", encoding="utf-8") as f: + json.dump(runtime_config, f, indent=2, sort_keys=True) + f.write("\n") + except OSError as exc: + logger.warning("Failed to write standalone runtime config %s: %s", config_path, exc) + return None + + return source_model_path From 09b395e4c9ef40f68e6f2597f047e3089839f2ef Mon Sep 17 00:00:00 2001 From: Cai Zeyong <878049625@qq.com> Date: Mon, 27 Jul 2026 16:08:06 +0800 Subject: [PATCH 20/50] =?UTF-8?q?tq=E9=80=82=E9=85=8D?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/transferqueue_integration_plan.md | 143 +++++++++ verl_speco/config/speco_base.yaml | 13 + verl_speco/integration/sglang_runtime.py | 20 ++ verl_speco/integration/task_runner.py | 7 + .../integration/transferqueue_bridge.py | 279 ++++++++++++++++++ verl_speco/workers/speco_worker.py | 33 ++- 6 files changed, 487 insertions(+), 8 deletions(-) create mode 100644 docs/transferqueue_integration_plan.md create mode 100644 verl_speco/integration/transferqueue_bridge.py diff --git a/docs/transferqueue_integration_plan.md b/docs/transferqueue_integration_plan.md new file mode 100644 index 00000000..56b7c3f8 --- /dev/null +++ b/docs/transferqueue_integration_plan.md @@ -0,0 +1,143 @@ +# verl-SpeCo TransferQueue 落地方案 + +> 目标:在**不修改上游 verl**的前提下,把 SpeCo online 训练里的逐样本特征流 +> 从「`SpecoRayPPOTrainer` driver 中转 + Ray object store」改为「TransferQueue +> 直传」,干掉 driver 这个数据瓶颈,并解锁流式消费与跨副本负载均衡。 +> +> 约束:仅 hook,与 SpeCo 现有 hook 模式一致;TQ 作为独立库使用,**不复用** verl +> 的 `main_ppo_sync` TQ 集成。 + +--- + +## 0. 现状:SpeCo online 特征流的 controller 瓶颈 + +SpeCo online 路径以 `SpecoRayPPOTrainer` 为 hub,所有跨进程大张量都被 driver 串行 +中转,介质是 Ray object store(`ray.put`/`ray.get`/`parallel_put`)。这与 verl 引入 +TQ 想解决的痛点 1:1 对应,只是 verl 干掉的是 `RayPPOTrainer`,我们要干掉的是 +SpeCo 在它之上加的 drafter 管线中转。 + +| # | 流向 | 当前机制 | 是否逐样本 | hook 位置(SpeCo 侧) | +|---|---|---|---|---| +| **a1** | target hidden states(SGLang 采集)-> drafter | `drafter_sample` 塞进 `DataProto.non_tensor_batch` -> driver pop/bucket -> `parallel_put` -> drafter `ray.get` | ✅ | `speco_ray_trainer.py` `generate_sequences_with_speco`;`sglang_adapter.py` `pop_drafter_samples`/`bucket_drafter_samples_by_replica`;`sglang_runtime.py` 组装 `drafter_sample` | +| **a2** | target hidden states(old-logprob hook)-> drafter | actor 前向 hook 截行 -> `ray.put` chunk -> driver 重打包 -> 分发 | ✅ | `oldlogprob_runtime.py` `_install_oldlogprob_hidden_hooks`/`_put_oldlogprob_hidden_refs`;`speco_ray_trainer.py` `_speco_collect_oldlogprob_features` | +| **b2** | target top-logprobs -> drafter(`use_logits=true`) | 随 a1 同一 side-channel | ✅ | `sglang_runtime.py` `target_logprobs`/`hidden_raw_target_logprobs` | +| **d** | rollout tokens -> drafter 训练集 | 随 a1 同一 side-channel(online)/`torch.save` 分片(offline) | ✅ | `sglang_runtime.py`;`speco_worker.py` `_store_rollout_sample` | +| b1 | target **lm_head 权重**(行)-> drafter `TargetHead` | ONE_TO_ALL Ray 分发 | ❌ 参数广播 | `rollout_publish.py` `export_actor_lm_head_weight`/`get_actor_lm_head_weight`;`speco_ray_trainer.py` `_speco_sync_target_lm_head_weight` | +| c | drafter 权重 -> rollout 引擎 | `ray.put` -> driver -> actor;vLLM 末段 ZMQ+SHM,SGLang 进程内 | ❌ 参数广播 | `speco_worker.py` `maybe_publish`;`rollout_publish.py` `update_draft_weights`;`vllm_runtime.py` `BucketedWeightSender` | + +**关键事实**:hidden states 跨进程前一律 CPU 物化(`oldlogprob_runtime.py`、 +`sglang_runtime.py`、`feature_store.py` 均 `.cpu()`),a1 路径下 driver 进程的 +host memory 会真正承载整批 hidden states 并做一次 Ray store 往返。这正是 TQ 要 +消除的往返。 + +--- + +## 1. 为什么不把"替换 feature_store"作为第一刀 + +`TorchShardFeatureStore`(`feature_store.py`)是 `torch.save` 分片 + JSONL manifest, +**只服务于 `collect_only`/`offline`**,不参与 online 热路径。替换它能统一离线存储 +抽象、换更快的分布式后端,但**不解决 controller 瓶颈**,性能收益有限。降级为 +可选尾项(见 §5 P3)。 + +--- + +## 2. 目标方案:TQ 直传逐样本特征流(a1 / a2 / b2 / d) + +### 2.1 角色映射 + +| TQ 角色 | SpeCo 对应 | +|---|---| +| Producer(写) | rollout worker(SGLang 路径,a1/b2/d)/ actor worker(old-logprob 路径,a2)——均在 SpeCo 既有 hook 内 | +| Consumer(读) | drafter worker `collect_rollout_features`(SpeCo 侧) | +| TransferQueueController(control plane) | SpeCo launcher 启动一个 Ray actor;drafter 经 `Sampler`/`StreamingDataLoader` 拉取 | +| Storage backend | `SimpleStorage`(ZMQ,跨节点 CPU 内存);进阶可切 `MooncakeStore`(RDMA,GPU-DRAM) | + +### 2.2 partition / key / 字段设计 + +- `partition_id`:`speco_train`(验证集用 `speco_val`)。 +- `key`:`{uid}_{session_id}_{index}`,与 verl TQ 一致;`uid` SpeCo 已有。 +- `tags`:`global_steps`、`source`∈{`rollout`,`oldlogprob`}、`replica_rank`/`owner_rank`、`status`、`prompt_len`/`response_len`/`seq_len`。ReplayBuffer/负载均衡按 tag 匹配。 +- `fields`(列):`input_ids`、`loss_mask`、`position_ids`、`hidden_states`、`last_hidden_states`/`target`、`target_logprobs`、`hidden_positions`、`prompts`、`responses`。与 `DraftFeatureSample`(`feature_store.py`)字段对齐,便于 online/offline 复用。 + +### 2.3 数据流(目标) + +``` +rollout/actor worker (SpeCo hook) + │ 生成/截取 hidden states 后,就地 tq.kv_batch_put(samples) + ▼ +TransferQueue (SimpleStorage, 跨节点 CPU 内存;可选 MooncakeStore RDMA) + │ control plane 按 sample 粒度追踪 ready 状态,Sampler 跨 drafter 副本均衡 + ▼ +drafter worker + │ tq.kv_batch_get / StreamingDataLoader 消费 → 喂入既有 DataBuffer / collect_online_data + ▼ +drafter 训练 (不变) +``` + +driver 只下发触发与轻量 key/meta,**不再承载 hidden states**。 + +--- + +## 3. 落地改动点(全部在 SpeCo 侧,hook-only) + +### 3.1 启动与配置 +- `draft_train_launcher.py` / `main.py`:`tq.init(config.transfer_queue)`;起 `TransferQueueController.remote(Sampler)`。 +- `config/speco_base.yaml`:新增 `drafter.transfer_queue` 块(backend、partition、enable 开关)。参考 verl `ppo_trainer.yaml` 的 `transfer_queue:` 结构,但**独立配置**,不复用 verl 的。 + +### 3.2 Producer 侧 +- **a1/b2/d(SGLang)**:`sglang_runtime.py` 组装 `drafter_sample` 处(~1594-1648),增加 `tq.kv_batch_put`;返回给 driver 的 `drafter_sample` 只保留 key/meta(或整段不再走 DataProto side-channel,driver 仅触发)。 +- **a2(old-logprob)**:`oldlogprob_runtime.py` `_put_oldlogprob_hidden_refs`(~216),把 `ray.put(hidden_chunk)` 换成 `tq.kv_batch_put`;`OLD_LOGPROB_HIDDEN_CHUNK_REFS_KEY` 改为 TQ key 列表。 + +### 3.3 Consumer 侧 +- `speco_worker.py` `collect_rollout_features`(~665):把 `_resolve_ray_object_ref`/`_resolve_hidden_state_chunks`(`ray.get`)换成 `tq.kv_batch_get`;`_dispatch_nd_compute`(~159)的 `parallel_put` 退化为只传 key(或 drafter 直接从 TQ Sampler 拉,driver 不参与分发)。 +- drafter 内部 `DataBuffer`/`collect_online_data`(`base_trainer.py`)保持不变,只是数据来源由 `ray.get` 改为 TQ get。 + +### 3.4 Driver 侧 +- `speco_ray_trainer.py`:`_speco_collect_rollout_features_rpc`/`speco_collect_rollout_features`(~351)、`_speco_collect_oldlogprob_features`(~1114)不再搬数据,只做触发/传 key;`bucket_drafter_samples_by_replica` 可由 TQ `Sampler` 替代(逐步迁移,先保留作回退)。 + +### 3.5 不改动 +- **b1(lm_head 权重)、c(drafter 权重)**:保持现状。与 verl 上游一致(权重不走 TQ),且 c 的 vLLM 末段已有专用 ZMQ+SHM 通道。 +- verl 本体:零改动。 + +--- + +## 4. 收益与边界(诚实评估) + +### 收益 +1. **去掉 driver 对 hidden states 的 host-memory 中转 + Ray store 往返**:producer 直存 TQ,consumer 直取,driver 不再承载整批特征。 +2. **流式消费**:drafter 在样本 ready 时即可消费,不必等整批 `generate_sequences` 返回,采集与训练可重叠。 +3. **跨 drafter 副本负载均衡**:TQ `Sampler`/`RankAwareSampler` 替代手写 `bucket_drafter_samples_by_replica`/`owner_rank` 分配。 +4. **(若采纳 P3)统一 online/collect_only/offline 存储**:同一 TQ partition,`collect_only` 写、`offline` 读,消掉 on-disk 分片层。 + +### 边界 / 不解决的事 +- 只优化**特征采集**这一子阶段,**不加速** rollout 本身、actor update、reward;e2e 增益取决于该子阶段在 step 中的占比。 SpeCo README 的 20% rollout / 11% e2e 提升来自 acceptance length,与本方案是不同机制,不要混为一谈。 +- **权重同步(b1/c)不放进 TQ**,与 verl 上游保持一致。 +- hidden states 跨进程前**仍需 CPU 物化**(现状如此);要避免物化需切 `MooncakeStore` RDMA,属进阶项。 +- 引入 TQ 依赖与一个 control-plane Ray actor,增加少量运维面。 + +### 风险 +- TQ 与 SpeCo 现有 `owner_rank`/`replica_rank` 路由语义需对齐(Sampler 要复刻「按 owner 分桶」语义,否则样本会错配 drafter 副本)。 +- old-logprob 的 chunk 拆分(`hidden_states_ref_chunks`)映射到 TQ 列式存储时,需保证 chunk meta 与 key 的一致性。 +- 回退路径:保留 `enable_transfer_queue=False` 时走原 Ray 路径,渐进切换。 + +--- + +## 5. 分阶段实施 + +| 阶段 | 范围 | 产出 | +|---|---|---| +| **P0** | a1(SGLang hidden states)走 TQ 直传;drafter `kv_batch_get` 消费;driver 仅触发 | 验证 controller-bypass 闭环 + 正确性 | +| **P1** | a2(old-logprob hidden states)走 TQ;chunk 拆分映射 TQ 列 | 覆盖第二条采集路径 | +| **P2** | b2(top-logprobs)+ d(tokens)随 a1 同 partition 传输;Sampler 替代手写 bucket | 完整特征流 + 跨副本均衡 | +| **P3(可选)** | `TorchShardFeatureStore` → TQ partition,统一 online/collect_only/offline | 离线工作流统一 | + +每个阶段保留 `enable_transfer_queue` 开关与原 Ray 路径回退。 + +--- + +## 6. 待确认决策 + +1. **TQ backend**:`SimpleStorage`(CPU 内存,默认)起步,还是直接上 `MooncakeStore`(RDMA,省 CPU 物化)?后者依赖 RDMA 网络,建议 P0 用 SimpleStorage。 +2. **drafter 消费模式**:`kv_batch_get`(主动拉,改动小)还是 `StreamingDataLoader`(全自动流式,改动大、收益高)?建议 P0 用前者,P2 再考虑后者。 +3. **driver 角色**:P0 先保留 driver 传 key(最小改动),还是直接让 drafter 从 TQ Sampler 自取(driver 彻底退出数据路径)?前者风险低,建议 P0 用前者。 +4. **是否做 P3**:离线统一是否在本次范围内,还是单独立项。 diff --git a/verl_speco/config/speco_base.yaml b/verl_speco/config/speco_base.yaml index f8eeb435..a01b8290 100644 --- a/verl_speco/config/speco_base.yaml +++ b/verl_speco/config/speco_base.yaml @@ -179,3 +179,16 @@ actor_rollout_ref: repeat: true prefetch_depth: 2 strict_schema: true + # TransferQueue transport for drafter features (P0: hidden states only). + # When enabled and the `transfer_queue` package is installed, the SGLang + # rollout server writes hidden states directly to TQ and the drafter + # worker reads them by key, bypassing the RayPPOTrainer driver and the + # Ray object store for the dominant tensor. Default off -> unchanged + # Ray path. Requires `pip install TransferQueue==0.1.7`. + transfer_queue: + enable: false + backend: + storage_backend: SimpleStorage + SimpleStorage: + total_storage_size: 100000 + num_data_storage_units: 8 diff --git a/verl_speco/integration/sglang_runtime.py b/verl_speco/integration/sglang_runtime.py index 121cf7d5..e4556863 100644 --- a/verl_speco/integration/sglang_runtime.py +++ b/verl_speco/integration/sglang_runtime.py @@ -1635,6 +1635,26 @@ async def generate( "global_step": collection_global_steps, "replica_rank": self.replica_rank, } + # P0: offload the dominant hidden_states tensor to + # TransferQueue so it bypasses the RayPPOTrainer driver and + # the Ray object store. The key rides with the sample dict; + # the drafter worker fetches by key. No-op when TQ disabled. + from verl_speco.integration.transferqueue_bridge import ( + configure_transfer_queue, + is_transfer_queue_enabled, + make_sample_key, + put_sample, + ) + configure_transfer_queue(training_cfg) + if is_transfer_queue_enabled(): + tq_key = make_sample_key(collection_global_steps, self.replica_rank, request_id) + put_sample( + tq_key, + {"hidden_states": hidden_states.unsqueeze(0).cpu()}, + tag={"global_step": collection_global_steps, "replica_rank": self.replica_rank}, + ) + drafter_sample["hidden_states_tq_key"] = tq_key + drafter_sample["hidden_states"] = None else: self._speco_log_missing_hidden_states_once( collection_global_steps=collection_global_steps, diff --git a/verl_speco/integration/task_runner.py b/verl_speco/integration/task_runner.py index b7fe041b..cb1c1835 100644 --- a/verl_speco/integration/task_runner.py +++ b/verl_speco/integration/task_runner.py @@ -269,5 +269,12 @@ def _run_with_speco_trainer(self, config): speco_worker_cls=speco_worker_cls, ) + # Bootstrap TransferQueue before spawning Ray actors so worker processes + # inherit the TQ environment (mirrors verl main_ppo_sync tq.init). No-op + # when transfer_queue.enable=false or the package is not installed. + from verl_speco.integration.transferqueue_bridge import init_transfer_queue + + init_transfer_queue(config) + trainer.init_workers() trainer.fit() diff --git a/verl_speco/integration/transferqueue_bridge.py b/verl_speco/integration/transferqueue_bridge.py new file mode 100644 index 00000000..e69b77f2 --- /dev/null +++ b/verl_speco/integration/transferqueue_bridge.py @@ -0,0 +1,279 @@ +"""TransferQueue bridge for SPECO drafter feature transport. + +This module lets SPECO route large per-sample drafter-training tensors (hidden +states, and later target logprobs) through TransferQueue (TQ) instead of +funneling them through the ``SpecoRayPPOTrainer`` driver process and the Ray +object store. It is the SpeCo-side analog of verl's ``transferqueue_utils.py``, +but used as a standalone transport library -- it does **not** depend on verl's +``main_ppo_sync`` TQ integration and does **not** modify upstream verl. + +Design (P0): +- Only the dominant tensor (``hidden_states``) is offloaded to TQ. The rest of + the ``drafter_sample`` dict (input_ids, prompts, responses, positions, + metadata scalars) keeps riding the existing DataProto side-channel + Ray + dispatch. The TQ key rides with the sample dict, so the driver is unchanged. +- Default ``enable: false`` -> behavior is bit-identical to the current Ray + path. TQ is only touched when explicitly enabled and the ``transfer_queue`` + package is importable. +- Producer = SGLang rollout server (``sglang_runtime.py``); consumer = drafter + worker (``speco_worker.py``). Both call ``configure_transfer_queue`` from + their respective drafter training config, then ``put_sample`` / ``get_sample``. + +Note: the TQ call sites follow the documented KV API (``kv_put`` / +``kv_batch_get`` / ``kv_clear``) of TransferQueue 0.1.7. When TQ is enabled, +these are exercised against the installed package; verify the exact signatures +against your TQ version on first run (the bridge fails loud, never silently). +""" + +from __future__ import annotations + +import logging +import os +import threading +from typing import Any, Optional + +import torch + +logger = logging.getLogger(__name__) +logger.setLevel(os.getenv("VERL_LOGGING_LEVEL", "WARN")) + +# Partition under which SPECO drafter samples are stored. +_SPECO_TQ_PARTITION = "speco_drafter_features" + +try: + import transfer_queue as tq # type: ignore + from transfer_queue import KVBatchMeta # noqa: F401 (re-exported for symmetry) + + _TQ_IMPORTABLE = True +except ImportError: + + _TQ_IMPORTABLE = False + + class KVBatchMeta: # type: ignore[no-redef] + """Stand-in used only when TransferQueue is not installed.""" + + class _MockTQ: + """Mock that raises on any use; only hit if enabled without TQ installed.""" + + def __getattr__(self, name: str) -> Any: + def _raise(*args: Any, **kwargs: Any) -> Any: + raise RuntimeError( + f"transfer_queue is not installed. Cannot call tq.{name}(). " + "Install with `pip install TransferQueue==0.1.7` or disable " + "actor_rollout_ref.rollout.drafter.training.transfer_queue.enable." + ) + + return _raise + + tq = _MockTQ() # type: ignore[assignment] + + +# --------------------------------------------------------------------------- +# State +# --------------------------------------------------------------------------- + +_state_lock = threading.Lock() +_state = { + "enabled": False, # config says enable=true + "configured": False, # configure_transfer_queue has run + "initialized": False, # tq.init() has run in this process + "config": None, # the transfer_queue sub-config (plain dict) +} + + +def configure_transfer_queue(training_cfg: Any) -> bool: + """Read the ``transfer_queue`` sub-config from the drafter training config. + + Called from both the SGLang server process (via env-serialized drafter + config) and the drafter worker process (via full hydra config). Idempotent. + Returns whether TQ transport is usable in this process. + """ + + global _state + tq_cfg = _extract_tq_config(training_cfg) + with _state_lock: + _state["config"] = tq_cfg + _state["enabled"] = bool(tq_cfg.get("enable")) if tq_cfg else False + _state["configured"] = True + usable = _state["enabled"] and _TQ_IMPORTABLE + if _state["enabled"] and not _TQ_IMPORTABLE: + logger.warning( + "[SpeCo TQ] transfer_queue.enable=true but transfer_queue package is " + "not installed; falling back to inline Ray transport." + ) + return usable + + +def _extract_tq_config(training_cfg: Any) -> Optional[dict]: + if training_cfg is None: + return None + transfer_queue_cfg = None + if hasattr(training_cfg, "get"): + transfer_queue_cfg = training_cfg.get("transfer_queue", None) + elif isinstance(training_cfg, dict): + transfer_queue_cfg = training_cfg.get("transfer_queue", None) + if transfer_queue_cfg is None: + return None + return _to_plain_dict(transfer_queue_cfg) + + +def _to_plain_dict(value: Any) -> dict: + """Convert OmegaConf DictConfig / nested mapping to a plain dict.""" + + if hasattr(value, "to_container"): + try: + import omegaconf + + return omegaconf.OmegaConf.to_container(value, resolve=True) # type: ignore[arg-type] + except Exception: # noqa: BLE001 + pass + if isinstance(value, dict): + return {k: _to_plain_dict(v) for k, v in value.items()} + return value + + +def is_transfer_queue_enabled() -> bool: + """True only if configured enabled AND the TQ package is importable.""" + + return bool(_state["enabled"]) and _TQ_IMPORTABLE + + +def init_transfer_queue(config: Any) -> bool: + """Cluster-wide TQ bootstrap, called once from the SpecoTaskRunner. + + Mirrors verl ``main_ppo_sync`` calling ``tq.init(config.transfer_queue)`` + in the TaskRunner before workers spawn. Ray actors (SGLang server, drafter + worker) inherit this process's environment, so their lazy ``tq.init()`` + (no-arg) connects to the already-started storage. Returns whether TQ is + usable; no-op (returns False) when disabled or not installed. + """ + + tq_cfg = _extract_tq_config(_drafter_training_cfg(config)) + if tq_cfg is None or not bool(tq_cfg.get("enable")) or not _TQ_IMPORTABLE: + return False + tq.init(_to_plain_dict(tq_cfg)) + with _state_lock: + _state["config"] = _to_plain_dict(tq_cfg) + _state["enabled"] = True + _state["initialized"] = True + logger.info("[SpeCo TQ] TransferQueue bootstrapped in task runner (partition=%s)", _SPECO_TQ_PARTITION) + return True + + +def _drafter_training_cfg(config: Any) -> Any: + try: + return config.actor_rollout_ref.rollout.drafter.training + except AttributeError: + return None + + +def _ensure_initialized() -> None: + """Lazily ``tq.init()`` once per worker process (mirrors verl TQ_INITIALIZED). + + No-arg init relies on env inheritance from the task-runner bootstrap. If a + future TQ version does not propagate config via env, switch this to + ``tq.init(_state["config"])``. + """ + + if _state["initialized"]: + return + with _state_lock: + if _state["initialized"]: + return + tq.init() + _state["initialized"] = True + + +# --------------------------------------------------------------------------- +# Key / put / get / clear +# --------------------------------------------------------------------------- + +def make_sample_key(global_step: Any, replica_rank: Any, request_id: Any) -> str: + """Build a deterministic, cluster-unique key for one drafter sample. + + Uniqueness space: (global_step, replica_rank, request_id). Each rollout + request produces exactly one drafter_sample, so this is unique per sample. + """ + + return f"speco:{global_step}:{replica_rank}:{request_id}" + + +def _to_tensordict(tensor_dict: dict) -> Any: + from tensordict import TensorDict + + return TensorDict(tensor_dict, batch_size=[]) + + +def put_sample( + key: str, + tensor_dict: dict, + *, + tag: Optional[dict] = None, +) -> None: + """Store a dict of CPU tensors under ``key`` in the SPECO TQ partition. + + ``tensor_dict`` values must be CPU ``torch.Tensor`` (or None). None values + are dropped. Raises if TQ is enabled but the call fails -- never silently + degrades, so a transport failure surfaces immediately rather than dropping + a sample. + """ + + if not is_transfer_queue_enabled(): + raise RuntimeError("put_sample called while TransferQueue is not enabled.") + payload = {k: v for k, v in tensor_dict.items() if torch.is_tensor(v)} + if not payload: + return + _ensure_initialized() + value = _to_tensordict(payload) + # KV API (TransferQueue docs): kv_put(key, value, partition_id, tag). + tq.kv_put(key, value, partition_id=_SPECO_TQ_PARTITION, tag=tag or {}) + + +def get_sample(key: str) -> dict: + """Retrieve the tensor dict stored under ``key`` and free it. + + Returns a plain ``{field: tensor}`` dict. Clears the key after read since + each sample is consumed by exactly one drafter owner replica. + """ + + if not is_transfer_queue_enabled(): + raise RuntimeError("get_sample called while TransferQueue is not enabled.") + _ensure_initialized() + # KV API: kv_batch_get(keys, partition_id) -> mapping key -> TensorDict + # (return shape is version-dependent; handle both dict and list forms). + result = tq.kv_batch_get([key], partition_id=_SPECO_TQ_PARTITION) + value = _extract_value(result, key) + try: + tq.kv_clear([key], partition_id=_SPECO_TQ_PARTITION) + except Exception: # noqa: BLE001 + logger.debug("[SpeCo TQ] kv_clear failed for key=%s (ignored)", key) + if value is None: + return {} + return _tensordict_to_dict(value) + + +def _extract_value(result: Any, key: str) -> Any: + if result is None: + return None + if isinstance(result, dict): + return result.get(key) + if isinstance(result, (list, tuple)): + return result[0] if len(result) > 0 else None + return result + + +def _tensordict_to_dict(value: Any) -> dict: + if hasattr(value, "items"): + return {k: v for k, v in value.items()} + return dict(value) + + +__all__ = [ + "KVBatchMeta", + "configure_transfer_queue", + "init_transfer_queue", + "is_transfer_queue_enabled", + "make_sample_key", + "put_sample", + "get_sample", +] diff --git a/verl_speco/workers/speco_worker.py b/verl_speco/workers/speco_worker.py index 63987715..7ec353e3 100644 --- a/verl_speco/workers/speco_worker.py +++ b/verl_speco/workers/speco_worker.py @@ -309,6 +309,14 @@ def __init__(self, config: DictConfig, role: str = "speco", device_name: Optiona self.publish_interval_steps = int(self.config.rollout.drafter.training.get("publish_interval_steps", 0)) self.train_steps_per_trigger = int(self.config.rollout.drafter.training.get("step", 100)) + # Configure TransferQueue transport for drafter features. No-op when + # disabled or when the transfer_queue package is not installed; the + # existing inline Ray path is used otherwise. Cached on the instance so + # collect_rollout_features can branch without re-reading config. + from verl_speco.integration.transferqueue_bridge import configure_transfer_queue + + self._speco_tq_enabled = configure_transfer_queue(self.config.rollout.drafter.training) + def _ensure_process_group_initialized(self): if not dist.is_initialized(): initialize_global_process_group_ray( @@ -703,15 +711,24 @@ def collect_rollout_features(self, samples: list[dict]): batch[key] = sample[key] hidden = sample.get("hidden_states") if hidden is None: - hidden_chunks = sample.get("hidden_states_ref_chunks") - if hidden_chunks: - expected_rows = None - hidden_positions = batch.get("hidden_positions") - if torch.is_tensor(hidden_positions): - expected_rows = int(hidden_positions.numel()) - hidden = _resolve_hidden_state_chunks(hidden_chunks, expected_rows=expected_rows) + tq_key = sample.get("hidden_states_tq_key") + if tq_key is not None and self._speco_tq_enabled: + # P0: hidden states were offloaded to TransferQueue by the + # rollout server; fetch by key (and free the storage). + from verl_speco.integration.transferqueue_bridge import get_sample + + payload = get_sample(tq_key) + hidden = payload.get("hidden_states") else: - hidden = _resolve_ray_object_ref(sample.get("hidden_states_ref")) + hidden_chunks = sample.get("hidden_states_ref_chunks") + if hidden_chunks: + expected_rows = None + hidden_positions = batch.get("hidden_positions") + if torch.is_tensor(hidden_positions): + expected_rows = int(hidden_positions.numel()) + hidden = _resolve_hidden_state_chunks(hidden_chunks, expected_rows=expected_rows) + else: + hidden = _resolve_ray_object_ref(sample.get("hidden_states_ref")) target_logprobs = sample.get("target_logprobs") if target_logprobs is None: target_logprobs = _resolve_ray_object_ref(sample.get("target_logprobs_ref")) From a83276e717f56d375a2fc007000168316c0fd1da Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Tue, 28 Jul 2026 14:21:42 +0800 Subject: [PATCH 21/50] Add token replay for standalone draft training --- tests/config/test_speco_config_overlay.py | 14 + tests/unit/test_draft_feature_store.py | 42 ++ tests/unit/test_target_feature_replay.py | 67 ++ verl_speco/config/draft_trainer.yaml | 13 + verl_speco/inspect_feature_store.py | 57 +- verl_speco/trainer/draft_dataset.py | 6 +- verl_speco/trainer/draft_training_loop.py | 43 +- verl_speco/trainer/feature_store.py | 212 +++++- verl_speco/trainer/target_feature_replay.py | 718 ++++++++++++++++++++ verl_speco/workers/speco_worker.py | 208 ++++-- 10 files changed, 1323 insertions(+), 57 deletions(-) create mode 100644 tests/unit/test_target_feature_replay.py create mode 100644 verl_speco/trainer/target_feature_replay.py diff --git a/tests/config/test_speco_config_overlay.py b/tests/config/test_speco_config_overlay.py index 598537de..9039fa67 100644 --- a/tests/config/test_speco_config_overlay.py +++ b/tests/config/test_speco_config_overlay.py @@ -60,6 +60,8 @@ def _copy_overlay_configs( def test_overlay_has_expected_default_drafter_shape() -> None: raw = OmegaConf.load(CONFIG_DIR / "speco_base.yaml") drafter = raw.actor_rollout_ref.rollout.drafter + standalone = OmegaConf.load(CONFIG_DIR / "draft_trainer.yaml") + standalone_training = standalone.actor_rollout_ref.rollout.drafter.training assert raw.speco.verl_base.version == "0.8.0" assert raw.speco.verl_base.branch == "release/v0.8.0" @@ -76,6 +78,11 @@ def test_overlay_has_expected_default_drafter_shape() -> None: assert drafter.training.warmup_style is None assert drafter.training.resume_trainer_state_from_checkpoint is True assert drafter.training.eagle1_num_hidden_layers == 1 + assert drafter.training.mode == "online" + assert drafter.training.feature_store.type == "torch_shard" + assert "target_feature_replay" not in drafter.training + assert standalone_training.target_feature_replay.cache.enabled is False + assert standalone_training.target_feature_replay.cache.max_size_gb == 0 def test_overlay_composes_with_release_upstream_verl(tmp_path: Path) -> None: @@ -119,6 +126,13 @@ def test_draft_trainer_composes_as_primary_config(tmp_path: Path) -> None: config = compose(config_name="draft_trainer") assert config.actor_rollout_ref.rollout.drafter.training.mode == "offline" + assert config.actor_rollout_ref.rollout.drafter.training.feature_store.type == ( + "torch_shard" + ) + assert ( + config.actor_rollout_ref.rollout.drafter.training.target_feature_replay.cache.enabled + is False + ) assert config.speco.draft_training.enable is True assert "trainer" in config assert "algorithm" in config diff --git a/tests/unit/test_draft_feature_store.py b/tests/unit/test_draft_feature_store.py index 071f0696..0cfdee9a 100644 --- a/tests/unit/test_draft_feature_store.py +++ b/tests/unit/test_draft_feature_store.py @@ -22,6 +22,8 @@ DraftFeatureDataLoader = draft_dataset.DraftFeatureDataLoader DraftFeatureDataLoaderConfig = draft_dataset.DraftFeatureDataLoaderConfig DraftFeatureSample = feature_store.DraftFeatureSample +DraftReplaySample = feature_store.DraftReplaySample +TokenReplayFeatureStore = feature_store.TokenReplayFeatureStore TorchShardFeatureStore = feature_store.TorchShardFeatureStore @@ -61,6 +63,46 @@ def test_torch_shard_feature_store_roundtrip(tmp_path): assert reader.get_metadata()["num_samples"] == 2 +def test_token_replay_feature_store_roundtrip(tmp_path): + sample = DraftReplaySample( + algorithm="DSPARK", + input_ids=torch.arange(12, dtype=torch.long), + loss_mask=torch.ones(12, dtype=torch.float32), + attention_mask=torch.ones(12, dtype=torch.bool), + position_ids=torch.arange(12, dtype=torch.long), + feature_positions=torch.arange(4, 10, dtype=torch.long), + draft_position_ids=torch.arange(5, 11, dtype=torch.long), + metadata={"target_model_path": "/target", "global_step": 3}, + ) + store = TokenReplayFeatureStore(tmp_path, max_samples_per_shard=1) + store.write_many([sample]) + store.close() + + reader = TokenReplayFeatureStore(tmp_path, read_only=True) + loaded = reader.read(next(reader.iter_keys())) + + assert loaded.algorithm == "DSPARK" + assert loaded.input_ids.dtype == torch.int32 + assert loaded.attention_mask.dtype == torch.bool + assert torch.equal(loaded.feature_positions, torch.arange(4, 10, dtype=torch.int32)) + assert "hidden_states" not in loaded.to_dict() + assert reader.get_metadata()["format"] == "token_replay" + + +def test_token_replay_rejects_non_contiguous_feature_positions(): + sample = DraftReplaySample( + input_ids=torch.arange(8), + loss_mask=torch.ones(8), + attention_mask=torch.ones(8, dtype=torch.bool), + position_ids=torch.arange(8), + feature_positions=torch.tensor([2, 4]), + draft_position_ids=torch.tensor([3, 5]), + ) + + with pytest.raises(ValueError, match="contiguous"): + sample.validate(strict=True) + + def test_feature_sample_normalizes_singleton_position_ids(): sample = DraftFeatureSample( input_ids=torch.tensor([1, 2, 3, 4], dtype=torch.long), diff --git a/tests/unit/test_target_feature_replay.py b/tests/unit/test_target_feature_replay.py new file mode 100644 index 00000000..9e7d9c6b --- /dev/null +++ b/tests/unit/test_target_feature_replay.py @@ -0,0 +1,67 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import pytest + +torch = pytest.importorskip("torch") + +from verl_speco.trainer.feature_store import DraftFeatureSample # noqa: E402 +from verl_speco.trainer.target_feature_replay import ( # noqa: E402 + BoundedReplayCache, + _hidden_capture_target, +) + + +def _feature_sample() -> DraftFeatureSample: + return DraftFeatureSample( + input_ids=torch.arange(8), + loss_mask=torch.ones(8), + hidden_states=torch.zeros(8, 16), + position_ids=torch.arange(1, 9), + ) + + +def test_hidden_capture_target_matches_transformers_hidden_state_indices(): + assert _hidden_capture_target(0, 36) == ("layer", 0) + assert _hidden_capture_target(34, 36) == ("layer", 34) + assert _hidden_capture_target(35, 36) == ("final", None) + + +def test_bounded_replay_cache_roundtrip(tmp_path): + cache = BoundedReplayCache( + tmp_path, + max_size_gb=0.01, + rank=1, + world_size=2, + ) + + assert cache.put("sample", _feature_sample()) is True + loaded = cache.get("sample") + + assert loaded is not None + assert torch.equal(loaded.input_ids, torch.arange(8)) + assert cache.metrics()["replay/cache_size_gb"] > 0 + assert (tmp_path / "rank00001" / "sample.pt").exists() + + +def test_bounded_replay_cache_disables_zero_budget(tmp_path): + cache = BoundedReplayCache( + tmp_path, + max_size_gb=0, + rank=0, + world_size=1, + ) + + assert cache.put("sample", _feature_sample()) is False + assert cache.get("sample") is None diff --git a/verl_speco/config/draft_trainer.yaml b/verl_speco/config/draft_trainer.yaml index c536d352..62f80853 100644 --- a/verl_speco/config/draft_trainer.yaml +++ b/verl_speco/config/draft_trainer.yaml @@ -44,3 +44,16 @@ actor_rollout_ref: master_addr: ${speco.draft_training.master_addr} master_port: ${speco.draft_training.master_port} standalone: ${speco.draft_training.standalone} + # Standalone/offline-only reconstruction for compact token_replay stores. + # The frozen target model is loaded lazily after the draft trainer starts. + target_feature_replay: + model_path: null + target_revision: null + dtype: bfloat16 + trust_remote_code: false + strict_target_model_path: false + logits_chunk_rows: 32 + cache: + enabled: false + path: null + max_size_gb: 0 diff --git a/verl_speco/inspect_feature_store.py b/verl_speco/inspect_feature_store.py index 4c09cba9..3e841ddf 100644 --- a/verl_speco/inspect_feature_store.py +++ b/verl_speco/inspect_feature_store.py @@ -28,7 +28,7 @@ def main() -> int: parser = argparse.ArgumentParser( - description="Inspect a torch_shard draft feature store." + description="Inspect a torch_shard or token_replay draft feature store." ) parser.add_argument( "path", @@ -133,15 +133,21 @@ def _sample_summary(sample: dict[str, Any]) -> dict[str, str]: keys = [ "input_ids", "loss_mask", + "attention_mask", "hidden_states", "last_hidden_states", "target_logprobs", "position_ids", + "feature_positions", + "draft_position_ids", ] return {key: _shape(sample[key]) for key in keys if key in sample} def _sample_issues(sample: dict[str, Any]) -> list[str]: + if sample.get("sample_type") == "token_replay" or "feature_positions" in sample: + return _replay_sample_issues(sample) + issues: list[str] = [] input_ids = _tensor(sample.get("input_ids"), "input_ids", issues) loss_mask = _tensor(sample.get("loss_mask"), "loss_mask", issues) @@ -232,5 +238,54 @@ def _sample_issues(sample: dict[str, Any]) -> list[str]: return issues +def _replay_sample_issues(sample: dict[str, Any]) -> list[str]: + issues: list[str] = [] + input_ids = _tensor(sample.get("input_ids"), "input_ids", issues) + loss_mask = _tensor(sample.get("loss_mask"), "loss_mask", issues) + attention_mask = _tensor(sample.get("attention_mask"), "attention_mask", issues) + position_ids = _tensor(sample.get("position_ids"), "position_ids", issues) + feature_positions = _tensor( + sample.get("feature_positions"), "feature_positions", issues + ) + draft_position_ids = _tensor( + sample.get("draft_position_ids"), "draft_position_ids", issues + ) + if input_ids is None: + return issues + + seq_len = int(input_ids.numel()) + for name, value in ( + ("loss_mask", loss_mask), + ("attention_mask", attention_mask), + ("position_ids", position_ids), + ): + if value is not None and int(value.numel()) != seq_len: + issues.append( + f"{name} length {int(value.numel())} does not match " + f"input_ids length {seq_len}" + ) + if feature_positions is None or draft_position_ids is None: + return issues + feature_positions = feature_positions.reshape(-1).long() + draft_position_ids = draft_position_ids.reshape(-1).long() + if int(feature_positions.numel()) == 0: + issues.append("feature_positions is empty") + return issues + if int(draft_position_ids.numel()) != int(feature_positions.numel()): + issues.append( + "draft_position_ids length does not match feature_positions length" + ) + if ( + int(feature_positions.min().item()) < 0 + or int(feature_positions.max().item()) >= seq_len + ): + issues.append(f"feature_positions fall outside input_ids length {seq_len}") + if int(feature_positions.numel()) > 1 and not bool( + torch.all(feature_positions[1:] == feature_positions[:-1] + 1).item() + ): + issues.append("feature_positions are not contiguous and increasing") + return issues + + if __name__ == "__main__": raise SystemExit(main()) diff --git a/verl_speco/trainer/draft_dataset.py b/verl_speco/trainer/draft_dataset.py index 80841c0b..923cec4f 100644 --- a/verl_speco/trainer/draft_dataset.py +++ b/verl_speco/trainer/draft_dataset.py @@ -18,7 +18,7 @@ from dataclasses import dataclass from typing import Iterator -from verl_speco.trainer.feature_store import DraftFeatureSample, DraftFeatureStore +from verl_speco.trainer.feature_store import DraftFeatureStore, DraftStoredSample @dataclass(frozen=True) @@ -52,7 +52,7 @@ def __init__(self, store: DraftFeatureStore, config: DraftFeatureDataLoaderConfi f"Invalid rank/world_size configuration: rank={rank}, world_size={world_size}" ) - def __iter__(self) -> Iterator[list[DraftFeatureSample]]: + def __iter__(self) -> Iterator[list[DraftStoredSample]]: epoch = 0 while True: keys = list( @@ -68,7 +68,7 @@ def __iter__(self) -> Iterator[list[DraftFeatureSample]]: rank_keys = keys[rank::world_size] if world_size > 1: rank_keys = rank_keys[: len(keys) // world_size] - batch: list[DraftFeatureSample] = [] + batch: list[DraftStoredSample] = [] for key in rank_keys: batch.append(self.store.read(key)) if len(batch) >= int(self.config.batch_size): diff --git a/verl_speco/trainer/draft_training_loop.py b/verl_speco/trainer/draft_training_loop.py index b0d965cf..7ce56ef3 100644 --- a/verl_speco/trainer/draft_training_loop.py +++ b/verl_speco/trainer/draft_training_loop.py @@ -52,10 +52,23 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: drafter_cfg = draft_config.rollout.drafter training_cfg = drafter_cfg.training feature_store_cfg = training_cfg.feature_store + feature_store_type = ( + str(feature_store_cfg.get("type", "torch_shard") or "torch_shard") + .strip() + .lower() + ) + training_mode = ( + str(training_cfg.get("mode", "offline") or "offline").strip().lower() + ) if not feature_store_cfg.get("path"): raise ValueError( "actor_rollout_ref.rollout.drafter.training.feature_store.path is required" ) + if feature_store_type == "token_replay" and training_mode != "offline": + raise ValueError( + "feature_store.type=token_replay is supported only by standalone " + "training.mode=offline" + ) _disable_standalone_sequence_parallel(draft_config) _configure_device(local_rank) @@ -89,6 +102,7 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: last_save_result: dict[str, Any] | None = None last_saved_step = 0 store = None + feature_replayer = None try: activated = await trainer.activate_training_model() if not activated: @@ -100,6 +114,20 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: last_saved_step = optimizer_step store = build_feature_store_from_config(feature_store_cfg, read_only=True) + if feature_store_type == "token_replay": + # Keep the large target model entirely outside online training imports + # and lifetime. The standalone loop materializes ordinary feature + # samples before handing them to the shared trainer. + from verl_speco.trainer.target_feature_replay import ( + TargetFeatureReplayer, + ) + + feature_replayer = TargetFeatureReplayer( + config, + rank=rank, + world_size=world_size, + device=trainer.runtime_device, + ) loader = DraftFeatureDataLoader( store, DraftFeatureDataLoaderConfig( @@ -116,8 +144,13 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: break step_started = time.perf_counter() attempted_batches += 1 + materialized_samples = ( + feature_replayer.materialize(samples) + if feature_replayer is not None + else samples + ) batch = trainer.prepare_training_batch_from_samples( - cast(list[Any], samples), + cast(list[Any], materialized_samples), step=optimizer_step, ) has_batch = batch is not None @@ -143,6 +176,8 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: attempted_batches=attempted_batches, step_elapsed_sec=time.perf_counter() - step_started, ) + if feature_replayer is not None: + step_metrics.update(feature_replayer.metrics()) _log_standalone_step_metrics(step_metrics, rank=rank) if save_interval > 0 and optimizer_step % save_interval == 0: last_save_result = _save_standalone_checkpoint(trainer, optimizer_step) @@ -158,6 +193,8 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: finally: if store is not None: store.close() + if feature_replayer is not None: + feature_replayer.close() await trainer.cleanup_training(clear_data=True) if dist.is_initialized(): dist.barrier() @@ -731,13 +768,15 @@ def _log_standalone_step_metrics(metrics: dict[str, float], *, rank: int) -> Non ("train/simulated_acc_len", "sim_acc_len"), ("train/lr", "lr"), ("perf/step_time", "step_time"), + ("replay/cache_hit_ratio", "cache_hit"), + ("replay/target_forward_time_total", "target_forward_total"), ): if key not in metrics: continue value = float(metrics[key]) if key == "train/lr": fields.append(f"{label}={value:.3e}") - elif key == "perf/step_time": + elif key in {"perf/step_time", "replay/target_forward_time_total"}: fields.append(f"{label}={value:.3f}s") else: fields.append(f"{label}={value:.4f}") diff --git a/verl_speco/trainer/feature_store.py b/verl_speco/trainer/feature_store.py index 2f500c9f..1bc34e78 100644 --- a/verl_speco/trainer/feature_store.py +++ b/verl_speco/trainer/feature_store.py @@ -202,12 +202,132 @@ def to_training_item(self) -> dict[str, Any]: return item +@dataclass +class DraftReplaySample: + """Compact token sample used to reconstruct target features offline.""" + + input_ids: torch.Tensor + loss_mask: torch.Tensor + attention_mask: torch.Tensor + position_ids: torch.Tensor + feature_positions: torch.Tensor + draft_position_ids: torch.Tensor + algorithm: str = "EAGLE3" + schema_version: int = SCHEMA_VERSION + metadata: dict[str, Any] = field(default_factory=dict) + + @classmethod + def from_dict( + cls, payload: dict[str, Any], *, strict: bool = True + ) -> "DraftReplaySample": + sample = cls( + schema_version=int(payload.get("schema_version", SCHEMA_VERSION)), + algorithm=str( + payload.get( + "algorithm", payload.get("metadata", {}).get("algorithm", "EAGLE3") + ) + ), + input_ids=payload["input_ids"], + loss_mask=payload["loss_mask"], + attention_mask=payload["attention_mask"], + position_ids=payload["position_ids"], + feature_positions=payload["feature_positions"], + draft_position_ids=payload["draft_position_ids"], + metadata=dict(payload.get("metadata") or {}), + ) + sample.validate(strict=strict) + return sample + + def validate(self, *, strict: bool = True) -> None: + if self.schema_version != SCHEMA_VERSION and strict: + raise ValueError( + f"Unsupported DraftReplaySample schema_version={self.schema_version}" + ) + tensor_fields = ( + "input_ids", + "loss_mask", + "attention_mask", + "position_ids", + "feature_positions", + "draft_position_ids", + ) + for name in tensor_fields: + value = getattr(self, name) + if not torch.is_tensor(value): + raise TypeError(f"DraftReplaySample.{name} must be a torch.Tensor") + if value.dim() > 1: + setattr(self, name, value.reshape(-1)) + + sequence_length = int(self.input_ids.numel()) + for name in ("loss_mask", "attention_mask", "position_ids"): + value = cast(torch.Tensor, getattr(self, name)) + if strict and int(value.numel()) != sequence_length: + raise ValueError( + f"DraftReplaySample input_ids/{name} length mismatch: " + f"{sequence_length} vs {int(value.numel())}" + ) + if strict and int(self.feature_positions.numel()) <= 0: + raise ValueError("DraftReplaySample.feature_positions must not be empty") + if strict and int(self.draft_position_ids.numel()) != int( + self.feature_positions.numel() + ): + raise ValueError( + "DraftReplaySample feature_positions/draft_position_ids length mismatch: " + f"{int(self.feature_positions.numel())} vs " + f"{int(self.draft_position_ids.numel())}" + ) + if int(self.feature_positions.numel()) > 0: + positions = self.feature_positions.detach().cpu().long() + if strict and ( + int(positions.min().item()) < 0 + or int(positions.max().item()) >= sequence_length + ): + raise ValueError( + "DraftReplaySample.feature_positions are outside input_ids: " + f"min={int(positions.min().item())} " + f"max={int(positions.max().item())} sequence_length={sequence_length}" + ) + if strict and int(positions.numel()) > 1: + deltas = positions[1:] - positions[:-1] + if not bool((deltas == 1).all().item()): + raise ValueError( + "DraftReplaySample.feature_positions must be contiguous and increasing" + ) + + def to_dict(self) -> dict[str, Any]: + self.validate(strict=False) + return { + "schema_version": self.schema_version, + "sample_type": "token_replay", + "algorithm": self.algorithm, + "input_ids": self.input_ids.detach().cpu().to(torch.int32).contiguous(), + "loss_mask": self.loss_mask.detach().cpu().to(torch.float16).contiguous(), + "attention_mask": self.attention_mask.detach().cpu().bool().contiguous(), + "position_ids": self.position_ids.detach() + .cpu() + .to(torch.int32) + .contiguous(), + "feature_positions": self.feature_positions.detach() + .cpu() + .to(torch.int32) + .contiguous(), + "draft_position_ids": self.draft_position_ids.detach() + .cpu() + .to(torch.int32) + .contiguous(), + "metadata": dict(self.metadata), + } + + +DraftStoredSample = DraftFeatureSample | DraftReplaySample + + class DraftFeatureStore(Protocol): def write_many( - self, samples: list[DraftFeatureSample | dict[str, Any]] + self, samples: list[DraftStoredSample | dict[str, Any]] ) -> list[str]: ... - def read(self, key: str) -> DraftFeatureSample: ... + def read(self, key: str) -> DraftStoredSample: ... def iter_keys(self, *, shuffle: bool = False, seed: int = 0) -> Iterator[str]: ... @@ -291,7 +411,7 @@ def __init__( self._write_metadata() def write_many( - self, samples: list[DraftFeatureSample | dict[str, Any]] + self, samples: list[DraftStoredSample | dict[str, Any]] ) -> list[str]: if self.read_only: raise RuntimeError("Cannot write to a read-only TorchShardFeatureStore") @@ -342,7 +462,7 @@ def flush_on_step(self, global_step: int | None, interval_steps: int) -> list[st return [] return self.flush() - def read(self, key: str) -> DraftFeatureSample: + def read(self, key: str) -> DraftStoredSample: shard_name, sample_index = _parse_key(key) shard = self._load_shard(shard_name) samples = shard.get("samples") or [] @@ -432,22 +552,83 @@ def _load_shard(self, shard_name: str) -> dict[str, Any]: return torch.load(path, map_location="cpu") +class TokenReplayFeatureStore(TorchShardFeatureStore): + """Compact shard store containing tokens and replay alignment metadata.""" + + def __init__( + self, + path: str | os.PathLike[str], + *, + max_samples_per_shard: int = 1024, + metadata: dict[str, Any] | None = None, + strict_schema: bool = True, + read_only: bool = False, + shard_prefix: str = "shard", + ): + replay_metadata = dict(metadata or {}) + replay_metadata["format"] = "token_replay" + super().__init__( + path, + max_samples_per_shard=max_samples_per_shard, + metadata=replay_metadata, + strict_schema=strict_schema, + read_only=read_only, + shard_prefix=shard_prefix, + ) + + def write_many( + self, samples: list[DraftStoredSample | dict[str, Any]] + ) -> list[str]: + if self.read_only: + raise RuntimeError("Cannot write to a read-only TokenReplayFeatureStore") + keys: list[str] = [] + for sample_like in samples: + sample = _coerce_replay_sample(sample_like, strict=self.strict_schema) + self._pending.append(sample.to_dict()) + keys.append(f"pending:{len(self._pending) - 1}") + if len(self._pending) >= self.max_samples_per_shard: + self.flush() + return keys + + def read(self, key: str) -> DraftReplaySample: + shard_name, sample_index = _parse_key(key) + shard = self._load_shard(shard_name) + samples = shard.get("samples") or [] + sample = samples[int(sample_index)] + return DraftReplaySample.from_dict(sample, strict=self.strict_schema) + + def build_feature_store_from_config( - feature_store_cfg, *, read_only: bool = False -) -> TorchShardFeatureStore: - store_type = str(feature_store_cfg.get("type", "torch_shard") or "torch_shard") - if store_type != "torch_shard": + feature_store_cfg, + *, + read_only: bool = False, + metadata: dict[str, Any] | None = None, + shard_prefix: str = "shard", +) -> DraftFeatureStore: + store_type = ( + str(feature_store_cfg.get("type", "torch_shard") or "torch_shard") + .strip() + .lower() + ) + store_cls: type[TorchShardFeatureStore] + if store_type == "torch_shard": + store_cls = TorchShardFeatureStore + elif store_type == "token_replay": + store_cls = TokenReplayFeatureStore + else: raise NotImplementedError(f"Unsupported draft feature store type: {store_type}") - return TorchShardFeatureStore( + return store_cls( feature_store_cfg.get("path"), max_samples_per_shard=int(feature_store_cfg.get("max_samples_per_shard", 1024)), + metadata=metadata, strict_schema=bool(feature_store_cfg.get("strict_schema", True)), read_only=read_only, + shard_prefix=shard_prefix, ) def _coerce_sample( - sample_like: DraftFeatureSample | dict[str, Any], *, strict: bool + sample_like: DraftStoredSample | dict[str, Any], *, strict: bool ) -> DraftFeatureSample: if isinstance(sample_like, DraftFeatureSample): sample_like.validate(strict=strict) @@ -457,6 +638,17 @@ def _coerce_sample( raise TypeError(f"Unsupported draft feature sample type: {type(sample_like)!r}") +def _coerce_replay_sample( + sample_like: DraftStoredSample | dict[str, Any], *, strict: bool +) -> DraftReplaySample: + if isinstance(sample_like, DraftReplaySample): + sample_like.validate(strict=strict) + return sample_like + if isinstance(sample_like, dict): + return DraftReplaySample.from_dict(sample_like, strict=strict) + raise TypeError(f"Unsupported draft replay sample type: {type(sample_like)!r}") + + def _cpu_tensor_tree(value: Any) -> Any: if torch.is_tensor(value): return value.detach().cpu().contiguous() diff --git a/verl_speco/trainer/target_feature_replay.py b/verl_speco/trainer/target_feature_replay.py new file mode 100644 index 00000000..7e9d8acd --- /dev/null +++ b/verl_speco/trainer/target_feature_replay.py @@ -0,0 +1,718 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Reconstruct standalone draft-training features from compact token samples.""" + +from __future__ import annotations + +import hashlib +import inspect +import json +import logging +import os +import tempfile +import time +from pathlib import Path +from typing import Any, Iterable, cast + +import torch +from torch import nn + +from verl_speco.integration.oldlogprob_layer_ids import ( + resolve_oldlogprob_aux_layer_ids, +) +from verl_speco.trainer.feature_store import DraftFeatureSample, DraftReplaySample + +logger = logging.getLogger(__name__) + + +def _config_value(config: Any, key: str, default: Any = None) -> Any: + if config is None: + return default + if hasattr(config, "get"): + return config.get(key, default) + return getattr(config, key, default) + + +def _parse_dtype(value: Any) -> torch.dtype: + normalized = str(value or "bfloat16").strip().lower() + dtypes = { + "bf16": torch.bfloat16, + "bfloat16": torch.bfloat16, + "fp16": torch.float16, + "float16": torch.float16, + "fp32": torch.float32, + "float32": torch.float32, + } + if normalized not in dtypes: + raise ValueError( + f"Unsupported target feature replay dtype {value!r}; " + "expected bfloat16, float16, or float32" + ) + return dtypes[normalized] + + +def _tensor_from_module_output(value: Any) -> torch.Tensor: + if torch.is_tensor(value): + return cast(torch.Tensor, value) + if isinstance(value, (tuple, list)) and value and torch.is_tensor(value[0]): + return cast(torch.Tensor, value[0]) + last_hidden_state = getattr(value, "last_hidden_state", None) + if torch.is_tensor(last_hidden_state): + return cast(torch.Tensor, last_hidden_state) + raise TypeError( + f"Target feature replay expected tensor-like module output, got {type(value)!r}" + ) + + +def _get_module_by_path(root: Any, path: str) -> Any: + current = root + for part in path.split("."): + if not part: + continue + current = getattr(current, part, None) + if current is None: + return None + return current + + +def _find_layers_and_final_norm(model: nn.Module) -> tuple[list[nn.Module], nn.Module]: + roots: list[nn.Module] = [model] + base_model = getattr(model, "base_model", None) + if isinstance(base_model, nn.Module) and base_model is not model: + roots.append(base_model) + + candidates = ( + ("model.layers", "model.norm"), + ("base_model.model.layers", "base_model.model.norm"), + ("model.decoder.layers", "model.decoder.final_layer_norm"), + ("transformer.h", "transformer.ln_f"), + ("gpt_neox.layers", "gpt_neox.final_layer_norm"), + ) + for root in roots: + for layers_path, norm_path in candidates: + layers = _get_module_by_path(root, layers_path) + norm = _get_module_by_path(root, norm_path) + if ( + isinstance(layers, (nn.ModuleList, list, tuple)) + and len(layers) > 0 + and isinstance(norm, nn.Module) + ): + return list(layers), norm + + for name, child in root.named_modules(): + if not isinstance(child, nn.ModuleList) or len(child) <= 0: + continue + if not (name.endswith("layers") or name.endswith("h")): + continue + parent_path = name.rsplit(".", 1)[0] if "." in name else "" + for norm_name in ("norm", "final_layer_norm", "ln_f"): + norm_path = f"{parent_path}.{norm_name}" if parent_path else norm_name + norm = _get_module_by_path(root, norm_path) + if isinstance(norm, nn.Module): + return list(child), norm + + raise RuntimeError( + "Target feature replay could not find transformer layers and final norm" + ) + + +def _hidden_capture_target(layer_id: int, num_layers: int) -> tuple[str, int | None]: + hidden_state_index = ( + int(layer_id) + 1 if int(layer_id) >= 0 else num_layers + 1 + int(layer_id) + ) + if hidden_state_index <= 0 or hidden_state_index > num_layers: + raise IndexError( + f"Target replay layer id {layer_id} resolved to hidden-state index " + f"{hidden_state_index}, but the model has {num_layers} layers" + ) + if hidden_state_index == num_layers: + return "final", None + return "layer", hidden_state_index - 1 + + +def _load_json_config(path: Any) -> dict[str, Any] | None: + if not path: + return None + config_path = os.path.join(os.fspath(path), "config.json") + try: + with open(config_path, encoding="utf-8") as config_file: + value = json.load(config_file) + except (OSError, json.JSONDecodeError): + return None + return value if isinstance(value, dict) else None + + +class BoundedReplayCache: + """Per-rank disk cache with a hard least-recently-used size budget.""" + + def __init__( + self, + path: str | os.PathLike[str], + *, + max_size_gb: float, + rank: int, + world_size: int, + ): + global_max_bytes = max(int(float(max_size_gb) * 1024**3), 0) + self.max_bytes = global_max_bytes // max(int(world_size), 1) + self.path = Path(path) / f"rank{int(rank):05d}" + self.path.mkdir(parents=True, exist_ok=True) + self._entries: dict[Path, tuple[int, float]] = {} + self._total_bytes = 0 + self._scan() + + @property + def enabled(self) -> bool: + return self.max_bytes > 0 + + def _scan(self) -> None: + self._entries = {} + self._total_bytes = 0 + for path in self.path.glob("*.pt"): + try: + stat = path.stat() + except OSError: + continue + size = int(stat.st_size) + self._entries[path] = (size, float(stat.st_mtime)) + self._total_bytes += size + + def get(self, key: str) -> DraftFeatureSample | None: + if not self.enabled: + return None + path = self.path / f"{key}.pt" + if not path.exists(): + return None + try: + try: + payload = torch.load(path, map_location="cpu", weights_only=False) + except TypeError: + payload = torch.load(path, map_location="cpu") + sample = DraftFeatureSample.from_dict(payload, strict=True) + now = time.time() + os.utime(path, (now, now)) + size = int(path.stat().st_size) + self._entries[path] = (size, now) + return sample + except Exception as exc: # noqa: BLE001 + logger.warning( + "Discard invalid target replay cache entry %s: %s", path, exc + ) + self._forget(path) + try: + path.unlink() + except OSError: + pass + return None + + def put(self, key: str, sample: DraftFeatureSample) -> bool: + if not self.enabled: + return False + path = self.path / f"{key}.pt" + if path.exists(): + return True + with tempfile.NamedTemporaryFile( + prefix=path.name, + suffix=".tmp", + dir=self.path, + delete=False, + ) as tmp_file: + tmp_path = Path(tmp_file.name) + try: + torch.save(sample.to_dict(), tmp_path) + size = int(tmp_path.stat().st_size) + if size > self.max_bytes: + return False + self._evict_until_fits(size) + os.replace(tmp_path, path) + now = time.time() + self._entries[path] = (size, now) + self._total_bytes += size + return True + except Exception as exc: # noqa: BLE001 + logger.warning( + "Failed to write target replay cache entry %s: %s", path, exc + ) + return False + finally: + if tmp_path.exists(): + try: + tmp_path.unlink() + except OSError: + pass + + def _evict_until_fits(self, incoming_size: int) -> None: + entries = sorted(self._entries.items(), key=lambda item: item[1][1]) + for path, _ in entries: + if self._total_bytes + incoming_size <= self.max_bytes: + break + try: + path.unlink() + except OSError: + continue + self._forget(path) + + def _forget(self, path: Path) -> None: + previous = self._entries.pop(path, None) + if previous is not None: + self._total_bytes = max(self._total_bytes - int(previous[0]), 0) + + def metrics(self) -> dict[str, float]: + return { + "replay/cache_size_gb": self._total_bytes / float(1024**3), + "replay/cache_budget_gb_per_rank": self.max_bytes / float(1024**3), + } + + +class TargetFeatureReplayer: + """Materialize target hidden states only for standalone token replay.""" + + def __init__( + self, + config: Any, + *, + rank: int, + world_size: int, + device: torch.device, + ): + self.config = config + self.rank = int(rank) + self.world_size = int(world_size) + self.device = torch.device(device) + self.draft_config = config.actor_rollout_ref + self.drafter_cfg = self.draft_config.rollout.drafter + self.training_cfg = self.drafter_cfg.training + self.replay_cfg = self.training_cfg.get("target_feature_replay", {}) or {} + configured_model_path = _config_value(self.replay_cfg, "model_path", None) + model_path = configured_model_path or self.draft_config.model.path + if not model_path: + raise ValueError( + "Token replay requires target_feature_replay.model_path or " + "actor_rollout_ref.model.path" + ) + self.model_path = os.fspath(model_path) + self.target_revision = str( + _config_value(self.replay_cfg, "target_revision", None) or self.model_path + ) + self.dtype = _parse_dtype(_config_value(self.replay_cfg, "dtype", "bfloat16")) + self.trust_remote_code = bool( + _config_value(self.replay_cfg, "trust_remote_code", False) + ) + self.strict_target_model_path = bool( + _config_value(self.replay_cfg, "strict_target_model_path", False) + ) + self.algorithm = str(self.drafter_cfg.speculative_algorithm).upper() + if self.algorithm not in {"EAGLE3", "DFLASH", "DSPARK"}: + raise ValueError( + f"Token replay does not support drafter algorithm {self.algorithm!r}" + ) + self.use_logits = bool(self.training_cfg.get("use_logits", False)) + self.logits_topk = int(self.training_cfg.get("logits_topk", 128) or 128) + self.logits_chunk_rows = max( + int(_config_value(self.replay_cfg, "logits_chunk_rows", 32) or 32), 1 + ) + + from transformers import AutoConfig + + self.target_config = AutoConfig.from_pretrained( + self.model_path, + trust_remote_code=self.trust_remote_code, + ) + self.target_num_hidden_layers = int( + getattr( + getattr(self.target_config, "text_config", self.target_config), + "num_hidden_layers", + ) + ) + model_configs = [ + value + for value in ( + _load_json_config(self.drafter_cfg.get("model_path", None)), + _load_json_config(self.drafter_cfg.get("checkpoint_path", None)), + ) + if value is not None + ] + layer_ids = resolve_oldlogprob_aux_layer_ids( + self.drafter_cfg, + target_num_hidden_layers=self.target_num_hidden_layers, + model_configs=model_configs, + ) + if not layer_ids: + raise RuntimeError( + "Token replay could not resolve target auxiliary layer ids" + ) + self.target_layer_ids = [int(layer_id) for layer_id in layer_ids] + dspark_l1_enabled = ( + self.algorithm == "DSPARK" + and float(self.training_cfg.get("dspark_l1_loss_alpha", 0.9) or 0.0) > 0 + ) + self.hidden_layout = ( + "dflash_aux_plus_last" + if dspark_l1_enabled + else "dflash_aux" + if self.algorithm in {"DFLASH", "DSPARK"} + else "eagle3_aux_plus_last" + ) + config_json = json.dumps( + self.target_config.to_dict(), sort_keys=True, default=str + ).encode() + self.target_config_fingerprint = hashlib.sha256(config_json).hexdigest() + + self.cache: BoundedReplayCache | None = None + cache_cfg = _config_value(self.replay_cfg, "cache", {}) or {} + if bool(_config_value(cache_cfg, "enabled", False)): + cache_path = _config_value(cache_cfg, "path", None) + if not cache_path: + feature_path = os.fspath(self.training_cfg.feature_store.path) + cache_path = f"{feature_path}.hidden_cache" + self.cache = BoundedReplayCache( + cache_path, + max_size_gb=float(_config_value(cache_cfg, "max_size_gb", 0.0) or 0.0), + rank=self.rank, + world_size=self.world_size, + ) + + self.model: nn.Module | None = None + self.layers: list[nn.Module] = [] + self.final_norm: nn.Module | None = None + self.backbone: nn.Module | None = None + self.output_embedding: nn.Module | None = None + self.cache_hits = 0 + self.cache_misses = 0 + self.materialized_samples = 0 + self.target_forward_seconds = 0.0 + + def materialize( + self, samples: Iterable[DraftReplaySample | DraftFeatureSample] + ) -> list[DraftFeatureSample]: + materialized: list[DraftFeatureSample] = [] + for sample in samples: + if isinstance(sample, DraftFeatureSample): + materialized.append(sample) + continue + if not isinstance(sample, DraftReplaySample): + raise TypeError( + f"Target feature replay expected DraftReplaySample, got {type(sample)!r}" + ) + self._validate_target_path(sample) + key = self._cache_key(sample) + cached = self.cache.get(key) if self.cache is not None else None + if cached is not None: + self.cache_hits += 1 + materialized.append(cached) + continue + self.cache_misses += 1 + replayed = self._materialize_one(sample) + if self.cache is not None: + self.cache.put(key, replayed) + materialized.append(replayed) + self.materialized_samples += len(materialized) + return materialized + + def _validate_target_path(self, sample: DraftReplaySample) -> None: + if sample.algorithm.upper() != self.algorithm: + raise ValueError( + "Token replay algorithm mismatch: " + f"sample={sample.algorithm!r} training={self.algorithm!r}" + ) + collected_layer_ids = sample.metadata.get("target_layer_ids") + if collected_layer_ids is not None: + normalized_layer_ids = ( + [int(collected_layer_ids)] + if isinstance(collected_layer_ids, int) + else [int(value) for value in collected_layer_ids] + ) + if normalized_layer_ids != self.target_layer_ids: + raise ValueError( + "Token replay target layer mismatch: " + f"collected={normalized_layer_ids} " + f"replay={self.target_layer_ids}" + ) + collected_layout = sample.metadata.get("hidden_states_layout") + if collected_layout and str(collected_layout) != self.hidden_layout: + raise ValueError( + "Token replay hidden layout mismatch: " + f"collected={collected_layout!r} replay={self.hidden_layout!r}" + ) + if not self.strict_target_model_path: + return + collected_path = sample.metadata.get("target_model_path") + if collected_path and os.path.normpath( + os.fspath(collected_path) + ) != os.path.normpath(self.model_path): + raise ValueError( + "Token replay target model path mismatch: " + f"collected={collected_path!r} replay={self.model_path!r}" + ) + + def _cache_key(self, sample: DraftReplaySample) -> str: + digest = hashlib.sha256() + contract = { + "target_revision": self.target_revision, + "target_config": self.target_config_fingerprint, + "algorithm": self.algorithm, + "target_layer_ids": self.target_layer_ids, + "hidden_layout": self.hidden_layout, + "dtype": str(self.dtype), + "use_logits": self.use_logits, + "logits_topk": self.logits_topk, + } + digest.update(json.dumps(contract, sort_keys=True).encode()) + for tensor in ( + sample.input_ids, + sample.attention_mask, + sample.position_ids, + sample.feature_positions, + sample.draft_position_ids, + sample.loss_mask, + ): + contiguous = tensor.detach().cpu().contiguous() + digest.update(str(contiguous.dtype).encode()) + digest.update(str(tuple(contiguous.shape)).encode()) + digest.update(contiguous.numpy().tobytes()) + return digest.hexdigest() + + def _ensure_model(self) -> None: + if self.model is not None: + return + from transformers import AutoModelForCausalLM + + logger.warning( + "Loading frozen target model for standalone token replay: path=%s dtype=%s device=%s", + self.model_path, + self.dtype, + self.device, + ) + model = AutoModelForCausalLM.from_pretrained( + self.model_path, + torch_dtype=self.dtype, + trust_remote_code=self.trust_remote_code, + low_cpu_mem_usage=True, + ) + model.eval() + model.requires_grad_(False) + model.to(self.device) + self.layers, self.final_norm = _find_layers_and_final_norm(model) + base_model_prefix = str(getattr(model, "base_model_prefix", "") or "") + backbone = ( + getattr(model, base_model_prefix, None) if base_model_prefix else None + ) + self.backbone = backbone if isinstance(backbone, nn.Module) else model + output_embedding = model.get_output_embeddings() + self.output_embedding = ( + output_embedding if isinstance(output_embedding, nn.Module) else None + ) + self.model = model + + def _materialize_one(self, sample: DraftReplaySample) -> DraftFeatureSample: + self._ensure_model() + assert self.backbone is not None + assert self.final_norm is not None + + feature_positions = sample.feature_positions.detach().cpu().long() + feature_end = int(feature_positions[-1].item()) + 1 + input_ids = sample.input_ids[:feature_end].to( + self.device, dtype=torch.long, non_blocking=True + ) + attention_mask = sample.attention_mask[:feature_end].to( + self.device, dtype=torch.long, non_blocking=True + ) + position_ids = sample.position_ids[:feature_end].to( + self.device, dtype=torch.long, non_blocking=True + ) + captures: dict[str, torch.Tensor] = {} + handles = [] + + def capture(key: str): + def hook(_module, _inputs, output): + captures[key] = _tensor_from_module_output(output) + + return hook + + aux_keys: list[str] = [] + modules: dict[str, nn.Module] = {} + for layer_id in self.target_layer_ids: + kind, layer_index = _hidden_capture_target( + layer_id, self.target_num_hidden_layers + ) + if kind == "final": + key = "final" + module = self.final_norm + else: + assert layer_index is not None + key = f"layer:{layer_index}" + module = self.layers[layer_index] + aux_keys.append(key) + modules[key] = module + include_final = self.hidden_layout in { + "eagle3_aux_plus_last", + "dflash_aux_plus_last", + } + need_final = include_final or (self.algorithm == "EAGLE3" and self.use_logits) + if need_final: + modules["final"] = self.final_norm + for key, module in modules.items(): + handles.append(module.register_forward_hook(capture(key))) + + started = time.perf_counter() + try: + forward_kwargs = { + "input_ids": input_ids.unsqueeze(0), + "attention_mask": attention_mask.unsqueeze(0), + "position_ids": position_ids.unsqueeze(0), + "use_cache": False, + "return_dict": True, + } + forward_kwargs = _supported_forward_kwargs( + self.backbone.forward, forward_kwargs + ) + with torch.inference_mode(): + self.backbone(**forward_kwargs) + finally: + for handle in handles: + handle.remove() + self.target_forward_seconds += time.perf_counter() - started + + required_keys = list(aux_keys) + if need_final: + required_keys.append("final") + missing = [key for key in required_keys if key not in captures] + if missing: + raise RuntimeError( + f"Target feature replay missed hidden-state captures: {missing}" + ) + + device_positions = feature_positions.to(self.device) + hidden_parts = [ + captures[key].squeeze(0).index_select(0, device_positions) + for key in aux_keys + ] + selected_final = ( + captures["final"].squeeze(0).index_select(0, device_positions) + if need_final + else None + ) + if include_final: + assert selected_final is not None + hidden_parts.append(selected_final) + hidden_states = torch.cat(hidden_parts, dim=-1).to( + device="cpu", dtype=self.dtype + ) + + target_logprobs = None + if self.algorithm == "EAGLE3" and self.use_logits: + assert selected_final is not None + target_logprobs = self._build_sparse_target_logprobs(selected_final[:-1]) + + selected_input_ids = sample.input_ids.index_select(0, feature_positions).long() + selected_loss_mask = sample.loss_mask.index_select(0, feature_positions).float() + metadata = dict(sample.metadata) + feature_start = int(feature_positions[0].item()) + feature_end = int(feature_positions[-1].item()) + 1 + metadata.update( + { + "source": "token_replay", + "target_model_path": self.model_path, + "target_revision": self.target_revision, + "target_config_fingerprint": self.target_config_fingerprint, + "target_layer_ids": list(self.target_layer_ids), + "hidden_states_layout": self.hidden_layout, + "feature_start": feature_start, + "feature_end": feature_end, + "hidden_position_start": feature_start, + "hidden_position_end": feature_end, + "hidden_positions": feature_positions, + "sequence_length": int(selected_input_ids.numel()), + "full_sequence_length": int(sample.input_ids.numel()), + "use_logits": self.use_logits, + } + ) + if target_logprobs is not None: + metadata["target_logprobs_position_start"] = feature_start + 1 + metadata["target_logprobs_position_end"] = feature_end + + return DraftFeatureSample( + algorithm=self.algorithm, + input_ids=selected_input_ids, + loss_mask=selected_loss_mask, + hidden_states=hidden_states, + target_logprobs=target_logprobs, + position_ids=sample.draft_position_ids.long(), + metadata=metadata, + ) + + def _build_sparse_target_logprobs( + self, final_hidden_states: torch.Tensor + ) -> torch.Tensor: + if self.output_embedding is None: + raise RuntimeError( + "EAGLE3 token replay with use_logits=true requires target output embeddings" + ) + rows: list[torch.Tensor] = [] + topk = max(self.logits_topk, 1) + with torch.inference_mode(): + for start in range( + 0, int(final_hidden_states.size(0)), self.logits_chunk_rows + ): + hidden = final_hidden_states[start : start + self.logits_chunk_rows] + logits = self.output_embedding(hidden).float() + local_topk = min(topk, int(logits.size(-1))) + values, ids = logits.topk(local_topk, dim=-1) + values = values - torch.logsumexp(logits, dim=-1, keepdim=True) + rows.append( + torch.stack((values, ids.to(dtype=values.dtype)), dim=-1).cpu() + ) + if not rows: + return torch.empty(0, topk, 2, dtype=torch.float32) + return torch.cat(rows, dim=0).contiguous() + + def metrics(self) -> dict[str, float]: + metrics = { + "replay/cache_hits_total": float(self.cache_hits), + "replay/cache_misses_total": float(self.cache_misses), + "replay/materialized_samples_total": float(self.materialized_samples), + "replay/target_forward_time_total": float(self.target_forward_seconds), + } + total = self.cache_hits + self.cache_misses + if total > 0: + metrics["replay/cache_hit_ratio"] = self.cache_hits / float(total) + if self.cache is not None: + metrics.update(self.cache.metrics()) + return metrics + + def close(self) -> None: + if self.model is None: + return + try: + self.model.to("cpu") + except Exception: # noqa: BLE001 + pass + self.model = None + self.layers = [] + self.final_norm = None + self.backbone = None + self.output_embedding = None + + +def _supported_forward_kwargs(forward: Any, kwargs: dict[str, Any]) -> dict[str, Any]: + try: + signature = inspect.signature(forward) + except (TypeError, ValueError): + return kwargs + if any( + parameter.kind == inspect.Parameter.VAR_KEYWORD + for parameter in signature.parameters.values() + ): + return kwargs + return {key: value for key, value in kwargs.items() if key in signature.parameters} diff --git a/verl_speco/workers/speco_worker.py b/verl_speco/workers/speco_worker.py index 9ca71d5b..87a40506 100644 --- a/verl_speco/workers/speco_worker.py +++ b/verl_speco/workers/speco_worker.py @@ -42,7 +42,13 @@ initialize_global_process_group_ray, set_numa_affinity, ) -from verl_speco.trainer.feature_store import DraftFeatureSample, TorchShardFeatureStore +from verl_speco.trainer.feature_store import ( + DraftFeatureSample, + DraftReplaySample, + TokenReplayFeatureStore, + TorchShardFeatureStore, + build_feature_store_from_config, +) logger = logging.getLogger(__file__) logger.setLevel(os.getenv("VERL_LOGGING_LEVEL", "WARN")) @@ -323,6 +329,7 @@ def __init__( self.trainer: Any = None self.feature_writer: Optional[TorchShardFeatureStore] = None self.feature_writer_path: Optional[str] = None + self.feature_writer_type: Optional[str] = None self.last_global_step: Optional[int] = None self.last_trained_step: Optional[int] = None self.training_process_group = None @@ -498,7 +505,7 @@ def init_model(self): def _store_rollout_sample( self, batch: dict, - hidden_states: torch.Tensor, + hidden_states: Optional[torch.Tensor], target_logprobs: Optional[torch.Tensor] = None, ): if ( @@ -510,6 +517,10 @@ def _store_rollout_sample( if self._drafter_training_mode() == "collect_only": self._write_rollout_feature_sample(batch, hidden_states, target_logprobs) return + if hidden_states is None: + raise RuntimeError( + "Online drafter training requires collected hidden states" + ) self.trainer.collect_online_data(batch, hidden_states, target_logprobs) def _drafter_training_mode(self) -> str: @@ -528,33 +539,54 @@ def _get_feature_writer(self) -> Optional[TorchShardFeatureStore]: path = _config_str(feature_store_cfg.get("path", None)) if not path: return None - if self.feature_writer is not None and self.feature_writer_path == path: + store_type = str( + feature_store_cfg.get("type", "torch_shard") or "torch_shard" + ).lower() + if ( + self.feature_writer is not None + and self.feature_writer_path == path + and self.feature_writer_type == store_type + ): return self.feature_writer + if self.feature_writer is not None: + self.feature_writer.close() model_cfg = self.config.get("model", None) target_model_path = ( _config_str(model_cfg.get("path", None)) if model_cfg is not None else "" ) - self.feature_writer = TorchShardFeatureStore( - path, - max_samples_per_shard=int( - feature_store_cfg.get("max_samples_per_shard", 1024) + self.feature_writer = cast( + TorchShardFeatureStore, + build_feature_store_from_config( + feature_store_cfg, + read_only=False, + metadata={ + "algorithm": str( + self.config.rollout.drafter.speculative_algorithm + ).upper(), + "target_model_path": target_model_path, + "drafter_model_path": _config_str( + self.config.rollout.drafter.get("model_path", None) + ), + "source": "rl_collect_only", + }, + shard_prefix=f"rank{int(self.rank):05d}_pid{int(os.getpid())}", ), - strict_schema=bool(feature_store_cfg.get("strict_schema", True)), - metadata={ - "algorithm": str( - self.config.rollout.drafter.speculative_algorithm - ).upper(), - "target_model_path": target_model_path, - "drafter_model_path": _config_str( - self.config.rollout.drafter.get("model_path", None) - ), - "source": "rl_collect_only", - }, - shard_prefix=f"rank{int(self.rank):05d}_pid{int(os.getpid())}", ) self.feature_writer_path = path + self.feature_writer_type = store_type return self.feature_writer + def _uses_token_replay_store(self) -> bool: + feature_store_cfg = self.config.rollout.drafter.training.get( + "feature_store", {} + ) + return ( + str(feature_store_cfg.get("type", "torch_shard") or "torch_shard") + .strip() + .lower() + == "token_replay" + ) + def _build_rollout_loss_mask( self, batch: dict, input_ids: torch.Tensor ) -> torch.Tensor: @@ -591,7 +623,7 @@ def _build_rollout_loss_mask( def _write_rollout_feature_sample( self, batch: dict, - hidden_states: torch.Tensor, + hidden_states: Optional[torch.Tensor], target_logprobs: Optional[torch.Tensor], ) -> None: writer = self._get_feature_writer() @@ -603,18 +635,25 @@ def _write_rollout_feature_sample( return full_input_ids = batch["input_ids"].detach().cpu().reshape(-1) full_loss_mask = self._build_rollout_loss_mask(batch, full_input_ids) - hidden_states = hidden_states.detach().cpu() - hidden_rows = int( - hidden_states.size(1) - if hidden_states.dim() == 3 and hidden_states.size(0) == 1 - else hidden_states.size(0) - ) hidden_positions = batch.get("hidden_positions") if torch.is_tensor(hidden_positions): hidden_positions = cast(torch.Tensor, hidden_positions) hidden_positions = hidden_positions.detach().cpu().long().reshape(-1) else: hidden_positions = None + if hidden_states is not None: + hidden_states = hidden_states.detach().cpu() + hidden_rows = int( + hidden_states.size(1) + if hidden_states.dim() == 3 and hidden_states.size(0) == 1 + else hidden_states.size(0) + ) + elif isinstance(writer, TokenReplayFeatureStore): + hidden_rows = self._token_replay_hidden_rows(batch, hidden_positions) + else: + raise RuntimeError( + "torch_shard feature collection requires collected hidden states" + ) feature_start, feature_end, position_ids = self._resolve_rollout_feature_window( full_input_ids, hidden_rows, @@ -687,16 +726,98 @@ def _write_rollout_feature_sample( ): if key in batch: metadata[key] = batch[key] - sample = DraftFeatureSample( - algorithm=str(self.config.rollout.drafter.speculative_algorithm).upper(), - input_ids=input_ids, - loss_mask=loss_mask, - hidden_states=hidden_states, - target_logprobs=target_logprobs, - position_ids=position_ids, - metadata=metadata, - ) - writer.write_many([sample]) + if isinstance(writer, TokenReplayFeatureStore): + full_attention_mask = self._replay_sequence_tensor( + batch.get("attention_mask"), + full_input_ids, + name="attention_mask", + default=torch.ones_like(full_input_ids, dtype=torch.bool), + ).bool() + default_position_ids = ( + full_attention_mask.long().cumsum(dim=0).sub(1).clamp_min(0) + ) + full_position_ids = self._replay_sequence_tensor( + batch.get("position_ids"), + full_input_ids, + name="position_ids", + default=default_position_ids, + ).long() + replay_metadata = { + key: value + for key, value in metadata.items() + if key + not in { + "hidden_positions", + "hidden_last_hidden_logprob_check", + "hidden_raw_topk_logprob_check", + "hidden_last_hidden_filter", + "hidden_last_hidden_select", + "target_logprobs_position_start", + "target_logprobs_position_end", + } + } + replay_metadata["source"] = "token_replay" + replay_sample = DraftReplaySample( + algorithm=algorithm, + input_ids=full_input_ids, + loss_mask=full_loss_mask, + attention_mask=full_attention_mask, + position_ids=full_position_ids, + feature_positions=position_ids - 1, + draft_position_ids=position_ids, + metadata=replay_metadata, + ) + writer.write_many([replay_sample]) + else: + assert hidden_states is not None + feature_sample = DraftFeatureSample( + algorithm=algorithm, + input_ids=input_ids, + loss_mask=loss_mask, + hidden_states=hidden_states, + target_logprobs=target_logprobs, + position_ids=position_ids, + metadata=metadata, + ) + writer.write_many([feature_sample]) + + @staticmethod + def _token_replay_hidden_rows( + batch: dict, hidden_positions: Optional[torch.Tensor] + ) -> int: + if hidden_positions is not None and int(hidden_positions.numel()) > 0: + return int(hidden_positions.numel()) + try: + start = int(batch["hidden_position_start"]) + end = int(batch["hidden_position_end"]) + except (KeyError, TypeError, ValueError) as exc: + raise ValueError( + "token_replay collection without a hidden tensor requires " + "hidden_positions or hidden_position_start/hidden_position_end" + ) from exc + if end <= start: + raise ValueError( + f"Invalid token_replay hidden window: start={start} end={end}" + ) + return end - start + + @staticmethod + def _replay_sequence_tensor( + value: Any, + input_ids: torch.Tensor, + *, + name: str, + default: torch.Tensor, + ) -> torch.Tensor: + if not torch.is_tensor(value): + return default + tensor = cast(torch.Tensor, value).detach().cpu().reshape(-1) + if int(tensor.numel()) != int(input_ids.numel()): + raise ValueError( + f"token_replay {name} cannot be normalized to one value per token: " + f"numel={int(tensor.numel())} tokens={int(input_ids.numel())}" + ) + return tensor @staticmethod def _align_rollout_target_logprobs( @@ -805,6 +926,8 @@ def collect_rollout_features(self, samples: list[dict]): "responses": sample["responses"], } for key in ( + "attention_mask", + "position_ids", "hidden_position_start", "hidden_position_end", "hidden_positions", @@ -828,8 +951,9 @@ def collect_rollout_features(self, samples: list[dict]): ): if key in sample: batch[key] = sample[key] - hidden = sample.get("hidden_states") - if hidden is None: + skip_hidden_payload = self._uses_token_replay_store() + hidden = None if skip_hidden_payload else sample.get("hidden_states") + if hidden is None and not skip_hidden_payload: hidden_chunks = sample.get("hidden_states_ref_chunks") if hidden_chunks: expected_rows = None @@ -842,12 +966,14 @@ def collect_rollout_features(self, samples: list[dict]): ) else: hidden = _resolve_ray_object_ref(sample.get("hidden_states_ref")) - target_logprobs = sample.get("target_logprobs") - if target_logprobs is None: + target_logprobs = ( + None if skip_hidden_payload else sample.get("target_logprobs") + ) + if target_logprobs is None and not skip_hidden_payload: target_logprobs = _resolve_ray_object_ref( sample.get("target_logprobs_ref") ) - if hidden is None: + if hidden is None and not skip_hidden_payload: continue self._store_rollout_sample( batch=batch, From 0b7422e9a3879ecdabbd874e31dbcda0fa78d65c Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Wed, 29 Jul 2026 15:58:31 +0800 Subject: [PATCH 22/50] Add vLLM hidden-state replay and linear LR decay --- .../integration/test_drafter_lr_scheduler.py | 84 ++++++ verl_speco/backends/lr_scheduler.py | 55 ++++ verl_speco/config/draft_trainer.yaml | 13 + verl_speco/config/speco_base.yaml | 2 + verl_speco/trainer/feature_store.py | 176 +++++++++++ verl_speco/trainer/target_feature_replay.py | 281 ++++++++++++++++++ verl_speco/vllm_hidden_states_generate.py | 144 +++++++++ 7 files changed, 755 insertions(+) create mode 100644 verl_speco/vllm_hidden_states_generate.py diff --git a/tests/integration/test_drafter_lr_scheduler.py b/tests/integration/test_drafter_lr_scheduler.py index d5272e1a..b7d7a83f 100644 --- a/tests/integration/test_drafter_lr_scheduler.py +++ b/tests/integration/test_drafter_lr_scheduler.py @@ -21,6 +21,7 @@ from verl_speco.backends.lr_scheduler import ( # noqa: E402 ClampedGlobalCosineLR, + LinearWarmupDecayLR, build_drafter_lr_scheduler, ) @@ -92,6 +93,55 @@ def test_scheduler_builder_uses_configured_global_cosine_values() -> None: assert optimizer.param_groups[0]["lr"] == pytest.approx(5e-6) +def test_linear_warmup_decay_reaches_zero_after_decay_steps() -> None: + optimizer = _optimizer() + scheduler = LinearWarmupDecayLR( + optimizer, + decay_steps=100, + min_lr_ratio=0.0, + warmup_steps=10, + ) + + assert optimizer.param_groups[0]["lr"] == pytest.approx(0.0) + + _step(optimizer, scheduler, 5) + assert optimizer.param_groups[0]["lr"] == pytest.approx(5e-6) + + _step(optimizer, scheduler, 5) + assert optimizer.param_groups[0]["lr"] == pytest.approx(1e-5) + + _step(optimizer, scheduler, 45) + assert optimizer.param_groups[0]["lr"] == pytest.approx(5e-6) + + _step(optimizer, scheduler, 45) + assert optimizer.param_groups[0]["lr"] == pytest.approx(0.0) + + _step(optimizer, scheduler, 10) + assert optimizer.param_groups[0]["lr"] == pytest.approx(0.0) + + +def test_scheduler_builder_uses_linear_warmup_decay_values() -> None: + optimizer = _optimizer(lr=2e-5) + scheduler = build_drafter_lr_scheduler( + optimizer, + { + "lr_scheduler_type": "linear", + "lr_decay_steps": 100, + "min_lr_ratio": 0.1, + "lr_warmup_steps": 10, + }, + ) + + _step(optimizer, scheduler, 10) + assert optimizer.param_groups[0]["lr"] == pytest.approx(2e-5) + + _step(optimizer, scheduler, 45) + assert optimizer.param_groups[0]["lr"] == pytest.approx(1.1e-5) + + _step(optimizer, scheduler, 45) + assert optimizer.param_groups[0]["lr"] == pytest.approx(2e-6) + + def test_scheduler_builder_resumes_from_successful_optimizer_steps() -> None: optimizer = _optimizer() scheduler = build_drafter_lr_scheduler( @@ -114,6 +164,27 @@ def test_scheduler_builder_resumes_from_successful_optimizer_steps() -> None: assert optimizer.param_groups[0]["lr"] == pytest.approx(1e-5 * expected_ratio) +def test_linear_scheduler_builder_resumes_from_successful_optimizer_steps() -> None: + optimizer = _optimizer() + scheduler = build_drafter_lr_scheduler( + optimizer, + { + "lr_scheduler_type": "linear", + "lr_decay_steps": 100, + "min_lr_ratio": 0.0, + "lr_warmup_steps": 10, + "_resume_optimizer_steps": 55, + }, + ) + + assert scheduler.last_epoch == 55 + assert optimizer.param_groups[0]["lr"] == pytest.approx(5e-6) + + _step(optimizer, scheduler, 1) + assert scheduler.last_epoch == 56 + assert optimizer.param_groups[0]["lr"] == pytest.approx(4.888888888888889e-6) + + def test_scheduler_builder_does_not_replace_explicit_invalid_decay() -> None: with pytest.raises(ValueError, match="lr_decay_steps"): build_drafter_lr_scheduler( @@ -136,3 +207,16 @@ def test_scheduler_builder_does_not_replace_explicit_invalid_decay() -> None: def test_clamped_global_cosine_rejects_invalid_config(kwargs, message) -> None: with pytest.raises(ValueError, match=message): ClampedGlobalCosineLR(_optimizer(), **kwargs) + + +@pytest.mark.parametrize( + ("kwargs", "message"), + [ + ({"decay_steps": 0}, "lr_decay_steps"), + ({"decay_steps": 10, "warmup_steps": 10}, "lr_warmup_steps"), + ({"min_lr_ratio": -0.1}, "min_lr_ratio"), + ], +) +def test_linear_warmup_decay_rejects_invalid_config(kwargs, message) -> None: + with pytest.raises(ValueError, match=message): + LinearWarmupDecayLR(_optimizer(), **kwargs) diff --git a/verl_speco/backends/lr_scheduler.py b/verl_speco/backends/lr_scheduler.py index ed811035..2980885f 100644 --- a/verl_speco/backends/lr_scheduler.py +++ b/verl_speco/backends/lr_scheduler.py @@ -83,6 +83,51 @@ def get_lr(self) -> list[float]: return [base_lr * ratio for base_lr in self.base_lrs] +class LinearWarmupDecayLR(LRScheduler): + """Linear warmup followed by linear decay over successful optimizer steps.""" + + def __init__( + self, + optimizer: Optimizer, + *, + decay_steps: int, + min_lr_ratio: float = 0.0, + warmup_steps: int = 0, + last_epoch: int = -1, + ) -> None: + self.decay_steps = int(decay_steps) + self.min_lr_ratio = float(min_lr_ratio) + self.warmup_steps = int(warmup_steps) + if self.decay_steps <= 0: + raise ValueError(f"lr_decay_steps must be positive, got {self.decay_steps}") + if self.warmup_steps < 0: + raise ValueError( + f"lr_warmup_steps must be non-negative, got {self.warmup_steps}" + ) + if self.warmup_steps >= self.decay_steps: + raise ValueError( + "lr_warmup_steps must be smaller than lr_decay_steps, " + f"got warmup={self.warmup_steps}, decay={self.decay_steps}" + ) + if not 0.0 <= self.min_lr_ratio <= 1.0: + raise ValueError(f"min_lr_ratio must be in [0, 1], got {self.min_lr_ratio}") + super().__init__(optimizer, last_epoch=last_epoch) + + def _lr_ratio(self, step: int) -> float: + step = max(int(step), 0) + if self.warmup_steps > 0 and step < self.warmup_steps: + return float(step) / float(self.warmup_steps) + + decay_span = self.decay_steps - self.warmup_steps + progress = min(max(step - self.warmup_steps, 0) / decay_span, 1.0) + linear_ratio = 1.0 - progress + return self.min_lr_ratio + (1.0 - self.min_lr_ratio) * linear_ratio + + def get_lr(self) -> list[float]: + ratio = self._lr_ratio(self.last_epoch) + return [base_lr * ratio for base_lr in self.base_lrs] + + def build_drafter_lr_scheduler(optimizer: Optimizer, train_cfg: Any) -> LRScheduler: """Build a drafter scheduler while retaining legacy warmup_style overrides.""" @@ -116,6 +161,16 @@ def build_drafter_lr_scheduler(optimizer: Optimizer, train_cfg: Any) -> LRSchedu num_cycles=float(0.5 if num_cycles is None else num_cycles), last_epoch=last_epoch, ) + if scheduler_type in {"linear", "linear_decay"}: + decay_steps = train_cfg.get("lr_decay_steps", train_cfg.get("step", 0)) + min_lr_ratio = train_cfg.get("min_lr_ratio", 0.0) + return LinearWarmupDecayLR( + optimizer, + decay_steps=int(decay_steps or 0), + min_lr_ratio=float(0.0 if min_lr_ratio is None else min_lr_ratio), + warmup_steps=warmup_steps, + last_epoch=last_epoch, + ) if scheduler_type in {"global_cosine", "clamped_global_cosine"}: decay_steps = train_cfg.get("lr_decay_steps", 100) min_lr_ratio = train_cfg.get("min_lr_ratio", 0.1) diff --git a/verl_speco/config/draft_trainer.yaml b/verl_speco/config/draft_trainer.yaml index 62f80853..6ba60569 100644 --- a/verl_speco/config/draft_trainer.yaml +++ b/verl_speco/config/draft_trainer.yaml @@ -47,12 +47,25 @@ actor_rollout_ref: # Standalone/offline-only reconstruction for compact token_replay stores. # The frozen target model is loaded lazily after the draft trainer starts. target_feature_replay: + backend: torch model_path: null target_revision: null dtype: bfloat16 trust_remote_code: false strict_target_model_path: false logits_chunk_rows: 32 + vllm_endpoint: http://localhost:8000/v1 + vllm_model: null + request_timeout: 120 + max_retries: 3 + on_generate: delete + require_arange_positions: true + offline_generation: + input_path: null + output_path: null + max_samples: 0 + batch_size: 1 + shuffle: false cache: enabled: false path: null diff --git a/verl_speco/config/speco_base.yaml b/verl_speco/config/speco_base.yaml index f8eeb435..ef798f74 100644 --- a/verl_speco/config/speco_base.yaml +++ b/verl_speco/config/speco_base.yaml @@ -84,6 +84,8 @@ actor_rollout_ref: train_batches_per_cycle: 4 lr: 1e-5 lr_warmup_steps: 0 + # Supported standalone drafter scheduler types: + # constant, cosine, linear, global_cosine. lr_scheduler_type: constant lr_decay_steps: 100 min_lr_ratio: 0.1 diff --git a/verl_speco/trainer/feature_store.py b/verl_speco/trainer/feature_store.py index 1bc34e78..cea2d862 100644 --- a/verl_speco/trainer/feature_store.py +++ b/verl_speco/trainer/feature_store.py @@ -598,6 +598,162 @@ def read(self, key: str) -> DraftReplaySample: return DraftReplaySample.from_dict(sample, strict=self.strict_schema) +class VllmSafetensorsFeatureStore(TorchShardFeatureStore): + """Feature store for vLLM-extracted hidden states saved as safetensors. + + Each sample is stored in an individual ``.safetensors`` file with the + non-tensor schema/metadata recorded in the shared manifest. This mirrors + the vLLM/speculators hidden-state extraction flow while exposing the same + ``DraftFeatureStore`` interface used by standalone training. + """ + + def __init__( + self, + path: str | os.PathLike[str], + *, + max_samples_per_shard: int = 1024, + metadata: dict[str, Any] | None = None, + strict_schema: bool = True, + read_only: bool = False, + shard_prefix: str = "hs", + ): + safetensors_metadata = dict(metadata or {}) + safetensors_metadata["format"] = "vllm_safetensors" + super().__init__( + path, + max_samples_per_shard=max_samples_per_shard, + metadata=safetensors_metadata, + strict_schema=strict_schema, + read_only=read_only, + shard_prefix=shard_prefix, + ) + + def write_many( + self, samples: list[DraftStoredSample | dict[str, Any]] + ) -> list[str]: + if self.read_only: + raise RuntimeError( + "Cannot write to a read-only VllmSafetensorsFeatureStore" + ) + try: + from safetensors.torch import save_file + except ImportError as exc: + raise RuntimeError( + "feature_store.type=vllm_safetensors requires safetensors" + ) from exc + + keys: list[str] = [] + for sample_like in samples: + sample = _coerce_sample(sample_like, strict=self.strict_schema) + sample_index = self._infer_next_shard_index() + file_name = f"{self.shard_prefix}_{sample_index:06d}.safetensors" + file_path = self.path / file_name + tensor_payload = self._sample_to_safetensors(sample) + with tempfile.NamedTemporaryFile( + prefix=file_path.name, + suffix=".tmp", + dir=file_path.parent, + delete=False, + ) as tmp_file: + tmp_name = tmp_file.name + try: + save_file(tensor_payload, tmp_name) + os.replace(tmp_name, file_path) + finally: + if os.path.exists(tmp_name): + os.remove(tmp_name) + + entry = { + "path": file_name, + "num_samples": 1, + "num_tokens": _sample_token_count(sample.to_dict()), + "sample": self._sample_manifest(sample), + } + with self.manifest_path.open("a", encoding="utf-8") as manifest_file: + manifest_file.write( + json.dumps(entry, ensure_ascii=True, sort_keys=True) + "\n" + ) + self._manifest.append(entry) + self._next_shard_index = sample_index + 1 + keys.append(f"{file_name}:0") + return keys + + def flush(self) -> list[str]: + return [] + + def read(self, key: str) -> DraftFeatureSample: + try: + from safetensors.torch import load_file + except ImportError as exc: + raise RuntimeError( + "feature_store.type=vllm_safetensors requires safetensors" + ) from exc + + file_name, sample_index = _parse_key(key) + if int(sample_index) != 0: + raise ValueError( + f"Invalid vllm_safetensors key {key!r}; expected sample index 0" + ) + entries = { + str(entry.get("path")): entry for entry in self._load_manifest() + } + entry = entries.get(file_name) + if entry is None: + raise KeyError(f"Missing vllm_safetensors manifest entry for {file_name}") + tensors = load_file(str(self.path / file_name), device="cpu") + manifest_sample = dict(entry.get("sample") or {}) + payload: dict[str, Any] = { + "schema_version": int(manifest_sample.get("schema_version", SCHEMA_VERSION)), + "algorithm": manifest_sample.get("algorithm", "EAGLE3"), + "input_ids": tensors["input_ids"], + "loss_mask": tensors["loss_mask"], + "hidden_states": tensors["hidden_states"], + "metadata": dict(manifest_sample.get("metadata") or {}), + } + for optional_key in ( + "last_hidden_states", + "target", + "target_logprobs", + "position_ids", + ): + if optional_key in tensors: + payload[optional_key] = tensors[optional_key] + return DraftFeatureSample.from_dict(payload, strict=self.strict_schema) + + def iter_keys(self, *, shuffle: bool = False, seed: int = 0) -> Iterator[str]: + keys = [f"{entry['path']}:0" for entry in self._load_manifest()] + if shuffle: + random.Random(int(seed)).shuffle(keys) + yield from keys + + def _sample_to_safetensors( + self, sample: DraftFeatureSample + ) -> dict[str, torch.Tensor]: + payload = sample.to_dict() + tensors = { + "input_ids": payload["input_ids"].long().contiguous(), + "loss_mask": payload["loss_mask"].float().contiguous(), + "hidden_states": payload["hidden_states"].contiguous(), + } + for optional_key in ( + "last_hidden_states", + "target", + "target_logprobs", + "position_ids", + ): + value = payload.get(optional_key) + if torch.is_tensor(value): + tensors[optional_key] = value.contiguous() + return tensors + + def _sample_manifest(self, sample: DraftFeatureSample) -> dict[str, Any]: + return { + "schema_version": sample.schema_version, + "algorithm": sample.algorithm, + "metadata": _json_safe_metadata(sample.metadata), + } + + def build_feature_store_from_config( feature_store_cfg, *, @@ -615,6 +771,8 @@ def build_feature_store_from_config( store_cls = TorchShardFeatureStore elif store_type == "token_replay": store_cls = TokenReplayFeatureStore + elif store_type in {"vllm_safetensors", "safetensors"}: + store_cls = VllmSafetensorsFeatureStore else: raise NotImplementedError(f"Unsupported draft feature store type: {store_type}") return store_cls( @@ -710,3 +868,21 @@ def _atomic_torch_save(payload: dict[str, Any], path: Path) -> None: finally: if os.path.exists(tmp_name): os.remove(tmp_name) + + +def _json_safe_metadata(value: Any) -> Any: + if torch.is_tensor(value): + if value.numel() <= 128: + return value.detach().cpu().tolist() + return { + "__tensor__": True, + "shape": list(value.shape), + "dtype": str(value.dtype), + } + if isinstance(value, dict): + return {str(key): _json_safe_metadata(item) for key, item in value.items()} + if isinstance(value, (list, tuple)): + return [_json_safe_metadata(item) for item in value] + if isinstance(value, (str, int, float, bool)) or value is None: + return value + return str(value) diff --git a/verl_speco/trainer/target_feature_replay.py b/verl_speco/trainer/target_feature_replay.py index 7e9d8acd..017bfcb2 100644 --- a/verl_speco/trainer/target_feature_replay.py +++ b/verl_speco/trainer/target_feature_replay.py @@ -153,6 +153,42 @@ def _load_json_config(path: Any) -> dict[str, Any] | None: return value if isinstance(value, dict) else None +def _wait_for_lock(lock_path: Path, timeout: float = 30.0) -> None: + if not lock_path.exists(): + return + try: + import fcntl + except ImportError: + deadline = time.monotonic() + float(timeout) + while lock_path.exists(): + if time.monotonic() >= deadline: + raise TimeoutError( + f"Timed out waiting for hidden-states lock: {lock_path}" + ) + time.sleep(0.1) + return + + fd = os.open(lock_path, os.O_RDONLY) + try: + deadline = time.monotonic() + float(timeout) + while True: + try: + fcntl.flock(fd, fcntl.LOCK_EX | fcntl.LOCK_NB) + break + except BlockingIOError: + if time.monotonic() >= deadline: + raise TimeoutError( + f"Timed out waiting for hidden-states lock: {lock_path}" + ) from None + time.sleep(0.1) + finally: + os.close(fd) + try: + lock_path.unlink() + except OSError: + pass + + class BoundedReplayCache: """Per-rank disk cache with a hard least-recently-used size budget.""" @@ -294,6 +330,16 @@ def __init__( self.drafter_cfg = self.draft_config.rollout.drafter self.training_cfg = self.drafter_cfg.training self.replay_cfg = self.training_cfg.get("target_feature_replay", {}) or {} + self.backend = ( + str(_config_value(self.replay_cfg, "backend", "torch") or "torch") + .strip() + .lower() + ) + if self.backend not in {"torch", "vllm_file"}: + raise ValueError( + f"Unsupported target_feature_replay.backend={self.backend!r}; " + "expected 'torch' or 'vllm_file'" + ) configured_model_path = _config_value(self.replay_cfg, "model_path", None) model_path = configured_model_path or self.draft_config.model.path if not model_path: @@ -322,6 +368,29 @@ def __init__( self.logits_chunk_rows = max( int(_config_value(self.replay_cfg, "logits_chunk_rows", 32) or 32), 1 ) + self.vllm_endpoint = str( + _config_value(self.replay_cfg, "vllm_endpoint", "http://localhost:8000/v1") + or "http://localhost:8000/v1" + ) + self.vllm_model = _config_value(self.replay_cfg, "vllm_model", None) + self.vllm_timeout = float( + _config_value(self.replay_cfg, "request_timeout", 120.0) or 120.0 + ) + self.vllm_max_retries = max( + int(_config_value(self.replay_cfg, "max_retries", 3) or 0), 0 + ) + self.vllm_on_generate = ( + str(_config_value(self.replay_cfg, "on_generate", "delete") or "delete") + .strip() + .lower() + ) + if self.vllm_on_generate not in {"delete", "keep"}: + raise ValueError( + "target_feature_replay.on_generate must be 'delete' or 'keep'" + ) + self.vllm_require_arange_positions = bool( + _config_value(self.replay_cfg, "require_arange_positions", True) + ) from transformers import AutoConfig @@ -388,10 +457,14 @@ def __init__( self.final_norm: nn.Module | None = None self.backbone: nn.Module | None = None self.output_embedding: nn.Module | None = None + self.vllm_client: Any | None = None + self.vllm_resolved_model: str | None = None self.cache_hits = 0 self.cache_misses = 0 self.materialized_samples = 0 self.target_forward_seconds = 0.0 + self.vllm_request_seconds = 0.0 + self.vllm_requests = 0 def materialize( self, samples: Iterable[DraftReplaySample | DraftFeatureSample] @@ -516,6 +589,11 @@ def _ensure_model(self) -> None: self.model = model def _materialize_one(self, sample: DraftReplaySample) -> DraftFeatureSample: + if self.backend == "vllm_file": + return self._materialize_one_vllm_file(sample) + return self._materialize_one_torch(sample) + + def _materialize_one_torch(self, sample: DraftReplaySample) -> DraftFeatureSample: self._ensure_model() assert self.backbone is not None assert self.final_norm is not None @@ -652,6 +730,204 @@ def hook(_module, _inputs, output): metadata=metadata, ) + def _materialize_one_vllm_file( + self, sample: DraftReplaySample + ) -> DraftFeatureSample: + if self.use_logits: + raise NotImplementedError( + "target_feature_replay.backend=vllm_file does not yet support " + "training.use_logits=true; use backend=torch for EAGLE3 logits." + ) + self._validate_vllm_positions(sample) + feature_positions = sample.feature_positions.detach().cpu().long() + feature_end = int(feature_positions[-1].item()) + 1 + prompt_ids = sample.input_ids[:feature_end].detach().cpu().long().tolist() + hidden_payload = self._request_vllm_hidden_states(prompt_ids) + try: + feature = self._feature_from_vllm_payload( + sample, + hidden_payload, + prompt_ids=prompt_ids, + source="token_replay_vllm_file", + ) + finally: + path = hidden_payload.get("_path") + if self.vllm_on_generate == "delete" and path: + try: + Path(os.fspath(path)).unlink(missing_ok=True) + except OSError: + logger.warning("Failed to delete vLLM hidden-states file %s", path) + return feature + + def _validate_vllm_positions(self, sample: DraftReplaySample) -> None: + if not self.vllm_require_arange_positions: + return + feature_positions = sample.feature_positions.detach().cpu().long() + feature_end = int(feature_positions[-1].item()) + 1 + expected = torch.arange(feature_end, dtype=torch.long) + actual = sample.position_ids[:feature_end].detach().cpu().long() + if not torch.equal(actual, expected): + raise ValueError( + "target_feature_replay.backend=vllm_file currently requires " + "position_ids to be contiguous arange positions for the replay prefix" + ) + + def _ensure_vllm_client(self) -> None: + if self.vllm_client is not None: + return + try: + import openai + except ImportError as exc: + raise RuntimeError( + "target_feature_replay.backend=vllm_file requires the openai package" + ) from exc + self.vllm_client = openai.OpenAI( + base_url=self.vllm_endpoint, + api_key="EMPTY", + max_retries=0, + ) + if self.vllm_model: + self.vllm_resolved_model = os.fspath(self.vllm_model) + else: + models = self.vllm_client.models.list() + self.vllm_resolved_model = models.data[0].id + + def _request_vllm_hidden_states(self, prompt_ids: list[int]) -> dict[str, Any]: + self._ensure_vllm_client() + assert self.vllm_client is not None + assert self.vllm_resolved_model is not None + last_error: Exception | None = None + started = time.perf_counter() + for attempt in range(self.vllm_max_retries + 1): + try: + response = self.vllm_client.completions.create( + model=self.vllm_resolved_model, + prompt=prompt_ids, + max_tokens=1, + extra_body={"return_token_ids": True}, + timeout=self.vllm_timeout, + ) + path = self._extract_hidden_states_path(response, prompt_ids) + payload = self._load_vllm_hidden_states(path) + payload["_path"] = path + self.vllm_requests += 1 + self.vllm_request_seconds += time.perf_counter() - started + return payload + except Exception as exc: # noqa: BLE001 + last_error = exc + if attempt >= self.vllm_max_retries: + break + time.sleep(float(2**attempt)) + self.vllm_request_seconds += time.perf_counter() - started + raise RuntimeError( + f"Failed to request vLLM hidden states after " + f"{self.vllm_max_retries + 1} attempts: {last_error}" + ) from last_error + + def _extract_hidden_states_path(self, response: Any, prompt_ids: list[int]) -> str: + choices = getattr(response, "choices", None) or [] + if choices: + prompt_token_ids = getattr(choices[0], "prompt_token_ids", None) + if prompt_token_ids is not None and list(prompt_token_ids) != prompt_ids: + raise ValueError( + "vLLM prompt_token_ids mismatch while extracting hidden states" + ) + kv_transfer_params = getattr(response, "kv_transfer_params", None) + if kv_transfer_params is None: + raise ValueError("vLLM response missing kv_transfer_params") + path = kv_transfer_params.get("hidden_states_path") + if not path: + raise ValueError("vLLM response missing hidden_states_path") + return os.fspath(path) + + def _load_vllm_hidden_states(self, path: str) -> dict[str, Any]: + try: + from safetensors.torch import load_file + except ImportError as exc: + raise RuntimeError("vLLM hidden-state replay requires safetensors") from exc + file_path = Path(path) + lock_path = Path(f"{path}.lock") + if lock_path.exists(): + _wait_for_lock(lock_path) + if not file_path.exists(): + raise FileNotFoundError(f"vLLM hidden-states file not found: {path}") + return dict(load_file(str(file_path), device="cpu")) + + def _feature_from_vllm_payload( + self, + sample: DraftReplaySample, + payload: dict[str, Any], + *, + prompt_ids: list[int], + source: str, + ) -> DraftFeatureSample: + token_ids = payload.get("token_ids") + hidden = payload.get("hidden_states") + if not torch.is_tensor(token_ids) or not torch.is_tensor(hidden): + raise ValueError( + "vLLM hidden-states payload must contain token_ids and hidden_states" + ) + if token_ids.detach().cpu().long().tolist() != prompt_ids: + raise ValueError("vLLM hidden-states token_ids do not match replay input") + if hidden.dim() != 3: + raise ValueError( + "vLLM hidden_states must have shape [seq, layers, hidden], " + f"got {tuple(hidden.shape)}" + ) + feature_positions = sample.feature_positions.detach().cpu().long() + expected_layers = len(self.target_layer_ids) + include_final = self.hidden_layout in { + "eagle3_aux_plus_last", + "dflash_aux_plus_last", + } + required_layers = expected_layers + (1 if include_final else 0) + if int(hidden.size(1)) < required_layers: + raise ValueError( + "vLLM hidden_states layer count is too small: " + f"got {int(hidden.size(1))}, need at least {required_layers}. " + "Start vLLM with target layer ids plus the final layer when the " + "training layout needs last hidden states." + ) + selected = hidden.index_select(0, feature_positions).to(dtype=self.dtype) + aux_hidden = selected[:, :expected_layers, :].flatten(1) + if include_final: + final_hidden = selected[:, required_layers - 1, :] + hidden_states = torch.cat([aux_hidden, final_hidden], dim=-1) + else: + hidden_states = aux_hidden + selected_input_ids = sample.input_ids.index_select(0, feature_positions).long() + selected_loss_mask = sample.loss_mask.index_select(0, feature_positions).float() + metadata = dict(sample.metadata) + feature_start = int(feature_positions[0].item()) + feature_end = int(feature_positions[-1].item()) + 1 + metadata.update( + { + "source": source, + "target_model_path": self.model_path, + "target_revision": self.target_revision, + "target_config_fingerprint": self.target_config_fingerprint, + "target_layer_ids": list(self.target_layer_ids), + "vllm_hidden_layers": int(hidden.size(1)), + "hidden_states_layout": self.hidden_layout, + "feature_start": feature_start, + "feature_end": feature_end, + "hidden_position_start": feature_start, + "hidden_position_end": feature_end, + "hidden_positions": feature_positions, + "sequence_length": int(selected_input_ids.numel()), + "full_sequence_length": int(sample.input_ids.numel()), + "use_logits": self.use_logits, + } + ) + return DraftFeatureSample( + algorithm=self.algorithm, + input_ids=selected_input_ids, + loss_mask=selected_loss_mask, + hidden_states=hidden_states.cpu().contiguous(), + position_ids=sample.draft_position_ids.long(), + metadata=metadata, + ) + def _build_sparse_target_logprobs( self, final_hidden_states: torch.Tensor ) -> torch.Tensor: @@ -684,6 +960,11 @@ def metrics(self) -> dict[str, float]: "replay/materialized_samples_total": float(self.materialized_samples), "replay/target_forward_time_total": float(self.target_forward_seconds), } + if self.backend == "vllm_file": + metrics["replay/vllm_requests_total"] = float(self.vllm_requests) + metrics["replay/vllm_request_time_total"] = float( + self.vllm_request_seconds + ) total = self.cache_hits + self.cache_misses if total > 0: metrics["replay/cache_hit_ratio"] = self.cache_hits / float(total) diff --git a/verl_speco/vllm_hidden_states_generate.py b/verl_speco/vllm_hidden_states_generate.py new file mode 100644 index 00000000..c775ca5e --- /dev/null +++ b/verl_speco/vllm_hidden_states_generate.py @@ -0,0 +1,144 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Generate vLLM safetensors draft features from token replay samples.""" + +from __future__ import annotations + +import logging +import os +from copy import deepcopy +from typing import Any + +import hydra +import torch +from omegaconf import OmegaConf, open_dict + +from verl_speco.trainer.draft_dataset import ( + DraftFeatureDataLoader, + DraftFeatureDataLoaderConfig, +) +from verl_speco.trainer.feature_store import build_feature_store_from_config +from verl_speco.trainer.target_feature_replay import TargetFeatureReplayer + +logger = logging.getLogger(__name__) + + +def _plain_config(value: Any) -> dict[str, Any]: + return dict(OmegaConf.to_container(value, resolve=True) or {}) + + +def generate_vllm_safetensors_features(config) -> dict[str, Any]: + """Materialize token replay samples through vLLM and save safetensors.""" + + draft_config = config.actor_rollout_ref + training_cfg = draft_config.rollout.drafter.training + replay_cfg = training_cfg.get("target_feature_replay", {}) or {} + generation_cfg = replay_cfg.get("offline_generation", {}) or {} + feature_store_cfg = training_cfg.feature_store + + input_path = generation_cfg.get("input_path", None) + output_path = generation_cfg.get("output_path", None) or feature_store_cfg.get( + "path", None + ) + if not input_path: + raise ValueError( + "target_feature_replay.offline_generation.input_path is required" + ) + if not output_path: + raise ValueError( + "target_feature_replay.offline_generation.output_path or " + "training.feature_store.path is required" + ) + + input_cfg = _plain_config(feature_store_cfg) + input_cfg.update({"type": "token_replay", "path": os.fspath(input_path)}) + + output_cfg = _plain_config(feature_store_cfg) + output_cfg.update({"type": "vllm_safetensors", "path": os.fspath(output_path)}) + + max_samples = int(generation_cfg.get("max_samples", 0) or 0) + batch_size = max(int(generation_cfg.get("batch_size", 1) or 1), 1) + shuffle = bool(generation_cfg.get("shuffle", False)) + seed = int(training_cfg.get("seed", 0) or 0) + + replay_config = deepcopy(config) + with open_dict(replay_config): + target_feature_replay = ( + replay_config.actor_rollout_ref.rollout.drafter.training.target_feature_replay + ) + target_feature_replay.backend = "vllm_file" + + input_store = build_feature_store_from_config(input_cfg, read_only=True) + output_store = build_feature_store_from_config( + output_cfg, + read_only=False, + metadata={ + "source_format": "token_replay", + "source_path": os.fspath(input_path), + "target_feature_backend": "vllm_file", + }, + ) + replayer = TargetFeatureReplayer( + replay_config, + rank=0, + world_size=1, + device=torch.device("cpu"), + ) + written = 0 + try: + loader = DraftFeatureDataLoader( + input_store, + DraftFeatureDataLoaderConfig( + batch_size=batch_size, + rank=0, + world_size=1, + shuffle=shuffle, + repeat=False, + seed=seed, + ), + ) + for samples in loader: + if max_samples > 0 and written >= max_samples: + break + if max_samples > 0: + samples = samples[: max_samples - written] + features = replayer.materialize(samples) + output_store.write_many(features) + written += len(features) + if written % 100 == 0: + logger.warning("Generated %s vLLM safetensors samples", written) + finally: + input_store.close() + output_store.close() + replayer.close() + + metrics = replayer.metrics() + result = { + "input_path": os.fspath(input_path), + "output_path": os.fspath(output_path), + "written_samples": written, + "metrics": metrics, + } + logger.warning("vLLM safetensors generation finished: %s", result) + return result + + +@hydra.main(config_path="config", config_name="draft_trainer", version_base=None) +def main(config): + logging.basicConfig(level=logging.INFO) + generate_vllm_safetensors_features(config) + + +if __name__ == "__main__": + main() From e224d5734e4de7f95cf22367397e9d1531a42d75 Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Fri, 31 Jul 2026 10:06:13 +0800 Subject: [PATCH 23/50] Fix vLLM hidden-state lock handling for token replay --- tests/unit/test_draft_training_loop.py | 17 +++ verl_speco/trainer/draft_training_loop.py | 133 +++++++++++++++++++- verl_speco/trainer/target_feature_replay.py | 121 +++++++++++++++--- 3 files changed, 246 insertions(+), 25 deletions(-) diff --git a/tests/unit/test_draft_training_loop.py b/tests/unit/test_draft_training_loop.py index c607697c..af34631b 100644 --- a/tests/unit/test_draft_training_loop.py +++ b/tests/unit/test_draft_training_loop.py @@ -25,6 +25,7 @@ from verl_speco.trainer.draft_training_loop import ( # noqa: E402 _rewrite_standalone_block_runtime_config, _save_standalone_checkpoint, + _should_log_batch_progress, _torch_load_cpu, ) @@ -43,6 +44,22 @@ def _save_checkpoint_async(self, step: int): return self.future +@pytest.mark.parametrize( + ("attempted_batches", "expected"), + [ + (1, True), + (2, True), + (3, True), + (4, False), + (99, False), + (100, True), + (101, False), + ], +) +def test_should_log_standalone_batch_progress(attempted_batches, expected): + assert _should_log_batch_progress(attempted_batches) is expected + + def test_standalone_checkpoint_schedules_without_waiting(): trainer = _FakeTrainer() diff --git a/verl_speco/trainer/draft_training_loop.py b/verl_speco/trainer/draft_training_loop.py index 7ce56ef3..b3ee29bb 100644 --- a/verl_speco/trainer/draft_training_loop.py +++ b/verl_speco/trainer/draft_training_loop.py @@ -41,6 +41,10 @@ logger = logging.getLogger(__name__) +def _should_log_batch_progress(attempted_batches: int) -> bool: + return attempted_batches <= 3 or attempted_batches % 100 == 0 + + def run_standalone_draft_training(config) -> dict[str, Any]: """Run independent draft training from a feature store.""" return asyncio.run(_run_standalone_draft_training_async(config)) @@ -48,6 +52,12 @@ def run_standalone_draft_training(config) -> dict[str, Any]: async def _run_standalone_draft_training_async(config) -> dict[str, Any]: rank, local_rank, world_size = _init_distributed() + logger.info( + "[standalone rank=%s] distributed runtime initialized local_rank=%s world_size=%s", + rank, + local_rank, + world_size, + ) draft_config = config.actor_rollout_ref drafter_cfg = draft_config.rollout.drafter training_cfg = drafter_cfg.training @@ -103,17 +113,42 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: last_saved_step = 0 store = None feature_replayer = None + current_stage = "activate_training_model" try: + stage_started = time.perf_counter() + logger.info( + "[standalone rank=%s] activating drafter model algorithm=%s", + rank, + drafter_cfg.speculative_algorithm, + ) activated = await trainer.activate_training_model() if not activated: raise RuntimeError( f"Failed to activate standalone drafter trainer on rank={rank}" ) + logger.info( + "[standalone rank=%s] drafter model activated elapsed=%.3fs", + rank, + time.perf_counter() - stage_started, + ) initial_optimizer_step = int(trainer.optimizer_steps_total) optimizer_step = initial_optimizer_step last_saved_step = optimizer_step + current_stage = "open_feature_store" + stage_started = time.perf_counter() + logger.info( + "[standalone rank=%s] opening feature store type=%s path=%s", + rank, + feature_store_type, + feature_store_cfg.get("path"), + ) store = build_feature_store_from_config(feature_store_cfg, read_only=True) + logger.info( + "[standalone rank=%s] feature store opened elapsed=%.3fs", + rank, + time.perf_counter() - stage_started, + ) if feature_store_type == "token_replay": # Keep the large target model entirely outside online training imports # and lifetime. The standalone loop materializes ordinary feature @@ -122,12 +157,26 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: TargetFeatureReplayer, ) + current_stage = "initialize_target_feature_replayer" + stage_started = time.perf_counter() + logger.info( + "[standalone rank=%s] initializing target feature replayer", + rank, + ) feature_replayer = TargetFeatureReplayer( config, rank=rank, world_size=world_size, device=trainer.runtime_device, ) + logger.info( + "[standalone rank=%s] target feature replayer initialized " + "backend=%s elapsed=%.3fs", + rank, + feature_replayer.backend, + time.perf_counter() - stage_started, + ) + current_stage = "create_dataloader" loader = DraftFeatureDataLoader( store, DraftFeatureDataLoaderConfig( @@ -139,21 +188,59 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: seed=int(training_cfg.get("seed", 0) or 0), ), ) + logger.info( + "[standalone rank=%s] dataloader ready batch_size_per_gpu=%s " + "shuffle=%s repeat=%s", + rank, + int(training_cfg.get("batch_size_per_gpu", 4)), + bool(feature_store_cfg.get("shuffle", True)), + bool(feature_store_cfg.get("repeat", True)), + ) for samples in loader: if max_steps > 0 and successful_steps >= max_steps: break step_started = time.perf_counter() attempted_batches += 1 - materialized_samples = ( - feature_replayer.materialize(samples) - if feature_replayer is not None - else samples - ) + log_batch_progress = _should_log_batch_progress(attempted_batches) + if log_batch_progress: + logger.info( + "[standalone rank=%s] batch=%s loaded samples=%s " + "successful_steps=%s", + rank, + attempted_batches, + len(samples), + successful_steps, + ) + current_stage = "materialize_target_features" + if feature_replayer is not None: + materialize_started = time.perf_counter() + if log_batch_progress: + logger.info( + "[standalone rank=%s] batch=%s materializing target features " + "backend=%s", + rank, + attempted_batches, + feature_replayer.backend, + ) + materialized_samples = feature_replayer.materialize(samples) + if log_batch_progress: + logger.info( + "[standalone rank=%s] batch=%s target features materialized " + "samples=%s elapsed=%.3fs", + rank, + attempted_batches, + len(materialized_samples), + time.perf_counter() - materialize_started, + ) + else: + materialized_samples = samples + current_stage = "prepare_training_batch" batch = trainer.prepare_training_batch_from_samples( cast(list[Any], materialized_samples), step=optimizer_step, ) has_batch = batch is not None + current_stage = "synchronize_batch_readiness" if not _all_ranks_true(has_batch, trainer.runtime_device): if rank == 0: logger.warning( @@ -163,7 +250,17 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: if batch is None: continue trainer.reset_training_metrics() + current_stage = "training_step" + if log_batch_progress: + logger.info( + "[standalone rank=%s] batch=%s starting drafter training step " + "optimizer_step=%s", + rank, + attempted_batches, + optimizer_step, + ) ok = await trainer.training_step_from_batch(batch, optimizer_step) + current_stage = "synchronize_training_step" if not _all_ranks_true(ok, trainer.runtime_device): continue successful_steps += 1 @@ -180,25 +277,51 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: step_metrics.update(feature_replayer.metrics()) _log_standalone_step_metrics(step_metrics, rank=rank) if save_interval > 0 and optimizer_step % save_interval == 0: + current_stage = "save_checkpoint" last_save_result = _save_standalone_checkpoint(trainer, optimizer_step) if _sync_any_rank_saved_checkpoint(last_save_result.get("saved")): last_saved_step = optimizer_step _barrier() + current_stage = "load_next_batch" final_save = bool(training_cfg.get("save_final_checkpoint", True)) if final_save and successful_steps > 0 and optimizer_step != last_saved_step: + current_stage = "save_final_checkpoint" last_save_result = _save_standalone_checkpoint( trainer, optimizer_step, wait=True ) _barrier() + except Exception: + logger.exception( + "[standalone rank=%s] training failed stage=%s attempted_batches=%s " + "successful_steps=%s optimizer_step=%s", + rank, + current_stage, + attempted_batches, + successful_steps, + optimizer_step, + ) + raise finally: + logger.info( + "[standalone rank=%s] cleanup starting stage=%s attempted_batches=%s " + "successful_steps=%s", + rank, + current_stage, + attempted_batches, + successful_steps, + ) if store is not None: store.close() if feature_replayer is not None: feature_replayer.close() + logger.info("[standalone rank=%s] cleaning trainer resources", rank) await trainer.cleanup_training(clear_data=True) if dist.is_initialized(): + logger.info("[standalone rank=%s] entering final process-group barrier", rank) dist.barrier() + logger.info("[standalone rank=%s] final process-group barrier complete", rank) dist.destroy_process_group() + logger.info("[standalone rank=%s] cleanup complete", rank) return { "rank": rank, diff --git a/verl_speco/trainer/target_feature_replay.py b/verl_speco/trainer/target_feature_replay.py index 017bfcb2..37af5537 100644 --- a/verl_speco/trainer/target_feature_replay.py +++ b/verl_speco/trainer/target_feature_replay.py @@ -173,7 +173,7 @@ def _wait_for_lock(lock_path: Path, timeout: float = 30.0) -> None: deadline = time.monotonic() + float(timeout) while True: try: - fcntl.flock(fd, fcntl.LOCK_EX | fcntl.LOCK_NB) + fcntl.flock(fd, fcntl.LOCK_SH | fcntl.LOCK_NB) break except BlockingIOError: if time.monotonic() >= deadline: @@ -181,6 +181,7 @@ def _wait_for_lock(lock_path: Path, timeout: float = 30.0) -> None: f"Timed out waiting for hidden-states lock: {lock_path}" ) from None time.sleep(0.1) + fcntl.flock(fd, fcntl.LOCK_UN) finally: os.close(fd) try: @@ -465,31 +466,57 @@ def __init__( self.target_forward_seconds = 0.0 self.vllm_request_seconds = 0.0 self.vllm_requests = 0 + logger.info( + "[target replay rank=%s] initialized backend=%s algorithm=%s " + "target_layers=%s hidden_layout=%s use_logits=%s endpoint=%s cache=%s", + self.rank, + self.backend, + self.algorithm, + self.target_layer_ids, + self.hidden_layout, + self.use_logits, + self.vllm_endpoint if self.backend == "vllm_file" else None, + self.cache is not None, + ) def materialize( self, samples: Iterable[DraftReplaySample | DraftFeatureSample] ) -> list[DraftFeatureSample]: materialized: list[DraftFeatureSample] = [] - for sample in samples: - if isinstance(sample, DraftFeatureSample): - materialized.append(sample) - continue - if not isinstance(sample, DraftReplaySample): - raise TypeError( - f"Target feature replay expected DraftReplaySample, got {type(sample)!r}" + for sample_index, sample in enumerate(samples): + try: + if isinstance(sample, DraftFeatureSample): + materialized.append(sample) + continue + if not isinstance(sample, DraftReplaySample): + raise TypeError( + "Target feature replay expected DraftReplaySample, " + f"got {type(sample)!r}" + ) + self._validate_target_path(sample) + key = self._cache_key(sample) + cached = self.cache.get(key) if self.cache is not None else None + if cached is not None: + self.cache_hits += 1 + materialized.append(cached) + continue + self.cache_misses += 1 + replayed = self._materialize_one(sample) + if self.cache is not None: + self.cache.put(key, replayed) + materialized.append(replayed) + except Exception: + metadata = getattr(sample, "metadata", {}) or {} + logger.exception( + "[target replay rank=%s] sample materialization failed " + "sample_index=%s algorithm=%s source=%s global_step=%s", + self.rank, + sample_index, + getattr(sample, "algorithm", None), + metadata.get("source"), + metadata.get("global_step"), ) - self._validate_target_path(sample) - key = self._cache_key(sample) - cached = self.cache.get(key) if self.cache is not None else None - if cached is not None: - self.cache_hits += 1 - materialized.append(cached) - continue - self.cache_misses += 1 - replayed = self._materialize_one(sample) - if self.cache is not None: - self.cache.put(key, replayed) - materialized.append(replayed) + raise self.materialized_samples += len(materialized) return materialized @@ -775,6 +802,13 @@ def _validate_vllm_positions(self, sample: DraftReplaySample) -> None: def _ensure_vllm_client(self) -> None: if self.vllm_client is not None: return + started = time.perf_counter() + logger.info( + "[target replay rank=%s] connecting to vLLM endpoint=%s configured_model=%s", + self.rank, + self.vllm_endpoint, + self.vllm_model, + ) try: import openai except ImportError as exc: @@ -791,6 +825,12 @@ def _ensure_vllm_client(self) -> None: else: models = self.vllm_client.models.list() self.vllm_resolved_model = models.data[0].id + logger.info( + "[target replay rank=%s] connected to vLLM model=%s elapsed=%.3fs", + self.rank, + self.vllm_resolved_model, + time.perf_counter() - started, + ) def _request_vllm_hidden_states(self, prompt_ids: list[int]) -> dict[str, Any]: self._ensure_vllm_client() @@ -798,8 +838,21 @@ def _request_vllm_hidden_states(self, prompt_ids: list[int]) -> dict[str, Any]: assert self.vllm_resolved_model is not None last_error: Exception | None = None started = time.perf_counter() + request_index = self.vllm_requests + 1 + log_request = request_index <= 2 or request_index % 100 == 0 for attempt in range(self.vllm_max_retries + 1): try: + attempt_started = time.perf_counter() + if log_request: + logger.info( + "[target replay rank=%s] vLLM request starting request=%s " + "attempt=%s/%s prompt_tokens=%s", + self.rank, + request_index, + attempt + 1, + self.vllm_max_retries + 1, + len(prompt_ids), + ) response = self.vllm_client.completions.create( model=self.vllm_resolved_model, prompt=prompt_ids, @@ -812,9 +865,37 @@ def _request_vllm_hidden_states(self, prompt_ids: list[int]) -> dict[str, Any]: payload["_path"] = path self.vllm_requests += 1 self.vllm_request_seconds += time.perf_counter() - started + if log_request: + hidden_states = payload.get("hidden_states") + hidden_shape = ( + tuple(hidden_states.shape) + if torch.is_tensor(hidden_states) + else None + ) + logger.info( + "[target replay rank=%s] vLLM request completed request=%s " + "attempt=%s path=%s hidden_shape=%s elapsed=%.3fs", + self.rank, + request_index, + attempt + 1, + path, + hidden_shape, + time.perf_counter() - attempt_started, + ) return payload except Exception as exc: # noqa: BLE001 last_error = exc + logger.warning( + "[target replay rank=%s] vLLM request failed request=%s " + "attempt=%s/%s prompt_tokens=%s elapsed=%.3fs error=%r", + self.rank, + request_index, + attempt + 1, + self.vllm_max_retries + 1, + len(prompt_ids), + time.perf_counter() - started, + exc, + ) if attempt >= self.vllm_max_retries: break time.sleep(float(2**attempt)) From 68a46f9cb1a051e8b3d939dea124217791300588 Mon Sep 17 00:00:00 2001 From: Cai Zeyong <878049625@qq.com> Date: Fri, 31 Jul 2026 12:59:09 +0800 Subject: [PATCH 24/50] =?UTF-8?q?actor=20hidden=20states=20tq=20=E5=AE=9E?= =?UTF-8?q?=E7=8E=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- verl_speco/integration/oldlogprob_runtime.py | 64 +++++++++- verl_speco/integration/sglang_runtime.py | 20 +++- verl_speco/integration/task_runner.py | 21 ++-- .../integration/transferqueue_bridge.py | 110 +++++++++++------- verl_speco/trainer/speco_ray_trainer.py | 4 + verl_speco/workers/speco_worker.py | 69 ++++++++--- 6 files changed, 217 insertions(+), 71 deletions(-) diff --git a/verl_speco/integration/oldlogprob_runtime.py b/verl_speco/integration/oldlogprob_runtime.py index 0c4da7bf..ccd30b9f 100644 --- a/verl_speco/integration/oldlogprob_runtime.py +++ b/verl_speco/integration/oldlogprob_runtime.py @@ -35,6 +35,9 @@ OLD_LOGPROB_HIDDEN_LAYOUT_KEY = "speco_oldlogprob_hidden_layout" OLD_LOGPROB_TIMING_KEY = "speco_oldlogprob_timing" OLD_LOGPROB_SELECTED_BATCH_INDICES_KEY = "speco_oldlogprob_selected_batch_indices" +# Stamped on the micro-batch by the collect plan so the actor-worker producer +# can build step-unique TransferQueue keys for old-logprob hidden chunks (P1). +OLD_LOGPROB_GLOBAL_STEP_KEY = "speco_oldlogprob_global_step" _TIMING_SELECT_US = 0 _TIMING_SP_MERGE_US = 1 @@ -180,6 +183,45 @@ def _oldlogprob_hidden_object_ref_enabled(micro_batch: Any) -> bool: return bool(value) +def _oldlogprob_global_step(micro_batch: Any) -> int: + """Read the global step stamped on the micro-batch by the collect plan.""" + + value = 0 + try: + from verl.utils import tensordict_utils as tu + + value = tu.get_non_tensor_data(data=micro_batch, key=OLD_LOGPROB_GLOBAL_STEP_KEY, default=0) + except Exception: # noqa: BLE001 + try: + value = micro_batch.get(OLD_LOGPROB_GLOBAL_STEP_KEY, 0) + except Exception: # noqa: BLE001 + if _tensor_key_present(micro_batch, OLD_LOGPROB_GLOBAL_STEP_KEY): + value = micro_batch[OLD_LOGPROB_GLOBAL_STEP_KEY] + value = getattr(value, "data", value) + try: + return int(value) + except (TypeError, ValueError): + return 0 + + +def _speco_tq_enabled_for_oldlogprob() -> bool: + """Configure (idempotently) and report whether TQ transport is usable here. + + The actor worker reaches its drafter training config through the same env + serialization the SGLang server uses (``SPECO_SGLANG_DRAFTER_CONFIG_ENV``). + """ + + from verl_speco.integration.transferqueue_bridge import ( + configure_transfer_queue, + is_transfer_queue_enabled, + ) + + drafter = _load_drafter_env() + training = _get_nested(drafter, ("training",), None) + configure_transfer_queue(training) + return is_transfer_queue_enabled() + + def _is_sparse_sp_non_source_context(context: dict[str, Any]) -> bool: return bool(context.get("sparse_sp_merge")) and int(context.get("sp_rank", 0) or 0) != 0 @@ -378,7 +420,27 @@ def _put_oldlogprob_hidden_refs(hidden_output: dict[str, Any], micro_batch: Any) offset += length hidden_chunk = torch.cat(tensors, dim=0).contiguous() if len(tensors) > 1 else tensors[0].contiguous() ray_put_started = time.perf_counter() - chunk_ref = ray.put(hidden_chunk) + # P1: when TQ is enabled, store the owner's concatenated hidden chunk + # in TransferQueue and carry the key in place of the Ray ObjectRef. + # The driver treats the token as opaque; the drafter consumer + # resolves "speco:" keys via TQ instead of ray.get. Falls back to + # ray.put otherwise (unchanged behavior). + if _speco_tq_enabled_for_oldlogprob(): + from verl_speco.integration.transferqueue_bridge import make_sample_key, put_sample + + tq_key = make_sample_key( + _oldlogprob_global_step(micro_batch), + int(owner), + f"chunk{len(chunk_refs)}", + ) + put_sample( + tq_key, + {"hidden": hidden_chunk}, + tag={"global_step": _oldlogprob_global_step(micro_batch), "owner": int(owner)}, + ) + chunk_ref = tq_key + else: + chunk_ref = ray.put(hidden_chunk) ray_put_us += (time.perf_counter() - ray_put_started) * 1_000_000.0 chunk_index = len(chunk_refs) chunk_refs.append(chunk_ref) diff --git a/verl_speco/integration/sglang_runtime.py b/verl_speco/integration/sglang_runtime.py index e4556863..90463b40 100644 --- a/verl_speco/integration/sglang_runtime.py +++ b/verl_speco/integration/sglang_runtime.py @@ -1648,13 +1648,31 @@ async def generate( configure_transfer_queue(training_cfg) if is_transfer_queue_enabled(): tq_key = make_sample_key(collection_global_steps, self.replica_rank, request_id) + tq_payload = {"hidden_states": hidden_states.unsqueeze(0).cpu()} + # P2: also offload the other large tensors so they bypass + # the driver too. The drafter consumer restores them from + # this same TQ payload. + if target_logprobs is not None: + tq_payload["target_logprobs"] = target_logprobs.unsqueeze(0).cpu() + if torch.is_tensor(hidden_raw_target_logprobs): + tq_payload["hidden_raw_target_logprobs"] = hidden_raw_target_logprobs.unsqueeze(0).cpu() + if torch.is_tensor(hidden_raw_target_logprobs_positions): + tq_payload["hidden_raw_target_logprobs_positions"] = ( + hidden_raw_target_logprobs_positions.unsqueeze(0).cpu() + ) put_sample( tq_key, - {"hidden_states": hidden_states.unsqueeze(0).cpu()}, + tq_payload, tag={"global_step": collection_global_steps, "replica_rank": self.replica_rank}, ) drafter_sample["hidden_states_tq_key"] = tq_key drafter_sample["hidden_states"] = None + if target_logprobs is not None: + drafter_sample["target_logprobs"] = None + if torch.is_tensor(hidden_raw_target_logprobs): + drafter_sample["hidden_raw_target_logprobs"] = None + if torch.is_tensor(hidden_raw_target_logprobs_positions): + drafter_sample["hidden_raw_target_logprobs_positions"] = None else: self._speco_log_missing_hidden_states_once( collection_global_steps=collection_global_steps, diff --git a/verl_speco/integration/task_runner.py b/verl_speco/integration/task_runner.py index cb1c1835..20c20012 100644 --- a/verl_speco/integration/task_runner.py +++ b/verl_speco/integration/task_runner.py @@ -270,11 +270,16 @@ def _run_with_speco_trainer(self, config): ) # Bootstrap TransferQueue before spawning Ray actors so worker processes - # inherit the TQ environment (mirrors verl main_ppo_sync tq.init). No-op - # when transfer_queue.enable=false or the package is not installed. - from verl_speco.integration.transferqueue_bridge import init_transfer_queue - - init_transfer_queue(config) - - trainer.init_workers() - trainer.fit() + # inherit the TQ environment (mirrors verl main_ppo_sync tq.init). The + # TaskRunner owns shutdown so a failed fit cannot leak the named + # controller/storage. No-op when transfer_queue.enable=false or the + # package is not installed. + from verl_speco.integration.transferqueue_bridge import close_transfer_queue, init_transfer_queue + + transfer_queue_started = init_transfer_queue(config) + try: + trainer.init_workers() + trainer.fit() + finally: + if transfer_queue_started: + close_transfer_queue() diff --git a/verl_speco/integration/transferqueue_bridge.py b/verl_speco/integration/transferqueue_bridge.py index e69b77f2..7062b7f4 100644 --- a/verl_speco/integration/transferqueue_bridge.py +++ b/verl_speco/integration/transferqueue_bridge.py @@ -1,28 +1,32 @@ """TransferQueue bridge for SPECO drafter feature transport. This module lets SPECO route large per-sample drafter-training tensors (hidden -states, and later target logprobs) through TransferQueue (TQ) instead of -funneling them through the ``SpecoRayPPOTrainer`` driver process and the Ray -object store. It is the SpeCo-side analog of verl's ``transferqueue_utils.py``, -but used as a standalone transport library -- it does **not** depend on verl's +states, target logprobs) through TransferQueue (TQ) instead of funneling them +through the ``SpecoRayPPOTrainer`` driver process and the Ray object store. It +is the SpeCo-side analog of verl's ``transferqueue_utils.py``, but used as a +standalone transport library -- it does **not** depend on verl's ``main_ppo_sync`` TQ integration and does **not** modify upstream verl. -Design (P0): -- Only the dominant tensor (``hidden_states``) is offloaded to TQ. The rest of - the ``drafter_sample`` dict (input_ids, prompts, responses, positions, - metadata scalars) keeps riding the existing DataProto side-channel + Ray - dispatch. The TQ key rides with the sample dict, so the driver is unchanged. +Design: +- P0: offload SGLang-collected ``hidden_states`` (a1 path). +- P1: offload old-logprob-collected ``hidden_states`` (a2 path) -- replaces the + ``ray.put`` chunk + driver relay. +- P2: also offload ``target_logprobs`` / ``hidden_raw_target_logprobs``. - Default ``enable: false`` -> behavior is bit-identical to the current Ray path. TQ is only touched when explicitly enabled and the ``transfer_queue`` package is importable. -- Producer = SGLang rollout server (``sglang_runtime.py``); consumer = drafter - worker (``speco_worker.py``). Both call ``configure_transfer_queue`` from - their respective drafter training config, then ``put_sample`` / ``get_sample``. + +A sample remains in TQ until task cleanup: a SPECO drafter replica contains +multiple TP/SP ranks and each rank reads the same sample (the owner-route +dispatch duplicates a DP bucket to all SP ranks of one replica; only the SP +leader is ``is_collect``). Deleting a sample after the first read would race the +remaining ranks. Garbage collection is therefore deferred to task teardown; a +finer-grained leader-clears-after-barrier is future work. Note: the TQ call sites follow the documented KV API (``kv_put`` / -``kv_batch_get`` / ``kv_clear``) of TransferQueue 0.1.7. When TQ is enabled, -these are exercised against the installed package; verify the exact signatures -against your TQ version on first run (the bridge fails loud, never silently). +``kv_batch_get`` / ``kv_close``) of TransferQueue 0.1.7. The exact signatures +(keyword names, return shapes) must be verified against the installed TQ +version on first run; the bridge fails loud, never silently. """ from __future__ import annotations @@ -78,6 +82,7 @@ def _raise(*args: Any, **kwargs: Any) -> Any: "configured": False, # configure_transfer_queue has run "initialized": False, # tq.init() has run in this process "config": None, # the transfer_queue sub-config (plain dict) + "owner": False, # this process created the task-level TQ system } @@ -142,10 +147,9 @@ def init_transfer_queue(config: Any) -> bool: """Cluster-wide TQ bootstrap, called once from the SpecoTaskRunner. Mirrors verl ``main_ppo_sync`` calling ``tq.init(config.transfer_queue)`` - in the TaskRunner before workers spawn. Ray actors (SGLang server, drafter - worker) inherit this process's environment, so their lazy ``tq.init()`` - (no-arg) connects to the already-started storage. Returns whether TQ is - usable; no-op (returns False) when disabled or not installed. + in the TaskRunner before workers spawn. Other Ray processes lazily call + ``tq.init()`` and connect to the named TransferQueue controller. Returns + whether TQ is usable; no-op (returns False) when disabled or not installed. """ tq_cfg = _extract_tq_config(_drafter_training_cfg(config)) @@ -156,6 +160,7 @@ def init_transfer_queue(config: Any) -> bool: _state["config"] = _to_plain_dict(tq_cfg) _state["enabled"] = True _state["initialized"] = True + _state["owner"] = True logger.info("[SpeCo TQ] TransferQueue bootstrapped in task runner (partition=%s)", _SPECO_TQ_PARTITION) return True @@ -170,9 +175,9 @@ def _drafter_training_cfg(config: Any) -> Any: def _ensure_initialized() -> None: """Lazily ``tq.init()`` once per worker process (mirrors verl TQ_INITIALIZED). - No-arg init relies on env inheritance from the task-runner bootstrap. If a - future TQ version does not propagate config via env, switch this to - ``tq.init(_state["config"])``. + A no-argument initialization discovers the named TransferQueue controller + on the connected Ray cluster. It deliberately does not create a separate + per-worker configuration. """ if _state["initialized"]: @@ -185,7 +190,7 @@ def _ensure_initialized() -> None: # --------------------------------------------------------------------------- -# Key / put / get / clear +# Key / put / get / close # --------------------------------------------------------------------------- def make_sample_key(global_step: Any, replica_rank: Any, request_id: Any) -> str: @@ -198,12 +203,6 @@ def make_sample_key(global_step: Any, replica_rank: Any, request_id: Any) -> str return f"speco:{global_step}:{replica_rank}:{request_id}" -def _to_tensordict(tensor_dict: dict) -> Any: - from tensordict import TensorDict - - return TensorDict(tensor_dict, batch_size=[]) - - def put_sample( key: str, tensor_dict: dict, @@ -224,34 +223,60 @@ def put_sample( if not payload: return _ensure_initialized() - value = _to_tensordict(payload) - # KV API (TransferQueue docs): kv_put(key, value, partition_id, tag). - tq.kv_put(key, value, partition_id=_SPECO_TQ_PARTITION, tag=tag or {}) + # Pass a plain single-sample dict of columns. TQ's kv_put adds its required + # batch dimension internally; constructing a scalar TensorDict here would be + # incorrect. (Exact kwarg names verified against TQ 0.1.7 on first run.) + tq.kv_put( + key=key, + partition_id=_SPECO_TQ_PARTITION, + fields=payload, + tag=tag or {}, + ) def get_sample(key: str) -> dict: - """Retrieve the tensor dict stored under ``key`` and free it. + """Retrieve one tensor dict without deleting it. - Returns a plain ``{field: tensor}`` dict. Clears the key after read since - each sample is consumed by exactly one drafter owner replica. + A drafter replica may execute this method on multiple TP/SP ranks; each + rank reads the same key. TQ storage is released once at task shutdown by + the process that initialized it, after every consumer has finished. """ if not is_transfer_queue_enabled(): raise RuntimeError("get_sample called while TransferQueue is not enabled.") _ensure_initialized() - # KV API: kv_batch_get(keys, partition_id) -> mapping key -> TensorDict - # (return shape is version-dependent; handle both dict and list forms). - result = tq.kv_batch_get([key], partition_id=_SPECO_TQ_PARTITION) + # TQ returns the stored sample (TensorDict-like). Return shape is version + # dependent, so handle both a direct value and a {key: value} mapping. + result = tq.kv_batch_get(keys=[key], partition_id=_SPECO_TQ_PARTITION) value = _extract_value(result, key) - try: - tq.kv_clear([key], partition_id=_SPECO_TQ_PARTITION) - except Exception: # noqa: BLE001 - logger.debug("[SpeCo TQ] kv_clear failed for key=%s (ignored)", key) if value is None: return {} return _tensordict_to_dict(value) +def close_transfer_queue() -> None: + """Close task-level TQ resources if this process initialized them. + + Only the TaskRunner owns controller/storage teardown. Worker-side lazy + clients must not call this because they may still be serving other ranks. + """ + + with _state_lock: + if not _state["owner"]: + return + _state["owner"] = False + _state["initialized"] = False + try: + tq.close() + except AttributeError: + # tq.close() is not part of the public TQ API on some versions; nothing + # to tear down explicitly. The named controller/storage actors are + # reaped when the Ray job exits. + logger.debug("[SpeCo TQ] tq.close() unavailable; skipping explicit teardown") + except Exception: # noqa: BLE001 + logger.debug("[SpeCo TQ] tq.close() raised; ignoring shutdown error") + + def _extract_value(result: Any, key: str) -> Any: if result is None: return None @@ -271,6 +296,7 @@ def _tensordict_to_dict(value: Any) -> dict: __all__ = [ "KVBatchMeta", "configure_transfer_queue", + "close_transfer_queue", "init_transfer_queue", "is_transfer_queue_enabled", "make_sample_key", diff --git a/verl_speco/trainer/speco_ray_trainer.py b/verl_speco/trainer/speco_ray_trainer.py index 58eb430c..c4f1ae40 100644 --- a/verl_speco/trainer/speco_ray_trainer.py +++ b/verl_speco/trainer/speco_ray_trainer.py @@ -26,6 +26,7 @@ from verl_speco.integration.oldlogprob_runtime import ( OLD_LOGPROB_AUX_LAYER_IDS_KEY, OLD_LOGPROB_COLLECT_MASK_KEY, + OLD_LOGPROB_GLOBAL_STEP_KEY, OLD_LOGPROB_HIDDEN_CAPTURE_IMPL_KEY, OLD_LOGPROB_HIDDEN_CHUNK_META_KEY, OLD_LOGPROB_HIDDEN_CHUNK_REFS_KEY, @@ -1868,6 +1869,9 @@ def compute_old_log_prob_without_collection(): self._speco_oldlogprob_hidden_layout(), ) tu.assign_non_tensor_data(batch_td, OLD_LOGPROB_HIDDEN_OBJECT_REF_KEY, True) + # Stamp the global step so the actor-worker producer can build + # step-unique TransferQueue keys for old-logprob hidden chunks (P1). + tu.assign_non_tensor_data(batch_td, OLD_LOGPROB_GLOBAL_STEP_KEY, self.global_steps) self._speco_last_oldlogprob_prepare_elapsed_sec = time.perf_counter() - prepare_started compute_started = time.perf_counter() diff --git a/verl_speco/workers/speco_worker.py b/verl_speco/workers/speco_worker.py index 7ec353e3..6afc2429 100644 --- a/verl_speco/workers/speco_worker.py +++ b/verl_speco/workers/speco_worker.py @@ -53,6 +53,18 @@ def _resolve_ray_object_ref(value): return value +def _resolve_tq_or_ray_ref(ref): + # P1: old-logprob chunk refs may be TransferQueue keys ("speco:" prefix) + # instead of Ray ObjectRefs. Fetch the stored chunk via TQ and fall back to + # ray.get otherwise. Multiple SP ranks of one replica read the same key; + # TQ samples are not cleared on read (see transferqueue_bridge.get_sample). + if isinstance(ref, str) and ref.startswith("speco:"): + from verl_speco.integration.transferqueue_bridge import get_sample + + return get_sample(ref).get("hidden") + return _resolve_ray_object_ref(ref) + + def _resolve_hidden_state_chunks(chunks, expected_rows: int | None = None): if not chunks: return None @@ -67,9 +79,9 @@ def _resolve_hidden_state_chunks(chunks, expected_rows: int | None = None): ref = chunk.get("ref") if ref is None: continue - cache_key = id(ref) + cache_key = ref if isinstance(ref, str) else id(ref) if cache_key not in resolved_cache: - resolved_cache[cache_key] = _resolve_ray_object_ref(ref) + resolved_cache[cache_key] = _resolve_tq_or_ray_ref(ref) tensor = resolved_cache[cache_key] if not torch.is_tensor(tensor): continue @@ -680,6 +692,32 @@ def collect_rollout_features(self, samples: list[dict]): for sample in samples: if not sample: continue + # P2: restore TransferQueue-offloaded tensors into the sample before + # building the batch. One fetch restores hidden_states plus the other + # large tensors (target_logprobs, raw target logprobs) so they all + # bypass the driver. a2 old-logprob samples (which carry + # hidden_states_ref_chunks, not a tq key) skip this and resolve below. + tq_key = sample.get("hidden_states_tq_key") + if tq_key is not None and self._speco_tq_enabled: + from verl_speco.integration.transferqueue_bridge import get_sample + + payload = get_sample(tq_key) + for _field in ( + "hidden_states", + "target_logprobs", + "hidden_raw_target_logprobs", + "hidden_raw_target_logprobs_positions", + ): + if payload.get(_field) is not None: + sample[_field] = payload[_field] + if sample.get("hidden_states") is None: + # Fail loud: a TQ key was produced but the payload is + # missing -> transport is broken. Do NOT silently drop the + # sample (that would corrupt training data). + raise RuntimeError( + f"[SpeCo TQ] drafter worker got empty hidden_states for " + f"key={tq_key}; TQ enabled but producer payload missing." + ) batch = { "input_ids": sample["input_ids"], "prompts": sample["prompts"], @@ -711,24 +749,17 @@ def collect_rollout_features(self, samples: list[dict]): batch[key] = sample[key] hidden = sample.get("hidden_states") if hidden is None: - tq_key = sample.get("hidden_states_tq_key") - if tq_key is not None and self._speco_tq_enabled: - # P0: hidden states were offloaded to TransferQueue by the - # rollout server; fetch by key (and free the storage). - from verl_speco.integration.transferqueue_bridge import get_sample - - payload = get_sample(tq_key) - hidden = payload.get("hidden_states") + # a2 old-logprob path: resolve Ray ObjectRefs (or TQ keys, P1) + # from per-owner chunks / single ref. + hidden_chunks = sample.get("hidden_states_ref_chunks") + if hidden_chunks: + expected_rows = None + hidden_positions = batch.get("hidden_positions") + if torch.is_tensor(hidden_positions): + expected_rows = int(hidden_positions.numel()) + hidden = _resolve_hidden_state_chunks(hidden_chunks, expected_rows=expected_rows) else: - hidden_chunks = sample.get("hidden_states_ref_chunks") - if hidden_chunks: - expected_rows = None - hidden_positions = batch.get("hidden_positions") - if torch.is_tensor(hidden_positions): - expected_rows = int(hidden_positions.numel()) - hidden = _resolve_hidden_state_chunks(hidden_chunks, expected_rows=expected_rows) - else: - hidden = _resolve_ray_object_ref(sample.get("hidden_states_ref")) + hidden = _resolve_ray_object_ref(sample.get("hidden_states_ref")) target_logprobs = sample.get("target_logprobs") if target_logprobs is None: target_logprobs = _resolve_ray_object_ref(sample.get("target_logprobs_ref")) From 99d0e40df35c9b7d0b2548bad283ff16bcc2046c Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Fri, 31 Jul 2026 16:08:33 +0800 Subject: [PATCH 25/50] Support JSONL token replay feature stores --- tests/unit/test_draft_feature_store.py | 117 ++++++++++ verl_speco/config/draft_trainer.yaml | 1 + verl_speco/inspect_jsonl_samples.py | 265 ++++++++++++++++++++++ verl_speco/trainer/draft_training_loop.py | 9 +- verl_speco/trainer/feature_store.py | 250 +++++++++++++++++++- verl_speco/vllm_hidden_states_generate.py | 5 +- 6 files changed, 638 insertions(+), 9 deletions(-) create mode 100644 verl_speco/inspect_jsonl_samples.py diff --git a/tests/unit/test_draft_feature_store.py b/tests/unit/test_draft_feature_store.py index 0cfdee9a..6f098bdf 100644 --- a/tests/unit/test_draft_feature_store.py +++ b/tests/unit/test_draft_feature_store.py @@ -12,6 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. import importlib +import json import pytest @@ -23,8 +24,11 @@ DraftFeatureDataLoaderConfig = draft_dataset.DraftFeatureDataLoaderConfig DraftFeatureSample = feature_store.DraftFeatureSample DraftReplaySample = feature_store.DraftReplaySample +JsonlTokenReplayFeatureStore = feature_store.JsonlTokenReplayFeatureStore TokenReplayFeatureStore = feature_store.TokenReplayFeatureStore TorchShardFeatureStore = feature_store.TorchShardFeatureStore +VllmSafetensorsFeatureStore = feature_store.VllmSafetensorsFeatureStore +build_feature_store_from_config = feature_store.build_feature_store_from_config def _sample(index: int = 0): @@ -103,6 +107,119 @@ def test_token_replay_rejects_non_contiguous_feature_positions(): sample.validate(strict=True) +def test_jsonl_token_replay_reads_input_ids_and_loss_mask(tmp_path): + path = tmp_path / "samples.jsonl" + row = { + "id": "sample-0", + "input_ids": list(range(10)), + "loss_mask": [0, 0, 0, 0, 1, 1, 1, 0, 0, 0], + "text": "ignored for replay", + "metadata": {"source": "unit"}, + } + path.write_text(json.dumps(row) + "\n", encoding="utf-8") + + store = JsonlTokenReplayFeatureStore(path, read_only=True, max_seq_len=4) + keys = list(store.iter_keys(shuffle=False)) + loaded = store.read(keys[0]) + + assert keys == ["samples.jsonl:0"] + assert loaded.algorithm == "EAGLE3" + assert torch.equal(loaded.input_ids, torch.arange(10)) + assert torch.equal(loaded.loss_mask, torch.tensor(row["loss_mask"], dtype=torch.float32)) + assert torch.equal(loaded.attention_mask, torch.ones(10, dtype=torch.bool)) + assert torch.equal(loaded.position_ids, torch.arange(10)) + assert torch.equal(loaded.feature_positions, torch.arange(3, 7)) + assert torch.equal(loaded.draft_position_ids, torch.arange(4, 8)) + assert loaded.metadata["source"] == "jsonl_token_replay" + assert loaded.metadata["id"] == "sample-0" + assert store.get_metadata()["format"] == "jsonl_token_replay" + assert store.get_metadata()["num_samples"] == 1 + + +def test_build_feature_store_from_config_supports_jsonl_token_replay(tmp_path): + path = tmp_path / "samples.jsonl" + path.write_text( + json.dumps({"input_ids": [1, 2, 3], "loss_mask": [0, 1, 1]}) + "\n", + encoding="utf-8", + ) + + store = build_feature_store_from_config( + { + "type": "jsonl_token_replay", + "path": path, + "max_seq_len": 8, + }, + read_only=True, + ) + loaded = store.read(next(store.iter_keys(shuffle=False))) + + assert isinstance(loaded, DraftReplaySample) + assert torch.equal(loaded.feature_positions, torch.arange(0, 3)) + + +def test_vllm_safetensors_feature_store_records_manifest_path_and_roundtrips(tmp_path): + pytest.importorskip("safetensors.torch") + hidden_positions = torch.arange(160, dtype=torch.long) + sample = DraftFeatureSample( + algorithm="DSPARK", + input_ids=torch.arange(160, dtype=torch.long), + loss_mask=torch.ones(160, dtype=torch.float32), + hidden_states=torch.randn(160, 16, dtype=torch.float32), + position_ids=torch.arange(10, 170, dtype=torch.long), + metadata={ + "source": "token_replay_vllm_file", + "global_step": 7, + "hidden_states_layout": "dflash_aux_plus_last", + "hidden_positions": hidden_positions, + "feature_start": 10, + "feature_end": 170, + }, + ) + store = VllmSafetensorsFeatureStore(tmp_path) + + keys = store.write_many([sample]) + store.close() + + manifest_lines = (tmp_path / "manifest.jsonl").read_text(encoding="utf-8").splitlines() + assert len(manifest_lines) == 1 + entry = json.loads(manifest_lines[0]) + assert entry["path"].endswith(".safetensors") + assert (tmp_path / entry["path"]).exists() + assert entry["sample"]["metadata"]["hidden_positions"] == { + "__tensor__": True, + "dtype": "torch.int64", + "shape": [160], + } + + reader = VllmSafetensorsFeatureStore(tmp_path, read_only=True) + loaded = reader.read(keys[0]) + + assert loaded.algorithm == "DSPARK" + assert torch.equal(loaded.input_ids, sample.input_ids) + assert torch.equal(loaded.position_ids, sample.position_ids) + assert torch.equal(loaded.hidden_states, sample.hidden_states) + assert torch.equal(loaded.metadata["hidden_positions"], hidden_positions) + assert reader.get_metadata()["format"] == "vllm_safetensors" + assert reader.get_metadata()["num_samples"] == 1 + + +def test_build_feature_store_from_config_supports_vllm_safetensors(tmp_path): + pytest.importorskip("safetensors.torch") + writer = build_feature_store_from_config( + {"type": "vllm_safetensors", "path": tmp_path} + ) + writer.write_many([_sample(0)]) + writer.close() + + reader = build_feature_store_from_config( + {"type": "vllm_safetensors", "path": tmp_path}, read_only=True + ) + loaded = reader.read(next(reader.iter_keys(shuffle=False))) + + assert isinstance(loaded, DraftFeatureSample) + assert torch.equal(loaded.input_ids, torch.tensor([1, 2, 3, 4])) + + def test_feature_sample_normalizes_singleton_position_ids(): sample = DraftFeatureSample( input_ids=torch.tensor([1, 2, 3, 4], dtype=torch.long), diff --git a/verl_speco/config/draft_trainer.yaml b/verl_speco/config/draft_trainer.yaml index 6ba60569..b5593b15 100644 --- a/verl_speco/config/draft_trainer.yaml +++ b/verl_speco/config/draft_trainer.yaml @@ -61,6 +61,7 @@ actor_rollout_ref: on_generate: delete require_arange_positions: true offline_generation: + input_type: token_replay input_path: null output_path: null max_samples: 0 diff --git a/verl_speco/inspect_jsonl_samples.py b/verl_speco/inspect_jsonl_samples.py new file mode 100644 index 00000000..217d638d --- /dev/null +++ b/verl_speco/inspect_jsonl_samples.py @@ -0,0 +1,265 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Inspect JSONL draft-training samples without loading model dependencies.""" + +from __future__ import annotations + +import argparse +import json +from collections import Counter, defaultdict +from pathlib import Path +from typing import Any + + +TOKEN_REPLAY_REQUIRED_KEYS = { + "input_ids", + "loss_mask", + "attention_mask", + "position_ids", + "feature_positions", + "draft_position_ids", +} +FEATURE_REQUIRED_KEYS = {"input_ids", "loss_mask", "hidden_states"} +INPUT_LOSS_REQUIRED_KEYS = {"input_ids", "loss_mask"} +VLLM_SAFETENSORS_MANIFEST_KEYS = {"path", "num_samples", "sample"} + + +def main() -> int: + parser = argparse.ArgumentParser( + description="Inspect a JSONL file and report whether it matches SPECO sample schemas." + ) + parser.add_argument("path", help="JSONL file to inspect.") + parser.add_argument( + "--max-lines", + type=int, + default=20, + help="Maximum number of JSONL rows to inspect.", + ) + parser.add_argument( + "--show-first", + action="store_true", + help="Print the first JSON object with long arrays summarized.", + ) + parser.add_argument( + "--strict-exit", + action="store_true", + help="Exit with code 1 when inspected rows do not match a known schema.", + ) + args = parser.parse_args() + + path = Path(args.path) + summaries = [] + key_counts: Counter[str] = Counter() + schema_counts: Counter[str] = Counter() + issues_by_schema: dict[str, list[str]] = defaultdict(list) + first_obj: dict[str, Any] | None = None + + with path.open(encoding="utf-8") as jsonl_file: + for line_number, line in enumerate(jsonl_file, start=1): + if len(summaries) >= int(args.max_lines): + break + line = line.strip() + if not line: + continue + try: + obj = json.loads(line) + except json.JSONDecodeError as exc: + summaries.append({"line": line_number, "schema": "invalid_json"}) + schema_counts["invalid_json"] += 1 + issues_by_schema["invalid_json"].append(f"line {line_number}: {exc}") + continue + if first_obj is None and isinstance(obj, dict): + first_obj = obj + if not isinstance(obj, dict): + summaries.append({"line": line_number, "schema": type(obj).__name__}) + schema_counts[type(obj).__name__] += 1 + continue + keys = set(obj) + key_counts.update(keys) + schema, issues = _classify(obj) + schema_counts[schema] += 1 + issues_by_schema[schema].extend( + f"line {line_number}: {issue}" for issue in issues + ) + summaries.append( + { + "line": line_number, + "schema": schema, + "keys": sorted(keys), + "shapes": _shape_summary(obj), + } + ) + + print(f"jsonl={path}") + print(f"inspected_lines={len(summaries)}") + print("schema_counts:") + for schema, count in schema_counts.most_common(): + print(f" {schema}: {count}") + print("key_counts:") + for key, count in key_counts.most_common(): + print(f" {key}: {count}") + print("sample_summaries:") + for summary in summaries[: min(len(summaries), 5)]: + print(json.dumps(summary, ensure_ascii=False, sort_keys=True)) + if issues_by_schema: + print("issues:") + for schema, issues in issues_by_schema.items(): + for issue in issues[:10]: + print(f" [{schema}] {issue}") + if args.show_first and first_obj is not None: + print("first_object:") + print( + json.dumps( + _compact_json(first_obj), ensure_ascii=False, indent=2, sort_keys=True + ) + ) + + invalid = any( + schema + not in { + "token_replay_jsonl", + "feature_jsonl", + "input_loss_jsonl", + "vllm_safetensors_manifest", + } + for schema in schema_counts + ) + return 1 if invalid and args.strict_exit else 0 + + +def _classify(obj: dict[str, Any]) -> tuple[str, list[str]]: + keys = set(obj) + if VLLM_SAFETENSORS_MANIFEST_KEYS.issubset(keys): + return "vllm_safetensors_manifest", [] + if TOKEN_REPLAY_REQUIRED_KEYS.issubset(keys): + return "token_replay_jsonl", _token_replay_issues(obj) + if FEATURE_REQUIRED_KEYS.issubset(keys): + return "feature_jsonl", _feature_issues(obj) + if INPUT_LOSS_REQUIRED_KEYS.issubset(keys): + return "input_loss_jsonl", _input_loss_issues(obj) + missing_token = sorted(TOKEN_REPLAY_REQUIRED_KEYS - keys) + missing_feature = sorted(FEATURE_REQUIRED_KEYS - keys) + return ( + "unknown", + [ + f"missing token_replay keys={missing_token}", + f"missing feature keys={missing_feature}", + ], + ) + + +def _token_replay_issues(obj: dict[str, Any]) -> list[str]: + issues = [] + input_len = _flat_len(obj.get("input_ids")) + for key in ("loss_mask", "attention_mask", "position_ids"): + value_len = _flat_len(obj.get(key)) + if input_len is not None and value_len != input_len: + issues.append(f"{key} length {value_len} != input_ids length {input_len}") + feature_len = _flat_len(obj.get("feature_positions")) + draft_len = _flat_len(obj.get("draft_position_ids")) + if feature_len is not None and draft_len != feature_len: + issues.append( + f"draft_position_ids length {draft_len} != feature_positions length {feature_len}" + ) + return issues + + +def _feature_issues(obj: dict[str, Any]) -> list[str]: + issues = [] + input_len = _flat_len(obj.get("input_ids")) + loss_len = _flat_len(obj.get("loss_mask")) + if input_len is not None and loss_len != input_len: + issues.append(f"loss_mask length {loss_len} != input_ids length {input_len}") + hidden_shape = _shape(obj.get("hidden_states")) + if not hidden_shape or hidden_shape[0] in {"scalar", "dict", "str", "none"}: + issues.append("hidden_states is not an array-like value") + return issues + + +def _input_loss_issues(obj: dict[str, Any]) -> list[str]: + issues = [] + input_len = _flat_len(obj.get("input_ids")) + loss_len = _flat_len(obj.get("loss_mask")) + if input_len is None: + issues.append("input_ids is not a JSON list") + if loss_len is None: + issues.append("loss_mask is not a JSON list") + if input_len is not None and loss_len is not None and loss_len != input_len: + issues.append(f"loss_mask length {loss_len} != input_ids length {input_len}") + return issues + + +def _shape_summary(obj: dict[str, Any]) -> dict[str, Any]: + return { + key: _shape(value) + for key, value in obj.items() + if key + in { + "input_ids", + "loss_mask", + "attention_mask", + "position_ids", + "feature_positions", + "draft_position_ids", + "hidden_states", + "target_logprobs", + "sample", + "path", + } + } + + +def _shape(value: Any) -> list[Any]: + if value is None: + return ["none"] + if isinstance(value, dict): + return ["dict", sorted(value)[:20]] + if isinstance(value, str): + return ["str", len(value)] + if not isinstance(value, list): + return ["scalar", type(value).__name__] + shape = [] + current = value + while isinstance(current, list): + shape.append(len(current)) + current = current[0] if current else None + return shape + + +def _flat_len(value: Any) -> int | None: + if not isinstance(value, list): + return None + current = value + while isinstance(current, list) and len(current) == 1: + current = current[0] + return len(current) if isinstance(current, list) else None + + +def _compact_json(value: Any) -> Any: + if isinstance(value, dict): + return {key: _compact_json(item) for key, item in value.items()} + if isinstance(value, list): + if len(value) > 16: + return { + "__list__": True, + "shape": _shape(value), + "head": [_compact_json(item) for item in value[:4]], + "tail": [_compact_json(item) for item in value[-4:]], + } + return [_compact_json(item) for item in value] + return value + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/verl_speco/trainer/draft_training_loop.py b/verl_speco/trainer/draft_training_loop.py index b3ee29bb..b471a5ac 100644 --- a/verl_speco/trainer/draft_training_loop.py +++ b/verl_speco/trainer/draft_training_loop.py @@ -70,14 +70,15 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: training_mode = ( str(training_cfg.get("mode", "offline") or "offline").strip().lower() ) + replay_feature_store_types = {"token_replay", "jsonl_token_replay", "jsonl"} if not feature_store_cfg.get("path"): raise ValueError( "actor_rollout_ref.rollout.drafter.training.feature_store.path is required" ) - if feature_store_type == "token_replay" and training_mode != "offline": + if feature_store_type in replay_feature_store_types and training_mode != "offline": raise ValueError( - "feature_store.type=token_replay is supported only by standalone " - "training.mode=offline" + f"feature_store.type={feature_store_type} is supported only by " + "standalone training.mode=offline" ) _disable_standalone_sequence_parallel(draft_config) @@ -149,7 +150,7 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: rank, time.perf_counter() - stage_started, ) - if feature_store_type == "token_replay": + if feature_store_type in replay_feature_store_types: # Keep the large target model entirely outside online training imports # and lifetime. The standalone loop materializes ordinary feature # samples before handing them to the shared trainer. diff --git a/verl_speco/trainer/feature_store.py b/verl_speco/trainer/feature_store.py index cea2d862..0b5838e2 100644 --- a/verl_speco/trainer/feature_store.py +++ b/verl_speco/trainer/feature_store.py @@ -598,6 +598,190 @@ def read(self, key: str) -> DraftReplaySample: return DraftReplaySample.from_dict(sample, strict=self.strict_schema) +class JsonlTokenReplayFeatureStore: + """Read ``input_ids``/``loss_mask`` JSONL rows as token replay samples.""" + + def __init__( + self, + path: str | os.PathLike[str], + *, + max_samples_per_shard: int = 1024, + metadata: dict[str, Any] | None = None, + strict_schema: bool = True, + read_only: bool = False, + shard_prefix: str = "shard", + max_seq_len: int = 512, + window_mode: str = "loss", + ): + if path is None: + raise ValueError("JsonlTokenReplayFeatureStore requires a non-empty path") + if not read_only: + raise RuntimeError("JsonlTokenReplayFeatureStore is read-only") + self.path = Path(path) + self.max_samples_per_shard = max(int(max_samples_per_shard), 1) + self.metadata = { + "schema_version": SCHEMA_VERSION, + "format": "jsonl_token_replay", + "created_by": "verl_speco", + "created_at": time.time(), + } + if metadata: + self.metadata.update(metadata) + self.strict_schema = bool(strict_schema) + self.read_only = bool(read_only) + self.shard_prefix = str(shard_prefix or "shard") + self.max_seq_len = int(max_seq_len or 0) + self.window_mode = str(window_mode or "loss").strip().lower() + if self.window_mode not in {"loss", "front", "full"}: + raise ValueError( + "feature_store.window_mode for jsonl_token_replay must be " + "'loss', 'front' or 'full'" + ) + self._files = self._resolve_jsonl_files() + self._line_offsets: dict[str, list[int]] = {} + self._keys = self._build_keys() + + def write_many( + self, samples: list[DraftStoredSample | dict[str, Any]] + ) -> list[str]: + raise RuntimeError("JsonlTokenReplayFeatureStore is read-only") + + def read(self, key: str) -> DraftReplaySample: + file_name, row_index = _parse_key(key) + file_path = self.path / file_name if self.path.is_dir() else self.path + offsets = self._line_offsets.get(file_name) + if offsets is None: + self._keys = self._build_keys() + offsets = self._line_offsets.get(file_name) + if offsets is None or int(row_index) >= len(offsets): + raise IndexError(f"JSONL row {row_index} not found in {file_path}") + payload = _load_jsonl_offset(file_path, offsets[int(row_index)]) + return self._row_to_replay_sample(payload, file_name, int(row_index)) + + def iter_keys(self, *, shuffle: bool = False, seed: int = 0) -> Iterator[str]: + keys = list(self._keys) + if shuffle: + random.Random(int(seed)).shuffle(keys) + yield from keys + + def get_metadata(self) -> dict[str, Any]: + metadata = dict(self.metadata) + metadata.update( + { + "num_files": len(self._files), + "num_samples": len(self._keys), + "max_seq_len": self.max_seq_len, + "window_mode": self.window_mode, + } + ) + return metadata + + def close(self) -> None: + return + + def _resolve_jsonl_files(self) -> list[Path]: + if self.path.is_file(): + return [self.path] + if self.path.is_dir(): + files = sorted(self.path.glob("*.jsonl")) + if files: + return files + raise FileNotFoundError(f"No JSONL file found at {self.path}") + + def _build_keys(self) -> list[str]: + keys: list[str] = [] + self._line_offsets = {} + for file_path in self._files: + file_key = ( + file_path.relative_to(self.path).as_posix() + if self.path.is_dir() + else file_path.name + ) + offsets: list[int] = [] + with file_path.open("rb") as jsonl_file: + while True: + offset = int(jsonl_file.tell()) + line = jsonl_file.readline() + if not line: + break + if line.strip(): + row_index = len(offsets) + offsets.append(offset) + keys.append(f"{file_key}:{row_index}") + self._line_offsets[file_key] = offsets + return keys + + def _row_to_replay_sample( + self, payload: dict[str, Any], file_name: str, line_index: int + ) -> DraftReplaySample: + input_ids = _json_list_tensor(payload, "input_ids", dtype=torch.long) + loss_mask = _json_list_tensor(payload, "loss_mask", dtype=torch.float32) + if int(input_ids.numel()) != int(loss_mask.numel()): + raise ValueError( + "jsonl_token_replay input_ids/loss_mask length mismatch: " + f"{int(input_ids.numel())} vs {int(loss_mask.numel())}" + ) + attention_mask = _optional_json_list_tensor( + payload, "attention_mask", dtype=torch.bool + ) + if attention_mask is None: + attention_mask = torch.ones_like(input_ids, dtype=torch.bool) + position_ids = _optional_json_list_tensor( + payload, "position_ids", dtype=torch.long + ) + if position_ids is None: + position_ids = attention_mask.long().cumsum(dim=0).sub(1).clamp_min(0) + + sequence_length = int(input_ids.numel()) + start, end = self._feature_window(loss_mask) + feature_positions = torch.arange(start, end, dtype=torch.long) + draft_position_ids = position_ids[start:end].long() + 1 + metadata = { + "source": "jsonl_token_replay", + "jsonl_path": file_name, + "jsonl_line": line_index, + "sequence_length": end - start, + "full_sequence_length": sequence_length, + "feature_start": start, + "feature_end": end, + "loss_tokens": int(loss_mask[start:end].sum().item()), + } + for key in ("id", "hash", "primary_id", "finish_reason"): + if key in payload: + metadata[key] = payload[key] + return DraftReplaySample( + algorithm=str(payload.get("algorithm", "EAGLE3")), + input_ids=input_ids, + loss_mask=loss_mask, + attention_mask=attention_mask, + position_ids=position_ids, + feature_positions=feature_positions, + draft_position_ids=draft_position_ids, + metadata=metadata, + ) + + def _feature_window(self, loss_mask: torch.Tensor) -> tuple[int, int]: + sequence_length = int(loss_mask.numel()) + if sequence_length <= 0: + raise ValueError("jsonl_token_replay input_ids must not be empty") + max_seq_len = self.max_seq_len if self.max_seq_len > 0 else sequence_length + max_seq_len = min(max_seq_len, sequence_length) + if self.window_mode == "full": + return 0, sequence_length + if self.window_mode == "front": + return 0, max_seq_len + + active = torch.nonzero(loss_mask.float() > 0, as_tuple=False).reshape(-1) + if int(active.numel()) <= 0: + return 0, max_seq_len + first_loss = int(active[0].item()) + start = max(first_loss - 1, 0) + end = min(start + max_seq_len, sequence_length) + if end <= start: + end = min(start + 1, sequence_length) + return start, end + + class VllmSafetensorsFeatureStore(TorchShardFeatureStore): """Feature store for vLLM-extracted hidden states saved as safetensors. @@ -710,6 +894,10 @@ def read(self, key: str) -> DraftFeatureSample: "hidden_states": tensors["hidden_states"], "metadata": dict(manifest_sample.get("metadata") or {}), } + if "metadata.hidden_positions" in tensors: + payload["metadata"]["hidden_positions"] = tensors[ + "metadata.hidden_positions" + ].long() for optional_key in ( "last_hidden_states", "target", @@ -730,10 +918,15 @@ def _sample_to_safetensors( self, sample: DraftFeatureSample ) -> dict[str, torch.Tensor]: payload = sample.to_dict() + hidden_states = payload["hidden_states"] + if not torch.is_tensor(hidden_states): + raise TypeError( + "feature_store.type=vllm_safetensors requires tensor hidden_states" + ) tensors = { "input_ids": payload["input_ids"].long().contiguous(), "loss_mask": payload["loss_mask"].float().contiguous(), - "hidden_states": payload["hidden_states"].contiguous(), + "hidden_states": hidden_states.contiguous(), } for optional_key in ( "last_hidden_states", @@ -744,6 +937,10 @@ def _sample_to_safetensors( value = payload.get(optional_key) if torch.is_tensor(value): tensors[optional_key] = value.contiguous() + metadata = payload.get("metadata") or {} + hidden_positions = metadata.get("hidden_positions") + if torch.is_tensor(hidden_positions): + tensors["metadata.hidden_positions"] = hidden_positions.long().contiguous() return tensors def _sample_manifest(self, sample: DraftFeatureSample) -> dict[str, Any]: @@ -766,13 +963,25 @@ def build_feature_store_from_config( .strip() .lower() ) - store_cls: type[TorchShardFeatureStore] if store_type == "torch_shard": - store_cls = TorchShardFeatureStore + store_cls: type[TorchShardFeatureStore] = TorchShardFeatureStore elif store_type == "token_replay": store_cls = TokenReplayFeatureStore elif store_type in {"vllm_safetensors", "safetensors"}: store_cls = VllmSafetensorsFeatureStore + elif store_type in {"jsonl_token_replay", "jsonl"}: + return JsonlTokenReplayFeatureStore( + feature_store_cfg.get("path"), + max_samples_per_shard=int( + feature_store_cfg.get("max_samples_per_shard", 1024) + ), + metadata=metadata, + strict_schema=bool(feature_store_cfg.get("strict_schema", True)), + read_only=read_only, + shard_prefix=shard_prefix, + max_seq_len=int(feature_store_cfg.get("max_seq_len", 512) or 0), + window_mode=str(feature_store_cfg.get("window_mode", "loss") or "loss"), + ) else: raise NotImplementedError(f"Unsupported draft feature store type: {store_type}") return store_cls( @@ -870,6 +1079,41 @@ def _atomic_torch_save(payload: dict[str, Any], path: Path) -> None: os.remove(tmp_name) +def _load_jsonl_offset(path: Path, offset: int) -> dict[str, Any]: + with path.open("rb") as jsonl_file: + jsonl_file.seek(int(offset)) + line = jsonl_file.readline().decode("utf-8").strip() + if not line: + raise ValueError(f"JSONL offset {offset} in {path} points to an empty line") + payload = json.loads(line) + if not isinstance(payload, dict): + raise TypeError( + f"JSONL offset {offset} in {path} must contain a JSON object" + ) + return payload + + +def _json_list_tensor( + payload: dict[str, Any], key: str, *, dtype: torch.dtype +) -> torch.Tensor: + if key not in payload: + raise KeyError(f"jsonl_token_replay sample missing required key {key!r}") + value = payload[key] + if not isinstance(value, list): + raise TypeError( + f"jsonl_token_replay {key} must be a JSON list, got {type(value).__name__}" + ) + return torch.tensor(value, dtype=dtype).reshape(-1) + + +def _optional_json_list_tensor( + payload: dict[str, Any], key: str, *, dtype: torch.dtype +) -> torch.Tensor | None: + if key not in payload or payload[key] is None: + return None + return _json_list_tensor(payload, key, dtype=dtype) + + def _json_safe_metadata(value: Any) -> Any: if torch.is_tensor(value): if value.numel() <= 128: diff --git a/verl_speco/vllm_hidden_states_generate.py b/verl_speco/vllm_hidden_states_generate.py index c775ca5e..d7e78e10 100644 --- a/verl_speco/vllm_hidden_states_generate.py +++ b/verl_speco/vllm_hidden_states_generate.py @@ -61,8 +61,9 @@ def generate_vllm_safetensors_features(config) -> dict[str, Any]: "training.feature_store.path is required" ) + input_type = str(generation_cfg.get("input_type", "token_replay") or "token_replay") input_cfg = _plain_config(feature_store_cfg) - input_cfg.update({"type": "token_replay", "path": os.fspath(input_path)}) + input_cfg.update({"type": input_type, "path": os.fspath(input_path)}) output_cfg = _plain_config(feature_store_cfg) output_cfg.update({"type": "vllm_safetensors", "path": os.fspath(output_path)}) @@ -84,7 +85,7 @@ def generate_vllm_safetensors_features(config) -> dict[str, Any]: output_cfg, read_only=False, metadata={ - "source_format": "token_replay", + "source_format": input_type, "source_path": os.fspath(input_path), "target_feature_backend": "vllm_file", }, From 3d12c4cd2061a5f2e0f369ecad1997b2359a0833 Mon Sep 17 00:00:00 2001 From: Cai Zeyong <878049625@qq.com> Date: Fri, 31 Jul 2026 17:06:26 +0800 Subject: [PATCH 26/50] =?UTF-8?q?TQ=20=E5=88=86=E6=94=AF=E5=8A=A0=E4=BA=86?= =?UTF-8?q?=20=5Fdensify=5Ftq=5Ftensor=EF=BC=8C=E6=8A=8A=20TQ=20=E5=8F=96?= =?UTF-8?q?=E5=9B=9E=E7=9A=84=20NestedTensor=20=E8=BF=98=E5=8E=9F=E6=88=90?= =?UTF-8?q?=E7=94=9F=E4=BA=A7=E7=AB=AF=E5=AD=98=E8=BF=9B=E5=8E=BB=E6=97=B6?= =?UTF-8?q?=E7=9A=84=20[rows,=20hidden]=20=E5=AF=86=E9=9B=86=E5=BD=A2?= =?UTF-8?q?=E5=BC=8F=EF=BC=88=E7=94=A8=20unbind+cat?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- verl_speco/workers/speco_worker.py | 28 +++++++++++++++++++++++++++- 1 file changed, 27 insertions(+), 1 deletion(-) diff --git a/verl_speco/workers/speco_worker.py b/verl_speco/workers/speco_worker.py index 6afc2429..8f192004 100644 --- a/verl_speco/workers/speco_worker.py +++ b/verl_speco/workers/speco_worker.py @@ -53,6 +53,32 @@ def _resolve_ray_object_ref(value): return value +def _densify_tq_tensor(tensor): + """Unwrap a tensor returned by TransferQueue into a plain dense tensor. + + TQ stores each ``put_sample`` payload inside a TensorDict and + ``kv_batch_get`` returns it with an added batch dimension, as a NestedTensor + (jagged on dim 0). The old-logprob chunk resolver slices + ``tensor[start:start + length]`` on dim 0, which NestedTensor does not + support (``slice(): not supported for NestedTensor on dim=0``). The producer + put a dense ``[rows, hidden]`` tensor, so flatten the NestedTensor back to + that 2-D form. Mirrors ``_speco_tensor_rows`` which uses ``tensor.unbind()`` + for the same nested-tensor case. + """ + if not torch.is_tensor(tensor): + return tensor + if tensor.is_nested: + parts = [p for p in tensor.unbind() if p.numel() > 0] + if not parts: + return None + tensor = torch.cat(parts, dim=0) + if tensor.dim() == 3: + tensor = tensor.squeeze(0) + elif tensor.dim() == 1: + tensor = tensor.unsqueeze(0) + return tensor.contiguous() + + def _resolve_tq_or_ray_ref(ref): # P1: old-logprob chunk refs may be TransferQueue keys ("speco:" prefix) # instead of Ray ObjectRefs. Fetch the stored chunk via TQ and fall back to @@ -61,7 +87,7 @@ def _resolve_tq_or_ray_ref(ref): if isinstance(ref, str) and ref.startswith("speco:"): from verl_speco.integration.transferqueue_bridge import get_sample - return get_sample(ref).get("hidden") + return _densify_tq_tensor(get_sample(ref).get("hidden")) return _resolve_ray_object_ref(ref) From adb86160123bfd982f08ef003db5365d33242fa5 Mon Sep 17 00:00:00 2001 From: Cai Zeyong <878049625@qq.com> Date: Mon, 3 Aug 2026 15:47:09 +0800 Subject: [PATCH 27/50] =?UTF-8?q?=E4=BF=AE=E5=A4=8Dtq=20=E5=A4=9A=E6=AC=A1?= =?UTF-8?q?=E4=BB=8Ecpu=E5=8F=96hidden=E7=9A=84bug=EF=BC=8C=E6=94=B9?= =?UTF-8?q?=E6=88=90=E5=8F=96=E4=B8=80=E6=AC=A1=EF=BC=8C=E5=89=A9=E4=BD=99?= =?UTF-8?q?=E4=BB=8Ecache=E5=8F=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- verl_speco/workers/speco_worker.py | 22 ++++++++++++++++------ 1 file changed, 16 insertions(+), 6 deletions(-) diff --git a/verl_speco/workers/speco_worker.py b/verl_speco/workers/speco_worker.py index 8f192004..2055bbed 100644 --- a/verl_speco/workers/speco_worker.py +++ b/verl_speco/workers/speco_worker.py @@ -91,10 +91,16 @@ def _resolve_tq_or_ray_ref(ref): return _resolve_ray_object_ref(ref) -def _resolve_hidden_state_chunks(chunks, expected_rows: int | None = None): +def _resolve_hidden_state_chunks(chunks, expected_rows: int | None = None, cache=None): if not chunks: return None - resolved_cache = {} + # Cross-sample cache: callers pass a per-step cache so multiple samples + # sharing the same owner chunk (same TQ key / ObjectRef) are fetched only + # once. Without this, ~16 samples in one owner each trigger a full TQ + # kv_batch_get (or ray.get) on the same ~400MB chunk. ray.get dedups at + # the object store; TQ get_sample does NOT, so the cache is essential. + if cache is None: + cache = {} pieces = [] full_rows = int(expected_rows or 0) hidden_size = None @@ -106,9 +112,9 @@ def _resolve_hidden_state_chunks(chunks, expected_rows: int | None = None): if ref is None: continue cache_key = ref if isinstance(ref, str) else id(ref) - if cache_key not in resolved_cache: - resolved_cache[cache_key] = _resolve_tq_or_ray_ref(ref) - tensor = resolved_cache[cache_key] + if cache_key not in cache: + cache[cache_key] = _resolve_tq_or_ray_ref(ref) + tensor = cache[cache_key] if not torch.is_tensor(tensor): continue start = int(chunk.get("chunk_start", 0) or 0) @@ -715,6 +721,10 @@ def _flush_rollout_features_for_step(self) -> None: def collect_rollout_features(self, samples: list[dict]): if not samples: return + # Per-step cross-sample cache for chunk fetches. Reset every collect + # call so keys (which carry global_step) never stale and the cache + # cannot grow unbounded. Shared across all samples in this step. + self._tq_chunk_cache = {} for sample in samples: if not sample: continue @@ -783,7 +793,7 @@ def collect_rollout_features(self, samples: list[dict]): hidden_positions = batch.get("hidden_positions") if torch.is_tensor(hidden_positions): expected_rows = int(hidden_positions.numel()) - hidden = _resolve_hidden_state_chunks(hidden_chunks, expected_rows=expected_rows) + hidden = _resolve_hidden_state_chunks(hidden_chunks, expected_rows=expected_rows, cache=self._tq_chunk_cache) else: hidden = _resolve_ray_object_ref(sample.get("hidden_states_ref")) target_logprobs = sample.get("target_logprobs") From 9d878c5ae6cc64ccfbf97d22a9d21f2a8b2f541b Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Tue, 4 Aug 2026 16:07:47 +0800 Subject: [PATCH 28/50] Support conversation JSONL token replay for standalone draft training --- tests/unit/test_draft_feature_store.py | 43 +++++++ tests/unit/test_draft_training_loop.py | 8 ++ tests/unit/test_target_feature_replay.py | 76 +++++++++++- verl_speco/config/speco_base.yaml | 5 + verl_speco/trainer/base_trainer.py | 2 + verl_speco/trainer/draft_training_loop.py | 23 ++++ verl_speco/trainer/feature_store.py | 130 +++++++++++++++++++- verl_speco/trainer/target_feature_replay.py | 85 ++++++++++--- 8 files changed, 352 insertions(+), 20 deletions(-) diff --git a/tests/unit/test_draft_feature_store.py b/tests/unit/test_draft_feature_store.py index 6f098bdf..458798c5 100644 --- a/tests/unit/test_draft_feature_store.py +++ b/tests/unit/test_draft_feature_store.py @@ -157,6 +157,49 @@ def test_build_feature_store_from_config_supports_jsonl_token_replay(tmp_path): assert torch.equal(loaded.feature_positions, torch.arange(0, 3)) +def test_jsonl_token_replay_reads_conversations_with_chat_template(tmp_path): + class FakeTokenizer: + def apply_chat_template( + self, messages, *, tokenize, add_generation_prompt + ): + assert tokenize is True + token_ids = [] + for message in messages: + role_id = {"user": 10, "assistant": 20, "system": 30}[message["role"]] + token_ids.extend([role_id, len(message["content"])]) + if add_generation_prompt: + token_ids.append(20) + return token_ids + + path = tmp_path / "samples.jsonl" + row = { + "id": "conv-0", + "conversations": [ + {"from": "human", "value": "question"}, + {"from": "assistant", "value": "answer"}, + ], + "algorithm": "DSPARK", + } + path.write_text(json.dumps(row) + "\n", encoding="utf-8") + + store = JsonlTokenReplayFeatureStore( + path, + read_only=True, + max_seq_len=8, + tokenizer_path="/target", + ) + store._tokenizer = FakeTokenizer() + loaded = store.read(next(store.iter_keys(shuffle=False))) + + assert loaded.algorithm == "DSPARK" + assert torch.equal(loaded.input_ids, torch.tensor([10, 8, 20, 6])) + assert torch.equal(loaded.loss_mask, torch.tensor([0.0, 0.0, 0.0, 1.0])) + assert torch.equal(loaded.feature_positions, torch.arange(2, 4)) + assert torch.equal(loaded.draft_position_ids, torch.arange(3, 5)) + assert loaded.metadata["source"] == "jsonl_conversations" + assert loaded.metadata["id"] == "conv-0" + + def test_vllm_safetensors_feature_store_records_manifest_path_and_roundtrips(tmp_path): pytest.importorskip("safetensors.torch") hidden_positions = torch.arange(160, dtype=torch.long) diff --git a/tests/unit/test_draft_training_loop.py b/tests/unit/test_draft_training_loop.py index af34631b..2faa009a 100644 --- a/tests/unit/test_draft_training_loop.py +++ b/tests/unit/test_draft_training_loop.py @@ -23,6 +23,7 @@ from verl_speco.trainer.standalone_checkpoint import rewrite_standalone_runtime_config # noqa: E402 from verl_speco.trainer.draft_training_loop import ( # noqa: E402 + _is_out_of_memory_error, _rewrite_standalone_block_runtime_config, _save_standalone_checkpoint, _should_log_batch_progress, @@ -60,6 +61,13 @@ def test_should_log_standalone_batch_progress(attempted_batches, expected): assert _should_log_batch_progress(attempted_batches) is expected +def test_is_out_of_memory_error_matches_npu_oom_message(): + error = RuntimeError("NPU out of memory. Tried to allocate 258.00 MiB") + + assert _is_out_of_memory_error(error) + assert not _is_out_of_memory_error(RuntimeError("bad batch")) + + def test_standalone_checkpoint_schedules_without_waiting(): trainer = _FakeTrainer() diff --git a/tests/unit/test_target_feature_replay.py b/tests/unit/test_target_feature_replay.py index 9e7d9c6b..818a4338 100644 --- a/tests/unit/test_target_feature_replay.py +++ b/tests/unit/test_target_feature_replay.py @@ -16,9 +16,10 @@ torch = pytest.importorskip("torch") -from verl_speco.trainer.feature_store import DraftFeatureSample # noqa: E402 +from verl_speco.trainer.feature_store import DraftFeatureSample, DraftReplaySample # noqa: E402 from verl_speco.trainer.target_feature_replay import ( # noqa: E402 BoundedReplayCache, + TargetFeatureReplayer, _hidden_capture_target, ) @@ -65,3 +66,76 @@ def test_bounded_replay_cache_disables_zero_budget(tmp_path): assert cache.put("sample", _feature_sample()) is False assert cache.get("sample") is None + + +def test_token_replay_algorithm_mismatch_is_warning_not_error(caplog): + replayer = TargetFeatureReplayer.__new__(TargetFeatureReplayer) + replayer.rank = 0 + replayer.algorithm = "DFLASH" + replayer.target_layer_ids = [1, 3] + replayer.hidden_layout = "dflash_aux" + replayer.strict_target_model_path = False + replayer._warned_replay_algorithm_mismatch = False + replayer._warned_replay_layer_mismatch = False + replayer._warned_replay_layout_mismatch = False + sample = DraftReplaySample( + algorithm="DSPARK", + input_ids=torch.arange(8), + loss_mask=torch.ones(8), + attention_mask=torch.ones(8, dtype=torch.bool), + position_ids=torch.arange(8), + feature_positions=torch.arange(2, 6), + draft_position_ids=torch.arange(3, 7), + metadata={ + "target_layer_ids": [2, 4], + "hidden_states_layout": "dflash_aux_plus_last", + }, + ) + + replayer._validate_target_path(sample) + + assert "token replay algorithm differs" in caplog.text + assert "token replay target layers differ" in caplog.text + assert "token replay hidden layout differs" in caplog.text + + +def test_vllm_payload_maps_suffix_hidden_rows_to_absolute_positions(): + replayer = TargetFeatureReplayer.__new__(TargetFeatureReplayer) + replayer.rank = 0 + replayer.target_layer_ids = [1, 3] + replayer.hidden_layout = "dflash_aux_plus_last" + replayer.dtype = torch.float32 + replayer.model_path = "/target" + replayer.target_revision = None + replayer.target_config_fingerprint = "unit" + replayer.use_logits = False + + sample = DraftReplaySample( + algorithm="DSPARK", + input_ids=torch.arange(10, dtype=torch.long), + loss_mask=torch.ones(10, dtype=torch.float32), + attention_mask=torch.ones(10, dtype=torch.bool), + position_ids=torch.arange(10, dtype=torch.long), + feature_positions=torch.arange(4, 10, dtype=torch.long), + draft_position_ids=torch.arange(5, 11, dtype=torch.long), + metadata={"global_step": 1}, + ) + hidden = torch.arange(5 * 3 * 4, dtype=torch.float32).reshape(5, 3, 4) + payload = { + "token_ids": torch.arange(10, dtype=torch.long), + "hidden_states": hidden, + } + + feature = replayer._feature_from_vllm_payload( + sample, + payload, + prompt_ids=list(range(10)), + source="token_replay_vllm_file", + ) + + assert torch.equal(feature.input_ids, torch.arange(5, 10)) + assert torch.equal(feature.position_ids, torch.arange(6, 11)) + assert feature.hidden_states.shape == (5, 12) + assert feature.metadata["feature_start"] == 5 + assert feature.metadata["feature_end"] == 10 + assert feature.metadata["vllm_hidden_position_offset"] == 5 diff --git a/verl_speco/config/speco_base.yaml b/verl_speco/config/speco_base.yaml index ef798f74..e47e3eaf 100644 --- a/verl_speco/config/speco_base.yaml +++ b/verl_speco/config/speco_base.yaml @@ -181,3 +181,8 @@ actor_rollout_ref: repeat: true prefetch_depth: 2 strict_schema: true + max_seq_len: 512 + window_mode: loss + tokenizer_path: null + trust_remote_code: false + train_on: last_assistant diff --git a/verl_speco/trainer/base_trainer.py b/verl_speco/trainer/base_trainer.py index 317bb8d4..bb064f44 100644 --- a/verl_speco/trainer/base_trainer.py +++ b/verl_speco/trainer/base_trainer.py @@ -4230,10 +4230,12 @@ async def training_step_from_batch( self, batch: dict[str, torch.Tensor], step: int ) -> bool: """Execute one optimizer step from a pre-built standalone batch.""" + self.last_standalone_training_error = None try: with torch.enable_grad(): return await self._training_step_on_batch(batch, step) except Exception as e: # noqa: BLE001 + self.last_standalone_training_error = e logger.exception(f"Standalone training step {step} failed with error: {e}") return False diff --git a/verl_speco/trainer/draft_training_loop.py b/verl_speco/trainer/draft_training_loop.py index b471a5ac..a54bbc71 100644 --- a/verl_speco/trainer/draft_training_loop.py +++ b/verl_speco/trainer/draft_training_loop.py @@ -45,6 +45,13 @@ def _should_log_batch_progress(attempted_batches: int) -> bool: return attempted_batches <= 3 or attempted_batches % 100 == 0 +def _is_out_of_memory_error(error: BaseException) -> bool: + message = str(error).lower() + if "out of memory" in message or "oom" in message: + return True + return error.__class__.__name__ in {"OutOfMemoryError", "CudaOutOfMemoryError"} + + def run_standalone_draft_training(config) -> dict[str, Any]: """Run independent draft training from a feature store.""" return asyncio.run(_run_standalone_draft_training_async(config)) @@ -144,6 +151,14 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: feature_store_type, feature_store_cfg.get("path"), ) + if feature_store_type in {"jsonl_token_replay", "jsonl"} and not ( + feature_store_cfg.get("tokenizer_path") + ): + tokenizer_path = draft_config.actor_rollout_ref.model.path + try: + feature_store_cfg.tokenizer_path = tokenizer_path + except AttributeError: + feature_store_cfg["tokenizer_path"] = tokenizer_path store = build_feature_store_from_config(feature_store_cfg, read_only=True) logger.info( "[standalone rank=%s] feature store opened elapsed=%.3fs", @@ -261,6 +276,14 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: optimizer_step, ) ok = await trainer.training_step_from_batch(batch, optimizer_step) + step_error = getattr(trainer, "last_standalone_training_error", None) + if step_error is not None and _is_out_of_memory_error(step_error): + raise RuntimeError( + "Standalone drafter training hit an unrecoverable OOM during " + f"batch={attempted_batches} optimizer_step={optimizer_step}. " + "Reduce batch_size_per_gpu, feature_store.max_seq_len, " + "dspark_num_anchors/block_size or disable DSpark L1 loss." + ) from step_error current_stage = "synchronize_training_step" if not _all_ranks_true(ok, trainer.runtime_device): continue diff --git a/verl_speco/trainer/feature_store.py b/verl_speco/trainer/feature_store.py index 0b5838e2..51b95cdc 100644 --- a/verl_speco/trainer/feature_store.py +++ b/verl_speco/trainer/feature_store.py @@ -599,7 +599,12 @@ def read(self, key: str) -> DraftReplaySample: class JsonlTokenReplayFeatureStore: - """Read ``input_ids``/``loss_mask`` JSONL rows as token replay samples.""" + """Read token replay JSONL rows as compact replay samples. + + Supported row formats: + - ``{"input_ids": [...], "loss_mask": [...]}`` + - ``{"conversations": [{"from": "human", "value": "..."}, ...]}`` + """ def __init__( self, @@ -612,6 +617,9 @@ def __init__( shard_prefix: str = "shard", max_seq_len: int = 512, window_mode: str = "loss", + tokenizer_path: str | os.PathLike[str] | None = None, + trust_remote_code: bool = False, + train_on: str = "last_assistant", ): if path is None: raise ValueError("JsonlTokenReplayFeatureStore requires a non-empty path") @@ -632,11 +640,20 @@ def __init__( self.shard_prefix = str(shard_prefix or "shard") self.max_seq_len = int(max_seq_len or 0) self.window_mode = str(window_mode or "loss").strip().lower() + self.tokenizer_path = os.fspath(tokenizer_path) if tokenizer_path else None + self.trust_remote_code = bool(trust_remote_code) + self.train_on = str(train_on or "last_assistant").strip().lower() + self._tokenizer: Any | None = None if self.window_mode not in {"loss", "front", "full"}: raise ValueError( "feature_store.window_mode for jsonl_token_replay must be " "'loss', 'front' or 'full'" ) + if self.train_on != "last_assistant": + raise ValueError( + "feature_store.train_on for jsonl_token_replay currently supports " + "only 'last_assistant'" + ) self._files = self._resolve_jsonl_files() self._line_offsets: dict[str, list[int]] = {} self._keys = self._build_keys() @@ -672,6 +689,8 @@ def get_metadata(self) -> dict[str, Any]: "num_samples": len(self._keys), "max_seq_len": self.max_seq_len, "window_mode": self.window_mode, + "train_on": self.train_on, + "tokenizer_path": self.tokenizer_path, } ) return metadata @@ -714,8 +733,20 @@ def _build_keys(self) -> list[str]: def _row_to_replay_sample( self, payload: dict[str, Any], file_name: str, line_index: int ) -> DraftReplaySample: - input_ids = _json_list_tensor(payload, "input_ids", dtype=torch.long) - loss_mask = _json_list_tensor(payload, "loss_mask", dtype=torch.float32) + row_source = "jsonl_token_replay" + if "input_ids" in payload or "loss_mask" in payload: + input_ids = _json_list_tensor(payload, "input_ids", dtype=torch.long) + loss_mask = _json_list_tensor(payload, "loss_mask", dtype=torch.float32) + elif "conversations" in payload: + input_ids, loss_mask = self._conversation_to_input_ids_and_loss_mask( + payload + ) + row_source = "jsonl_conversations" + else: + raise KeyError( + "jsonl_token_replay sample must contain input_ids/loss_mask or " + "conversations" + ) if int(input_ids.numel()) != int(loss_mask.numel()): raise ValueError( "jsonl_token_replay input_ids/loss_mask length mismatch: " @@ -737,7 +768,7 @@ def _row_to_replay_sample( feature_positions = torch.arange(start, end, dtype=torch.long) draft_position_ids = position_ids[start:end].long() + 1 metadata = { - "source": "jsonl_token_replay", + "source": row_source, "jsonl_path": file_name, "jsonl_line": line_index, "sequence_length": end - start, @@ -760,6 +791,66 @@ def _row_to_replay_sample( metadata=metadata, ) + def _ensure_tokenizer(self) -> Any: + if self._tokenizer is not None: + return self._tokenizer + if not self.tokenizer_path: + raise ValueError( + "feature_store.tokenizer_path is required when " + "jsonl_token_replay reads conversations rows" + ) + try: + from transformers import AutoTokenizer + except ImportError as exc: + raise RuntimeError( + "jsonl_token_replay conversations rows require transformers" + ) from exc + self._tokenizer = AutoTokenizer.from_pretrained( + self.tokenizer_path, + trust_remote_code=self.trust_remote_code, + ) + return self._tokenizer + + def _conversation_to_input_ids_and_loss_mask( + self, payload: dict[str, Any] + ) -> tuple[torch.Tensor, torch.Tensor]: + conversations = payload.get("conversations") + if not isinstance(conversations, list) or not conversations: + raise ValueError( + "jsonl_token_replay conversations must be a non-empty list" + ) + messages = [_conversation_item_to_message(item) for item in conversations] + if not messages or messages[-1]["role"] != "assistant": + raise ValueError( + "jsonl_token_replay train_on=last_assistant requires the final " + "conversation item to be assistant" + ) + + tokenizer = self._ensure_tokenizer() + prompt_ids = tokenizer.apply_chat_template( + messages[:-1], + tokenize=True, + add_generation_prompt=True, + ) + full_ids = tokenizer.apply_chat_template( + messages, + tokenize=True, + add_generation_prompt=False, + ) + if not isinstance(prompt_ids, list) or not isinstance(full_ids, list): + raise TypeError("tokenizer.apply_chat_template must return token id lists") + if len(full_ids) <= len(prompt_ids): + raise ValueError( + "jsonl_token_replay conversations produced no assistant tokens" + ) + loss_mask = [0.0] * len(prompt_ids) + [1.0] * ( + len(full_ids) - len(prompt_ids) + ) + return ( + torch.tensor(full_ids, dtype=torch.long).reshape(-1), + torch.tensor(loss_mask, dtype=torch.float32).reshape(-1), + ) + def _feature_window(self, loss_mask: torch.Tensor) -> tuple[int, int]: sequence_length = int(loss_mask.numel()) if sequence_length <= 0: @@ -981,6 +1072,12 @@ def build_feature_store_from_config( shard_prefix=shard_prefix, max_seq_len=int(feature_store_cfg.get("max_seq_len", 512) or 0), window_mode=str(feature_store_cfg.get("window_mode", "loss") or "loss"), + tokenizer_path=feature_store_cfg.get("tokenizer_path"), + trust_remote_code=bool(feature_store_cfg.get("trust_remote_code", False)), + train_on=str( + feature_store_cfg.get("train_on", "last_assistant") + or "last_assistant" + ), ) else: raise NotImplementedError(f"Unsupported draft feature store type: {store_type}") @@ -1114,6 +1211,31 @@ def _optional_json_list_tensor( return _json_list_tensor(payload, key, dtype=dtype) +def _normalize_conversation_role(value: Any) -> str: + role = str(value or "").strip().lower() + if role in {"human", "user"}: + return "user" + if role in {"assistant", "gpt"}: + return "assistant" + if role == "system": + return "system" + return role + + +def _conversation_item_to_message(item: Any) -> dict[str, str]: + if not isinstance(item, dict): + raise TypeError( + "jsonl_token_replay conversations entries must be JSON objects" + ) + role = _normalize_conversation_role(item.get("role", item.get("from"))) + content = item.get("content", item.get("value", "")) + if role not in {"system", "user", "assistant"}: + raise ValueError( + f"Unsupported conversation role for jsonl_token_replay: {role!r}" + ) + return {"role": role, "content": str(content)} + + def _json_safe_metadata(value: Any) -> Any: if torch.is_tensor(value): if value.numel() <= 128: diff --git a/verl_speco/trainer/target_feature_replay.py b/verl_speco/trainer/target_feature_replay.py index 37af5537..0667fff1 100644 --- a/verl_speco/trainer/target_feature_replay.py +++ b/verl_speco/trainer/target_feature_replay.py @@ -466,6 +466,9 @@ def __init__( self.target_forward_seconds = 0.0 self.vllm_request_seconds = 0.0 self.vllm_requests = 0 + self._warned_replay_algorithm_mismatch = False + self._warned_replay_layer_mismatch = False + self._warned_replay_layout_mismatch = False logger.info( "[target replay rank=%s] initialized backend=%s algorithm=%s " "target_layers=%s hidden_layout=%s use_logits=%s endpoint=%s cache=%s", @@ -522,10 +525,16 @@ def materialize( def _validate_target_path(self, sample: DraftReplaySample) -> None: if sample.algorithm.upper() != self.algorithm: - raise ValueError( - "Token replay algorithm mismatch: " - f"sample={sample.algorithm!r} training={self.algorithm!r}" - ) + if not self._warned_replay_algorithm_mismatch: + logger.warning( + "[target replay rank=%s] token replay algorithm differs from " + "training algorithm; using the training algorithm for " + "materialized features sample=%s training=%s", + self.rank, + sample.algorithm, + self.algorithm, + ) + self._warned_replay_algorithm_mismatch = True collected_layer_ids = sample.metadata.get("target_layer_ids") if collected_layer_ids is not None: normalized_layer_ids = ( @@ -534,17 +543,28 @@ def _validate_target_path(self, sample: DraftReplaySample) -> None: else [int(value) for value in collected_layer_ids] ) if normalized_layer_ids != self.target_layer_ids: - raise ValueError( - "Token replay target layer mismatch: " - f"collected={normalized_layer_ids} " - f"replay={self.target_layer_ids}" - ) + if not self._warned_replay_layer_mismatch: + logger.warning( + "[target replay rank=%s] token replay target layers differ " + "from replay target layers; recomputing hidden states with " + "the training configuration collected=%s replay=%s", + self.rank, + normalized_layer_ids, + self.target_layer_ids, + ) + self._warned_replay_layer_mismatch = True collected_layout = sample.metadata.get("hidden_states_layout") if collected_layout and str(collected_layout) != self.hidden_layout: - raise ValueError( - "Token replay hidden layout mismatch: " - f"collected={collected_layout!r} replay={self.hidden_layout!r}" - ) + if not self._warned_replay_layout_mismatch: + logger.warning( + "[target replay rank=%s] token replay hidden layout differs " + "from replay layout; recomputing hidden states with the " + "training configuration collected=%s replay=%s", + self.rank, + collected_layout, + self.hidden_layout, + ) + self._warned_replay_layout_mismatch = True if not self.strict_target_model_path: return collected_path = sample.metadata.get("target_model_path") @@ -956,6 +976,7 @@ def _feature_from_vllm_payload( f"got {tuple(hidden.shape)}" ) feature_positions = sample.feature_positions.detach().cpu().long() + hidden_position_offset = max(len(prompt_ids) - int(hidden.size(0)), 0) expected_layers = len(self.target_layer_ids) include_final = self.hidden_layout in { "eagle3_aux_plus_last", @@ -969,7 +990,36 @@ def _feature_from_vllm_payload( "Start vLLM with target layer ids plus the final layer when the " "training layout needs last hidden states." ) - selected = hidden.index_select(0, feature_positions).to(dtype=self.dtype) + relative_feature_positions = feature_positions - hidden_position_offset + feature_keep_mask = ( + (relative_feature_positions >= 0) + & (relative_feature_positions < int(hidden.size(0))) + ) + filtered_feature_positions = not bool(feature_keep_mask.all().item()) + if filtered_feature_positions: + dropped = int((~feature_keep_mask).sum().item()) + logger.warning( + "[target replay rank=%s] dropping vLLM feature positions outside " + "hidden rows dropped=%s hidden_rows=%s hidden_offset=%s " + "feature_min=%s feature_max=%s", + self.rank, + dropped, + int(hidden.size(0)), + hidden_position_offset, + int(feature_positions.min().item()), + int(feature_positions.max().item()), + ) + feature_positions = feature_positions[feature_keep_mask] + relative_feature_positions = relative_feature_positions[feature_keep_mask] + if int(feature_positions.numel()) <= 0: + raise ValueError( + "vLLM hidden_states contain no rows for replay feature positions: " + f"hidden_rows={int(hidden.size(0))}, " + f"hidden_position_offset={hidden_position_offset}" + ) + selected = hidden.index_select(0, relative_feature_positions).to( + dtype=self.dtype + ) aux_hidden = selected[:, :expected_layers, :].flatten(1) if include_final: final_hidden = selected[:, required_layers - 1, :] @@ -978,6 +1028,9 @@ def _feature_from_vllm_payload( hidden_states = aux_hidden selected_input_ids = sample.input_ids.index_select(0, feature_positions).long() selected_loss_mask = sample.loss_mask.index_select(0, feature_positions).float() + draft_position_ids = sample.draft_position_ids.detach().cpu().long() + if filtered_feature_positions: + draft_position_ids = draft_position_ids[feature_keep_mask] metadata = dict(sample.metadata) feature_start = int(feature_positions[0].item()) feature_end = int(feature_positions[-1].item()) + 1 @@ -989,6 +1042,8 @@ def _feature_from_vllm_payload( "target_config_fingerprint": self.target_config_fingerprint, "target_layer_ids": list(self.target_layer_ids), "vllm_hidden_layers": int(hidden.size(1)), + "vllm_hidden_rows": int(hidden.size(0)), + "vllm_hidden_position_offset": hidden_position_offset, "hidden_states_layout": self.hidden_layout, "feature_start": feature_start, "feature_end": feature_end, @@ -1005,7 +1060,7 @@ def _feature_from_vllm_payload( input_ids=selected_input_ids, loss_mask=selected_loss_mask, hidden_states=hidden_states.cpu().contiguous(), - position_ids=sample.draft_position_ids.long(), + position_ids=draft_position_ids, metadata=metadata, ) From c929576dfd4226ed2d81a2bf29104298c88fc0ee Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Tue, 4 Aug 2026 19:26:11 +0800 Subject: [PATCH 29/50] Improve standalone draft token replay training --- tests/unit/test_draft_feature_store.py | 70 +++++++++++++++++++ tests/unit/test_draft_training_loop.py | 16 +++++ verl_speco/backends/dflash_trainer_backend.py | 11 ++- verl_speco/trainer/base_trainer.py | 10 +++ verl_speco/trainer/draft_training_loop.py | 35 +++++++++- verl_speco/trainer/feature_store.py | 29 +++++++- 6 files changed, 165 insertions(+), 6 deletions(-) diff --git a/tests/unit/test_draft_feature_store.py b/tests/unit/test_draft_feature_store.py index 458798c5..c3f95a22 100644 --- a/tests/unit/test_draft_feature_store.py +++ b/tests/unit/test_draft_feature_store.py @@ -200,6 +200,76 @@ def apply_chat_template( assert loaded.metadata["id"] == "conv-0" +def test_jsonl_token_replay_accepts_tensor_chat_template_output(tmp_path): + class TensorTokenizer: + def apply_chat_template( + self, messages, *, tokenize, add_generation_prompt + ): + token_ids = [] + for message in messages: + token_ids.extend([1 if message["role"] == "user" else 2, 3]) + if add_generation_prompt: + token_ids.append(2) + return torch.tensor([token_ids], dtype=torch.long) + + path = tmp_path / "samples.jsonl" + row = { + "conversations": [ + {"from": "human", "value": "question"}, + {"from": "assistant", "value": "answer"}, + ], + } + path.write_text(json.dumps(row) + "\n", encoding="utf-8") + + store = JsonlTokenReplayFeatureStore( + path, + read_only=True, + tokenizer_path="/target", + ) + store._tokenizer = TensorTokenizer() + loaded = store.read(next(store.iter_keys(shuffle=False))) + + assert torch.equal(loaded.input_ids, torch.tensor([1, 3, 2, 3])) + assert torch.equal(loaded.loss_mask, torch.tensor([0.0, 0.0, 0.0, 1.0])) + + +def test_jsonl_token_replay_accepts_batch_encoding_chat_template_output(tmp_path): + class FakeBatchEncoding: + def __init__(self, input_ids): + self.data = {"input_ids": input_ids} + + class BatchEncodingTokenizer: + def apply_chat_template( + self, messages, *, tokenize, add_generation_prompt + ): + token_ids = [] + for message in messages: + token_ids.extend([1 if message["role"] == "user" else 2, 3]) + if add_generation_prompt: + token_ids.append(2) + return FakeBatchEncoding([token_ids]) + + path = tmp_path / "samples.jsonl" + row = { + "conversations": [ + {"from": "human", "value": "question"}, + {"from": "assistant", "value": "answer"}, + ], + } + path.write_text(json.dumps(row) + "\n", encoding="utf-8") + + store = JsonlTokenReplayFeatureStore( + path, + read_only=True, + tokenizer_path="/target", + ) + store._tokenizer = BatchEncodingTokenizer() + loaded = store.read(next(store.iter_keys(shuffle=False))) + + assert torch.equal(loaded.input_ids, torch.tensor([1, 3, 2, 3])) + assert torch.equal(loaded.loss_mask, torch.tensor([0.0, 0.0, 0.0, 1.0])) + + def test_vllm_safetensors_feature_store_records_manifest_path_and_roundtrips(tmp_path): pytest.importorskip("safetensors.torch") hidden_positions = torch.arange(160, dtype=torch.long) diff --git a/tests/unit/test_draft_training_loop.py b/tests/unit/test_draft_training_loop.py index 2faa009a..d8ebb3ee 100644 --- a/tests/unit/test_draft_training_loop.py +++ b/tests/unit/test_draft_training_loop.py @@ -23,12 +23,14 @@ from verl_speco.trainer.standalone_checkpoint import rewrite_standalone_runtime_config # noqa: E402 from verl_speco.trainer.draft_training_loop import ( # noqa: E402 + _contains_replay_samples, _is_out_of_memory_error, _rewrite_standalone_block_runtime_config, _save_standalone_checkpoint, _should_log_batch_progress, _torch_load_cpu, ) +from verl_speco.trainer.feature_store import DraftReplaySample # noqa: E402 class _FakeTrainer: @@ -68,6 +70,20 @@ def test_is_out_of_memory_error_matches_npu_oom_message(): assert not _is_out_of_memory_error(RuntimeError("bad batch")) +def test_contains_replay_samples_detects_draft_replay_sample(): + sample = DraftReplaySample( + input_ids=torch.arange(4), + loss_mask=torch.ones(4), + attention_mask=torch.ones(4, dtype=torch.bool), + position_ids=torch.arange(4), + feature_positions=torch.arange(1, 3), + draft_position_ids=torch.arange(2, 4), + ) + + assert _contains_replay_samples([sample]) + assert not _contains_replay_samples([{"input_ids": [1, 2]}]) + + def test_standalone_checkpoint_schedules_without_waiting(): trainer = _FakeTrainer() diff --git a/verl_speco/backends/dflash_trainer_backend.py b/verl_speco/backends/dflash_trainer_backend.py index 058fabb1..aa0f0026 100644 --- a/verl_speco/backends/dflash_trainer_backend.py +++ b/verl_speco/backends/dflash_trainer_backend.py @@ -570,8 +570,15 @@ def forward( loss_sum_per_position = ( loss_per_token.view(bsz, n_blocks, self.block_size) * binary_weights ).sum(dim=(0, 1)) + correct_3d = correct.view(bsz, n_blocks, self.block_size) + pred_valid_3d = binary_weights[:, :, 1:].bool() + pred_correct_3d = correct_3d[:, :, 1:] & pred_valid_3d + simulated_accept_length_sum = ( + pred_correct_3d.float().cumprod(dim=-1).sum() + ) + simulated_accept_block_count = pred_valid_3d.any(dim=-1).float().sum() correct_per_position = ( - correct.view(bsz, n_blocks, self.block_size).float().sum(dim=(0, 1)) + correct_3d.float().sum(dim=(0, 1)) ) loss_per_position = loss_sum_per_position / count_per_pos acc_per_position = correct_per_position / count_per_pos @@ -584,6 +591,8 @@ def forward( "quality_token_count": quality_token_count, "valid_token_count": binary_eval_mask.sum().float(), "weighted_token_count": flat_weights.sum().float(), + "simulated_accept_length_sum": simulated_accept_length_sum, + "simulated_accept_block_count": simulated_accept_block_count, "sanitized_rows": sanitized_rows, "masked_rows": masked_rows, "loss_sum_per_position": loss_sum_per_position, diff --git a/verl_speco/trainer/base_trainer.py b/verl_speco/trainer/base_trainer.py index bb064f44..9a99595d 100644 --- a/verl_speco/trainer/base_trainer.py +++ b/verl_speco/trainer/base_trainer.py @@ -690,6 +690,14 @@ def get_training_metrics(self) -> dict[str, float]: metrics[f"{prefix}/top5_acc"] = ( sums.get(f"{prefix}/top5_correct_count", 0.0) / quality_tokens ) + simulated_accept_blocks = sums.get( + f"{prefix}/simulated_accept_block_count", 0.0 + ) + if simulated_accept_blocks > 0: + metrics[f"{prefix}/simulated_acc_len"] = ( + sums.get(f"{prefix}/simulated_accept_length_sum", 0.0) + / simulated_accept_blocks + ) ce_tokens = sums.get(f"{prefix}/ce_weighted_token_count", 0.0) if ce_tokens > 0: metrics[f"{prefix}/ce_loss"] = ( @@ -781,6 +789,8 @@ def _record_dflash_training_metrics(self, loss_dict: dict[str, Any]) -> None: "quality_token_count": f"{prefix}/quality_token_count", "valid_token_count": f"{prefix}/valid_token_count", "weighted_token_count": f"{prefix}/weighted_token_count", + "simulated_accept_length_sum": f"{prefix}/simulated_accept_length_sum", + "simulated_accept_block_count": f"{prefix}/simulated_accept_block_count", "ce_loss_sum": f"{prefix}/ce_loss_sum", "ce_weighted_token_count": f"{prefix}/ce_weighted_token_count", "l1_loss_sum": f"{prefix}/l1_loss_sum", diff --git a/verl_speco/trainer/draft_training_loop.py b/verl_speco/trainer/draft_training_loop.py index a54bbc71..84c07625 100644 --- a/verl_speco/trainer/draft_training_loop.py +++ b/verl_speco/trainer/draft_training_loop.py @@ -35,7 +35,10 @@ DraftFeatureDataLoader, DraftFeatureDataLoaderConfig, ) -from verl_speco.trainer.feature_store import build_feature_store_from_config +from verl_speco.trainer.feature_store import ( + DraftReplaySample, + build_feature_store_from_config, +) from verl_speco.trainer.standalone_checkpoint import rewrite_standalone_runtime_config logger = logging.getLogger(__name__) @@ -52,6 +55,10 @@ def _is_out_of_memory_error(error: BaseException) -> bool: return error.__class__.__name__ in {"OutOfMemoryError", "CudaOutOfMemoryError"} +def _contains_replay_samples(samples: list[Any]) -> bool: + return any(isinstance(sample, DraftReplaySample) for sample in samples) + + def run_standalone_draft_training(config) -> dict[str, Any]: """Run independent draft training from a feature store.""" return asyncio.run(_run_standalone_draft_training_async(config)) @@ -154,7 +161,7 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: if feature_store_type in {"jsonl_token_replay", "jsonl"} and not ( feature_store_cfg.get("tokenizer_path") ): - tokenizer_path = draft_config.actor_rollout_ref.model.path + tokenizer_path = draft_config.model.path try: feature_store_cfg.tokenizer_path = tokenizer_path except AttributeError: @@ -228,6 +235,24 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: successful_steps, ) current_stage = "materialize_target_features" + if feature_replayer is None and _contains_replay_samples(samples): + logger.warning( + "[standalone rank=%s] feature store type=%s yielded replay " + "samples without an initialized target feature replayer; " + "initializing replayer lazily", + rank, + feature_store_type, + ) + from verl_speco.trainer.target_feature_replay import ( + TargetFeatureReplayer, + ) + + feature_replayer = TargetFeatureReplayer( + config, + rank=rank, + world_size=world_size, + device=trainer.runtime_device, + ) if feature_replayer is not None: materialize_started = time.perf_counter() if log_batch_progress: @@ -880,7 +905,11 @@ def _standalone_step_metrics( metrics["train/avg_loss"] = avg_loss if avg_acc is not None: metrics["train/avg_acc"] = avg_acc - if pred_accuracies: + if f"{prefix}/simulated_acc_len" in raw_metrics: + metrics["train/simulated_acc_len"] = float( + raw_metrics[f"{prefix}/simulated_acc_len"] + ) + elif pred_accuracies: metrics["train/simulated_acc_len"] = _simulated_accept_length( pred_accuracies ) diff --git a/verl_speco/trainer/feature_store.py b/verl_speco/trainer/feature_store.py index 51b95cdc..8ebdd156 100644 --- a/verl_speco/trainer/feature_store.py +++ b/verl_speco/trainer/feature_store.py @@ -837,8 +837,8 @@ def _conversation_to_input_ids_and_loss_mask( tokenize=True, add_generation_prompt=False, ) - if not isinstance(prompt_ids, list) or not isinstance(full_ids, list): - raise TypeError("tokenizer.apply_chat_template must return token id lists") + prompt_ids = _token_ids_to_list(prompt_ids) + full_ids = _token_ids_to_list(full_ids) if len(full_ids) <= len(prompt_ids): raise ValueError( "jsonl_token_replay conversations produced no assistant tokens" @@ -1236,6 +1236,31 @@ def _conversation_item_to_message(item: Any) -> dict[str, str]: return {"role": role, "content": str(content)} +def _token_ids_to_list(value: Any) -> list[int]: + if hasattr(value, "data") and isinstance(getattr(value, "data", None), dict): + data = getattr(value, "data") + if "input_ids" in data: + value = data["input_ids"] + elif isinstance(value, dict) and "input_ids" in value: + value = value["input_ids"] + if torch.is_tensor(value): + value = value.detach().cpu().tolist() + elif hasattr(value, "tolist"): + value = value.tolist() + if ( + isinstance(value, list) + and len(value) == 1 + and isinstance(value[0], (list, tuple)) + ): + value = value[0] + if not isinstance(value, (list, tuple)): + raise TypeError( + "tokenizer.apply_chat_template must return token ids as a list, " + f"tuple, tensor, or tolist()-compatible value; got {type(value).__name__}" + ) + return [int(token_id) for token_id in value] + + def _json_safe_metadata(value: Any) -> Any: if torch.is_tensor(value): if value.numel() <= 128: From 7dc84502d0556a311edf57ab7805e10d2ba66643 Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Thu, 13 Aug 2026 20:06:37 +0800 Subject: [PATCH 30/50] feat(draft-training): add producer pipeline and Mooncake transfer --- README.md | 86 +++ tests/config/test_speco_config_overlay.py | 10 + tests/unit/test_draft_training_loop.py | 55 +- tests/unit/test_mooncake_transfer.py | 73 +++ tests/unit/test_target_feature_pipeline.py | 80 +++ tests/unit/test_target_feature_replay.py | 67 +++ tools/run_qwen3-4b_drafter_dspark_mooncake.sh | 72 +++ verl_speco/config/draft_trainer.yaml | 26 + .../mooncake_hidden_states_connector.py | 329 ++++++++++++ verl_speco/trainer/draft_training_loop.py | 133 ++++- verl_speco/trainer/mooncake_transfer.py | 216 ++++++++ verl_speco/trainer/target_feature_pipeline.py | 196 +++++++ verl_speco/trainer/target_feature_replay.py | 498 ++++++++++++++++-- 13 files changed, 1778 insertions(+), 63 deletions(-) create mode 100644 tests/unit/test_mooncake_transfer.py create mode 100644 tests/unit/test_target_feature_pipeline.py create mode 100644 tools/run_qwen3-4b_drafter_dspark_mooncake.sh create mode 100644 verl_speco/integration/mooncake_hidden_states_connector.py create mode 100644 verl_speco/trainer/mooncake_transfer.py create mode 100644 verl_speco/trainer/target_feature_pipeline.py diff --git a/README.md b/README.md index 15322db9..8f75310c 100644 --- a/README.md +++ b/README.md @@ -313,6 +313,92 @@ python -m verl_speco.inspect_feature_store /path/to/features \ --strict-exit ``` +### Producer/Mooncake pipeline + +For token-only stores, standalone training can overlap target-model inference, +hidden-state transfer, and drafter optimization. vLLM acts as the producer, +Mooncake transfers the extracted tensors without per-sample files, and a +bounded rank-local queue prefetches complete batches while FSDP trains the +current batch. This path is disabled by default and does not change online +PPO/drafter training. + +Install the Mooncake package for the target hardware and start a Mooncake +master. For example, use `mooncake-transfer-engine` on CUDA or +`mooncake-transfer-engine-npu` on Ascend. Export the same connection settings +in the vLLM and training processes. The connector requires vLLM 0.23 or newer: + +```bash +# CUDA; use mooncake-transfer-engine-npu instead on Ascend. +pip install 'mooncake-transfer-engine>=0.3.10.post1' + +export MOONCAKE_MASTER_SERVER=127.0.0.1:50051 +export MOONCAKE_METADATA_SERVER=http://127.0.0.1:8090/metadata +export MOONCAKE_PROTOCOL=tcp +export MOONCAKE_GLOBAL_SEGMENT_SIZE=$((16 * 1024 * 1024 * 1024)) +export MOONCAKE_LOCAL_BUFFER_SIZE=$((2 * 1024 * 1024 * 1024)) + +mooncake_master \ + --enable_http_metadata_server=true \ + --http_metadata_server_host=0.0.0.0 \ + --http_metadata_server_port=8090 +``` + +Start vLLM with SpeCo's store-only hidden-state connector. The final layer must +be appended to the auxiliary layer list when the selected drafter loss needs +last-hidden-state supervision: + +```bash +vllm serve /path/to/target_model \ + --port 8000 \ + --tensor-parallel-size 2 \ + --speculative-config '{"method":"extract_hidden_states","num_speculative_tokens":1,"draft_model_config":{"hf_config":{"eagle_aux_hidden_state_layer_ids":[1,9,17,25,33,36]}}}' \ + --kv-transfer-config '{"kv_connector":"SpeCoMooncakeHiddenStatesConnector","kv_connector_module_path":"verl_speco.integration.mooncake_hidden_states_connector","kv_role":"kv_producer"}' \ + --no-enable-chunked-prefill +``` + +Enable the pipeline in the standalone command: + +```bash +python -m verl_speco.draft_train_launcher \ + speco.draft_training.nproc_per_node=4 \ + actor_rollout_ref.model.path=/path/to/target_model \ + actor_rollout_ref.rollout.drafter.speculative_algorithm=DSPARK \ + actor_rollout_ref.rollout.drafter.training.mode=offline \ + actor_rollout_ref.rollout.drafter.training.feature_store.type=jsonl_token_replay \ + actor_rollout_ref.rollout.drafter.training.feature_store.path=/path/to/data.jsonl \ + actor_rollout_ref.rollout.drafter.training.target_feature_replay.backend=vllm_mooncake \ + actor_rollout_ref.rollout.drafter.training.target_feature_replay.vllm_endpoint=http://127.0.0.1:8000/v1 \ + actor_rollout_ref.rollout.drafter.training.target_feature_pipeline.enabled=true \ + actor_rollout_ref.rollout.drafter.training.target_feature_pipeline.concurrency=16 \ + actor_rollout_ref.rollout.drafter.training.target_feature_pipeline.transfer_concurrency=8 \ + actor_rollout_ref.rollout.drafter.training.target_feature_pipeline.producer_prefetch_depth=4 \ + actor_rollout_ref.rollout.drafter.training.target_feature_pipeline.prefetch_depth=2 +``` + +To spread replay requests across independent vLLM deployments, replace the +single endpoint override with an endpoint pool: + +```bash +'actor_rollout_ref.rollout.drafter.training.target_feature_replay.vllm_endpoints=["http://host1:8000/v1","http://host2:8000/v1"]' \ +actor_rollout_ref.rollout.drafter.training.target_feature_replay.endpoint_cooldown=5 +``` + +The standalone producer routes each request to the healthy endpoint with the +fewest in-flight requests. A failed request is retried on another endpoint when +the configured retry budget permits it. The original `vllm_endpoint` option +remains supported and is used when `vllm_endpoints` is unset. + +`concurrency` and `transfer_concurrency` are global budgets divided across all +training ranks. Start with request concurrency between 16 and 32, then increase +only while vLLM throughput rises. +`producer_prefetch_depth` bounds batches that have outstanding HTTP work; +`transfer_concurrency` bounds simultaneous Mooncake GETs. A +`prefetch_depth` of 2 normally hides transfer latency without retaining too +many large hidden-state batches. The standalone metrics include producer queue +depth, consumer wait time, vLLM request time, and Mooncake GET time. +The complete three-process launcher is +[`tools/run_qwen3-4b_drafter_dspark_mooncake.sh`](./tools/run_qwen3-4b_drafter_dspark_mooncake.sh). + ## Configuration SPECO-specific options live under: diff --git a/tests/config/test_speco_config_overlay.py b/tests/config/test_speco_config_overlay.py index 9039fa67..3c44aa91 100644 --- a/tests/config/test_speco_config_overlay.py +++ b/tests/config/test_speco_config_overlay.py @@ -83,6 +83,16 @@ def test_overlay_has_expected_default_drafter_shape() -> None: assert "target_feature_replay" not in drafter.training assert standalone_training.target_feature_replay.cache.enabled is False assert standalone_training.target_feature_replay.cache.max_size_gb == 0 + assert standalone_training.target_feature_replay.vllm_endpoints is None + assert standalone_training.target_feature_replay.endpoint_cooldown == 5 + assert standalone_training.target_feature_pipeline.enabled is False + assert standalone_training.target_feature_pipeline.concurrency == 16 + assert standalone_training.target_feature_pipeline.transfer_concurrency == 8 + assert standalone_training.target_feature_pipeline.producer_prefetch_depth == 4 + assert standalone_training.target_feature_pipeline.prefetch_depth == 2 + assert ( + standalone_training.target_feature_replay.mooncake.protocol == "tcp" + ) def test_overlay_composes_with_release_upstream_verl(tmp_path: Path) -> None: diff --git a/tests/unit/test_draft_training_loop.py b/tests/unit/test_draft_training_loop.py index eb3622c3..d4abdfd3 100644 --- a/tests/unit/test_draft_training_loop.py +++ b/tests/unit/test_draft_training_loop.py @@ -23,17 +23,18 @@ from omegaconf import OmegaConf # noqa: E402 -from verl_speco.trainer.standalone_checkpoint import rewrite_standalone_runtime_config # noqa: E402 from verl_speco.trainer.draft_training_loop import ( # noqa: E402 _build_backend, _contains_replay_samples, _is_out_of_memory_error, + _next_batch_across_ranks, _rewrite_standalone_block_runtime_config, _save_standalone_checkpoint, _should_log_batch_progress, _torch_load_cpu, ) from verl_speco.trainer.feature_store import DraftReplaySample # noqa: E402 +from verl_speco.trainer.standalone_checkpoint import rewrite_standalone_runtime_config # noqa: E402 class _FakeTrainer: @@ -786,3 +787,55 @@ def fake_load(path, **kwargs): {"map_location": "cpu", "weights_only": True}, {"map_location": "cpu"}, ] + + +def test_next_batch_across_ranks_returns_local_batch_without_distributed(): + batch = [object()] + + assert _next_batch_across_ranks( + iter([batch]), rank=0, device=torch.device("cpu") + ) is batch + + +def test_next_batch_across_ranks_returns_none_when_source_is_exhausted(): + assert ( + _next_batch_across_ranks( + iter(()), rank=0, device=torch.device("cpu") + ) + is None + ) + + +def test_next_batch_across_ranks_preserves_local_producer_error(): + def broken_source(): + raise ValueError("producer failed") + yield [] + + with pytest.raises(RuntimeError, match="failed on rank=0") as exc_info: + _next_batch_across_ranks( + iter(broken_source()), rank=0, device=torch.device("cpu") + ) + + assert isinstance(exc_info.value.__cause__, ValueError) + + +def test_next_batch_across_ranks_stops_for_remote_rank_failure(monkeypatch): + monkeypatch.setattr( + "verl_speco.trainer.draft_training_loop.dist.is_initialized", lambda: True + ) + monkeypatch.setattr( + "verl_speco.trainer.draft_training_loop.dist.get_world_size", lambda: 2 + ) + + def fake_all_reduce(state, op): + del op + state[0] = 1 + + monkeypatch.setattr( + "verl_speco.trainer.draft_training_loop.dist.all_reduce", fake_all_reduce + ) + + with pytest.raises(RuntimeError, match="failed on another rank"): + _next_batch_across_ranks( + iter([[object()]]), rank=1, device=torch.device("cpu") + ) diff --git a/tests/unit/test_mooncake_transfer.py b/tests/unit/test_mooncake_transfer.py new file mode 100644 index 00000000..2cb342ea --- /dev/null +++ b/tests/unit/test_mooncake_transfer.py @@ -0,0 +1,73 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); + +import sys +import types + +import pytest + +torch = pytest.importorskip("torch") +pytest.importorskip("safetensors") + +from verl_speco.trainer.mooncake_transfer import ( # noqa: E402 + MooncakeTensorStore, + MooncakeTransferConfig, + _parse_size, +) + + +class _RawStore: + objects = {} + + def setup(self, **kwargs): + self.setup_kwargs = kwargs + return 0 + + def put(self, key, payload): + self.objects[key] = payload + return 0 + + def get(self, key): + return self.objects.get(key) + + def remove(self, key, force): + self.objects.pop(key, None) + + def close(self): + return None + + +def test_mooncake_tensor_store_roundtrip(monkeypatch): + store_module = types.ModuleType("mooncake.store") + store_module.MooncakeDistributedStore = _RawStore + mooncake_module = types.ModuleType("mooncake") + mooncake_module.store = store_module + monkeypatch.setitem(sys.modules, "mooncake", mooncake_module) + monkeypatch.setitem(sys.modules, "mooncake.store", store_module) + + config = MooncakeTransferConfig( + local_hostname="localhost", + metadata_server="P2PHANDSHAKE", + master_server_address="127.0.0.1:50051", + global_segment_size=_parse_size("64MB"), + local_buffer_size=_parse_size("128MB"), + protocol="tcp", + device_name="", + get_timeout=1, + get_poll_interval=0.01, + ) + store = MooncakeTensorStore(config) + expected = { + "token_ids": torch.arange(4), + "hidden_states": torch.arange(24, dtype=torch.bfloat16).reshape(2, 3, 4), + } + + metadata = store.put("sample", expected) + actual = store.get("sample") + + assert metadata["tensor_shapes"]["hidden_states"] == (2, 3, 4) + assert torch.equal(actual["token_ids"], expected["token_ids"]) + assert torch.equal(actual["hidden_states"], expected["hidden_states"]) + store.remove("sample") + assert "sample" not in _RawStore.objects diff --git a/tests/unit/test_target_feature_pipeline.py b/tests/unit/test_target_feature_pipeline.py new file mode 100644 index 00000000..67e8a227 --- /dev/null +++ b/tests/unit/test_target_feature_pipeline.py @@ -0,0 +1,80 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); + +import time + +import pytest + +torch = pytest.importorskip("torch") + +from verl_speco.trainer.feature_store import DraftFeatureSample # noqa: E402 +from verl_speco.trainer.target_feature_pipeline import ( # noqa: E402 + TargetFeatureProducer, +) + + +def _feature(value: int) -> DraftFeatureSample: + return DraftFeatureSample( + input_ids=torch.tensor([value]), + loss_mask=torch.ones(1), + hidden_states=torch.zeros(1, 4), + position_ids=torch.ones(1, dtype=torch.long), + ) + + +class _Replayer: + backend = "vllm_file" + + def materialize(self, samples): + time.sleep(0.01) + return [_feature(int(samples[0]))] + + +def test_target_feature_producer_preserves_batch_order_and_prefetches(): + producer = TargetFeatureProducer( + [[0, 1], [2, 3]], + _Replayer(), + rank=0, + concurrency=2, + transfer_concurrency=1, + producer_prefetch_depth=2, + prefetch_depth=2, + queue_timeout=2, + ) + try: + batches = list(producer) + finally: + producer.close() + + assert [[int(x.input_ids[0]) for x in batch] for batch in batches] == [ + [0, 1], + [2, 3], + ] + assert producer.metrics()["producer/samples_total"] == 4 + + +class _FailingReplayer: + backend = "vllm_file" + + def materialize(self, samples): + raise ValueError("broken replay") + + +def test_target_feature_producer_propagates_background_failure(): + producer = TargetFeatureProducer( + [[0]], + _FailingReplayer(), + rank=0, + concurrency=1, + transfer_concurrency=1, + producer_prefetch_depth=1, + prefetch_depth=1, + queue_timeout=2, + ) + try: + with pytest.raises(RuntimeError, match="producer failed") as exc_info: + next(producer) + assert isinstance(exc_info.value.__cause__, ValueError) + finally: + producer.close() diff --git a/tests/unit/test_target_feature_replay.py b/tests/unit/test_target_feature_replay.py index 818a4338..a608611a 100644 --- a/tests/unit/test_target_feature_replay.py +++ b/tests/unit/test_target_feature_replay.py @@ -12,6 +12,9 @@ # See the License for the specific language governing permissions and # limitations under the License. +import threading +from types import SimpleNamespace + import pytest torch = pytest.importorskip("torch") @@ -20,7 +23,9 @@ from verl_speco.trainer.target_feature_replay import ( # noqa: E402 BoundedReplayCache, TargetFeatureReplayer, + _VllmEndpointState, _hidden_capture_target, + _normalize_vllm_endpoints, ) @@ -39,6 +44,68 @@ def test_hidden_capture_target_matches_transformers_hidden_state_indices(): assert _hidden_capture_target(35, 36) == ("final", None) +def test_normalize_vllm_endpoints_prefers_pool_and_deduplicates(): + assert _normalize_vllm_endpoints( + { + "vllm_endpoint": "http://legacy:8000/v1", + "vllm_endpoints": [ + "http://host1:8000/v1/", + "http://host2:8000/v1", + "http://host1:8000/v1", + ], + } + ) == ["http://host1:8000/v1", "http://host2:8000/v1"] + + +def test_vllm_request_fails_over_to_another_endpoint(monkeypatch): + class _Completions: + def __init__(self, error=None): + self.error = error + self.calls = 0 + + def create(self, **kwargs): + self.calls += 1 + if self.error is not None: + raise self.error + return SimpleNamespace(choices=[]) + + failed = _Completions(RuntimeError("endpoint down")) + healthy = _Completions() + replayer = TargetFeatureReplayer.__new__(TargetFeatureReplayer) + replayer.rank = 0 + replayer.vllm_timeout = 1 + replayer.vllm_max_retries = 1 + replayer.vllm_endpoint_cooldown = 5 + replayer.vllm_requests = 0 + replayer.vllm_request_seconds = 0.0 + replayer._metrics_lock = threading.Lock() + replayer._endpoint_lock = threading.Lock() + replayer._vllm_clients_initialized = True + replayer._vllm_endpoint_states = [ + _VllmEndpointState( + index=0, + url="http://host1:8000/v1", + client=SimpleNamespace(completions=failed), + model="target", + ), + _VllmEndpointState( + index=1, + url="http://host2:8000/v1", + client=SimpleNamespace(completions=healthy), + model="target", + ), + ] + monkeypatch.setattr("verl_speco.trainer.target_feature_replay.time.sleep", lambda _: None) + + response = replayer._request_vllm_response([1, 2, 3]) + + assert response.choices == [] + assert failed.calls == 1 + assert healthy.calls == 1 + assert replayer._vllm_endpoint_states[0].failures == 1 + assert replayer._vllm_endpoint_states[1].requests == 1 + + def test_bounded_replay_cache_roundtrip(tmp_path): cache = BoundedReplayCache( tmp_path, diff --git a/tools/run_qwen3-4b_drafter_dspark_mooncake.sh b/tools/run_qwen3-4b_drafter_dspark_mooncake.sh new file mode 100644 index 00000000..10398a51 --- /dev/null +++ b/tools/run_qwen3-4b_drafter_dspark_mooncake.sh @@ -0,0 +1,72 @@ +set -euo pipefail +set -x + +# Run each stage in a separate shell: RUN_STAGE=master, vllm, then train. +# On Ascend replace CUDA_VISIBLE_DEVICES below with ASCEND_RT_VISIBLE_DEVICES. +RUN_STAGE=${RUN_STAGE:-train} +MODEL_PATH=${MODEL_PATH:-/path/to/Qwen3-4B} +DATA_PATH=${DATA_PATH:-/path/to/token_replay.jsonl} +DRAFT_CKPTS_DIR=${DRAFT_CKPTS_DIR:-/path/to/draft_checkpoints} +VLLM_DEVICES=${VLLM_DEVICES:-0,1} +TRAIN_DEVICES=${TRAIN_DEVICES:-2,3,4,5} +VLLM_TP=${VLLM_TP:-2} +TRAIN_GPUS=${TRAIN_GPUS:-4} + +export MOONCAKE_MASTER_SERVER=${MOONCAKE_MASTER_SERVER:-127.0.0.1:50051} +export MOONCAKE_METADATA_SERVER=${MOONCAKE_METADATA_SERVER:-http://127.0.0.1:8090/metadata} +export MOONCAKE_PROTOCOL=${MOONCAKE_PROTOCOL:-tcp} +export MOONCAKE_GLOBAL_SEGMENT_SIZE=${MOONCAKE_GLOBAL_SEGMENT_SIZE:-17179869184} +export MOONCAKE_LOCAL_BUFFER_SIZE=${MOONCAKE_LOCAL_BUFFER_SIZE:-2147483648} + +if [ "${RUN_STAGE}" = "master" ]; then + exec mooncake_master \ + --enable_http_metadata_server=true \ + --http_metadata_server_host=0.0.0.0 \ + --http_metadata_server_port=8090 +fi + +if [ "${RUN_STAGE}" = "vllm" ]; then + CUDA_VISIBLE_DEVICES=${VLLM_DEVICES} exec vllm serve "${MODEL_PATH}" \ + --host 0.0.0.0 --port 8000 \ + --tensor-parallel-size "${VLLM_TP}" \ + --gpu-memory-utilization 0.85 \ + --speculative-config '{"method":"extract_hidden_states","num_speculative_tokens":1,"draft_model_config":{"hf_config":{"eagle_aux_hidden_state_layer_ids":[1,9,17,25,33,36]}}}' \ + --kv-transfer-config '{"kv_connector":"SpeCoMooncakeHiddenStatesConnector","kv_connector_module_path":"verl_speco.integration.mooncake_hidden_states_connector","kv_role":"kv_producer"}' \ + --no-enable-chunked-prefill +fi + +if [ "${RUN_STAGE}" != "train" ]; then + echo "RUN_STAGE must be master, vllm, or train" >&2 + exit 2 +fi + +CUDA_VISIBLE_DEVICES=${TRAIN_DEVICES} PYTHONUNBUFFERED=1 \ +python3 -m verl_speco.draft_train_launcher \ + speco.draft_training.nproc_per_node=${TRAIN_GPUS} \ + speco.draft_training.nnodes=1 \ + actor_rollout_ref.model.path=${MODEL_PATH} \ + actor_rollout_ref.actor.strategy=fsdp2 \ + actor_rollout_ref.rollout.drafter.enable=True \ + actor_rollout_ref.rollout.drafter.enable_drafter_training=True \ + actor_rollout_ref.rollout.drafter.checkpoint_path=${DRAFT_CKPTS_DIR} \ + actor_rollout_ref.rollout.drafter.speculative_algorithm=DSPARK \ + actor_rollout_ref.rollout.drafter.training.mode=offline \ + actor_rollout_ref.rollout.drafter.training.feature_store.type=jsonl_token_replay \ + actor_rollout_ref.rollout.drafter.training.feature_store.path=${DATA_PATH} \ + actor_rollout_ref.rollout.drafter.training.feature_store.shuffle=True \ + actor_rollout_ref.rollout.drafter.training.feature_store.repeat=True \ + actor_rollout_ref.rollout.drafter.training.target_feature_replay.backend=vllm_mooncake \ + actor_rollout_ref.rollout.drafter.training.target_feature_replay.vllm_endpoint=http://127.0.0.1:8000/v1 \ + actor_rollout_ref.rollout.drafter.training.target_feature_replay.on_generate=delete \ + actor_rollout_ref.rollout.drafter.training.target_feature_pipeline.enabled=True \ + actor_rollout_ref.rollout.drafter.training.target_feature_pipeline.concurrency=16 \ + actor_rollout_ref.rollout.drafter.training.target_feature_pipeline.transfer_concurrency=8 \ + actor_rollout_ref.rollout.drafter.training.target_feature_pipeline.producer_prefetch_depth=4 \ + actor_rollout_ref.rollout.drafter.training.target_feature_pipeline.prefetch_depth=2 \ + actor_rollout_ref.rollout.drafter.training.batch_size_per_gpu=4 \ + actor_rollout_ref.rollout.drafter.training.max_steps=1000 \ + actor_rollout_ref.rollout.drafter.training.save_interval_steps=100 \ + actor_rollout_ref.rollout.drafter.training.lr=1e-5 \ + actor_rollout_ref.rollout.drafter.training.lr_warmup_steps=50 \ + actor_rollout_ref.rollout.drafter.training.warmup_style=cosine \ + "$@" diff --git a/verl_speco/config/draft_trainer.yaml b/verl_speco/config/draft_trainer.yaml index b5593b15..e17d7862 100644 --- a/verl_speco/config/draft_trainer.yaml +++ b/verl_speco/config/draft_trainer.yaml @@ -55,11 +55,27 @@ actor_rollout_ref: strict_target_model_path: false logits_chunk_rows: 32 vllm_endpoint: http://localhost:8000/v1 + # Optional endpoint pool. When non-null, this takes precedence over + # vllm_endpoint and requests use least-inflight routing with failover. + vllm_endpoints: null vllm_model: null request_timeout: 120 max_retries: 3 + endpoint_cooldown: 5 on_generate: delete require_arange_positions: true + mooncake: + local_hostname: ${oc.env:MOONCAKE_LOCAL_HOSTNAME,localhost} + metadata_server: ${oc.env:MOONCAKE_METADATA_SERVER,http://localhost:8090/metadata} + master_server_address: ${oc.env:MOONCAKE_MASTER_SERVER,localhost:50051} + # Consumer ranks only need receive buffers. The vLLM producer owns + # the large segment through its MOONCAKE_GLOBAL_SEGMENT_SIZE env. + global_segment_size: 64MB + local_buffer_size: 1GB + protocol: tcp + device_name: "" + get_timeout: 120 + get_poll_interval: 0.02 offline_generation: input_type: token_replay input_path: null @@ -71,3 +87,13 @@ actor_rollout_ref: enabled: false path: null max_size_gb: 0 + # Overlap vLLM generation and Mooncake/file transfer with FSDP training. + # This pipeline is standalone-only and accepts vLLM replay backends. + target_feature_pipeline: + enabled: false + # Global budgets. Standalone divides them across torchrun ranks. + concurrency: 16 + transfer_concurrency: 8 + producer_prefetch_depth: 4 + prefetch_depth: 2 + queue_timeout: 300 diff --git a/verl_speco/integration/mooncake_hidden_states_connector.py b/verl_speco/integration/mooncake_hidden_states_connector.py new file mode 100644 index 00000000..4dce80ec --- /dev/null +++ b/verl_speco/integration/mooncake_hidden_states_connector.py @@ -0,0 +1,329 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +"""vLLM connector that publishes extracted target features through Mooncake. + +Load this module with ``kv_connector_module_path``. It is intentionally kept +outside normal SpeCo imports because its API is tied to vLLM V1 internals. +""" + +from __future__ import annotations + +import logging +import os +import re +import socket +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Any + +import torch +from vllm.config import VllmConfig, get_layers_from_vllm_config +from vllm.distributed.kv_transfer.kv_connector.v1.base import ( + KVConnectorBase_V1, + KVConnectorMetadata, + KVConnectorRole, + SupportsHMA, +) +from vllm.v1.attention.backend import AttentionMetadata +from vllm.v1.core.sched.output import SchedulerOutput + +from verl_speco.trainer.mooncake_transfer import ( + MooncakeTensorStore, + MooncakeTransferConfig, +) + +if TYPE_CHECKING: + from vllm.v1.core.kv_cache_manager import KVCacheBlocks + from vllm.v1.kv_cache_interface import KVCacheConfig + from vllm.v1.request import Request + +logger = logging.getLogger(__name__) + + +def _validate_vllm_version() -> None: + try: + from importlib.metadata import version + + from packaging.version import Version + + installed = Version(version("vllm")) + except Exception: # noqa: BLE001 + return + if installed < Version("0.23.0"): + raise RuntimeError( + "SpeCoMooncakeHiddenStatesConnector requires vLLM >= 0.23.0; " + f"found {installed}. The connector uses the V1 HMA hidden-state API." + ) + + +def _safe_key(key: str) -> str: + value = re.sub(r"[^a-zA-Z0-9_-]", "_", key) + return f"k{value}" if value and value[0].isdigit() else value + + +def _slot_mapping( + block_ids: list[int], page_size: int, num_tokens: int, device: torch.device +) -> torch.Tensor: + blocks = torch.tensor(block_ids, dtype=torch.int64, device=device) + offsets = torch.arange(page_size, dtype=torch.int64, device=device) + return (blocks.unsqueeze(1) * page_size + offsets).flatten()[:num_tokens] + + +@dataclass +class _RequestMetadata: + request_id: str + token_ids: torch.Tensor + block_ids: list[int] = field(default_factory=list) + + +@dataclass +class SpeCoMooncakeConnectorMetadata(KVConnectorMetadata): + requests: list[_RequestMetadata] = field(default_factory=list) + + def add(self, request_id: str, token_ids: list[int], block_ids: list[int]) -> None: + self.requests.append( + _RequestMetadata( + request_id=request_id, + token_ids=torch.tensor(token_ids, dtype=torch.long), + block_ids=list(block_ids), + ) + ) + + +class SpeCoMooncakeHiddenStatesConnector(KVConnectorBase_V1, SupportsHMA): + """Store-only connector for vLLM ``extract_hidden_states`` output.""" + + @property + def prefer_cross_layer_blocks(self) -> bool: + return False + + def __init__( + self, + vllm_config: VllmConfig, + role: KVConnectorRole, + kv_cache_config: "KVCacheConfig | None" = None, + ): + _validate_vllm_version() + super().__init__( + vllm_config=vllm_config, + role=role, + kv_cache_config=kv_cache_config, + ) + speculative = vllm_config.speculative_config + if speculative is None: + raise ValueError( + "SpeCoMooncakeHiddenStatesConnector requires extract_hidden_states" + ) + hf_config = speculative.draft_model_config.hf_config + self._layer_ids = list( + getattr(hf_config, "eagle_aux_hidden_state_layer_ids", []) + ) + self._hidden_size = int(vllm_config.model_config.get_hidden_size()) + self._training_layers = max(len(self._layer_ids) - 1, 1) + self._cache_layers: list[str] = [] + self._cache_group_id = self._find_cache_group(kv_cache_config) + self._active_requests: dict[str, Any] = {} + self._request_blocks: dict[str, list[int]] = {} + self._response_metadata: dict[str, dict[str, Any]] = {} + configured_prefix = os.getenv("SPECO_MOONCAKE_KEY_PREFIX") + self._key_prefix = _safe_key( + configured_prefix or f"{socket.gethostname()}_{os.getpid()}" + ) + self._store: MooncakeTensorStore | None = None + self._store_setup_attempted = False + self._tp_rank: int | None = None + + @staticmethod + def _find_cache_group(kv_cache_config: "KVCacheConfig | None") -> int | None: + if kv_cache_config is None: + return None + for index, group in enumerate(kv_cache_config.kv_cache_groups): + if any("cache_only_layers" in name for name in group.layer_names): + return index + return None + + def _get_tp_rank(self) -> int: + if self._tp_rank is None: + try: + from vllm.distributed import get_tensor_model_parallel_rank + + self._tp_rank = int(get_tensor_model_parallel_rank()) + except Exception: # noqa: BLE001 + self._tp_rank = 0 + return self._tp_rank + + def _ensure_store(self) -> MooncakeTensorStore | None: + if self._store_setup_attempted: + return self._store + self._store_setup_attempted = True + if self._get_tp_rank() != 0: + return None + try: + store = MooncakeTensorStore(MooncakeTransferConfig.from_mapping()) + store.setup() + self._store = store + except Exception: # noqa: BLE001 + logger.exception("Failed to initialize SpeCo Mooncake connector") + return self._store + + def start_load_kv(self, *args: Any, **kwargs: Any) -> None: + return None + + def wait_for_layer_load(self, layer_name: str) -> None: + return None + + def wait_for_save(self) -> None: + return None + + def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]) -> None: + from vllm.model_executor.models.extract_hidden_states import ( + CacheOnlyAttentionLayer, + ) + + layers = get_layers_from_vllm_config( + self._vllm_config, CacheOnlyAttentionLayer, list(kv_caches) + ) + self._cache_layers = list(layers) + if len(self._cache_layers) != 1: + raise RuntimeError( + "Expected one extract_hidden_states cache layer, got " + f"{self._cache_layers}" + ) + + def save_kv_layer( + self, + layer_name: str, + kv_layer: torch.Tensor, + attn_metadata: AttentionMetadata, + **kwargs: Any, + ) -> None: + if layer_name not in self._cache_layers: + return + from vllm.model_executor.models.extract_hidden_states import ( + CacheOnlyAttentionMetadata, + ) + + if not isinstance(attn_metadata, CacheOnlyAttentionMetadata): + raise TypeError( + "Expected CacheOnlyAttentionMetadata for extracted hidden states" + ) + metadata = self._get_connector_metadata() + if not isinstance(metadata, SpeCoMooncakeConnectorMetadata): + raise TypeError("Unexpected connector metadata type") + store = self._ensure_store() + if store is None: + return + page_size = int(kv_layer.shape[1]) + for request in metadata.requests: + num_tokens = int(request.token_ids.numel()) + positions = _slot_mapping( + request.block_ids, page_size, num_tokens, kv_layer.device + ) + if int(positions.numel()) < num_tokens: + continue + all_hidden = kv_layer.flatten(0, 1)[positions][:num_tokens].reshape( + num_tokens, -1 + ) + split_at = self._training_layers * self._hidden_size + training_hidden = all_hidden[:, :split_at].reshape( + num_tokens, self._training_layers, self._hidden_size + ) + last_hidden = all_hidden[:, -self._hidden_size :].unsqueeze(1) + hidden_states = torch.cat((training_hidden, last_hidden), dim=1).to( + torch.bfloat16 + ) + key = f"{self._key_prefix}_{_safe_key(request.request_id)}" + result = store.put( + key, + { + "hidden_states": hidden_states, + "token_ids": request.token_ids, + }, + ) + response = self._response_metadata.get(request.request_id) + if response is not None: + response.update(result) + + def get_num_new_matched_tokens( + self, request: "Request", num_computed_tokens: int + ) -> tuple[int | None, bool]: + return 0, False + + def update_state_after_alloc( + self, + request: "Request", + blocks: "KVCacheBlocks", + num_external_tokens: int, + ) -> None: + if num_external_tokens != 0: + raise ValueError("SpeCo Mooncake connector is store-only") + + def build_connector_meta( + self, scheduler_output: SchedulerOutput + ) -> KVConnectorMetadata: + metadata = SpeCoMooncakeConnectorMetadata() + for request in scheduler_output.scheduled_new_reqs: + token_ids = request.prompt_token_ids or [] + group_id = self._cache_group_id + if group_id is None: + group_id = max( + range(len(request.block_ids)), + key=lambda index: len(request.block_ids[index]), + ) + self._cache_group_id = group_id + blocks = list(request.block_ids[group_id]) + metadata.add(request.req_id, token_ids, blocks) + self._active_requests[request.req_id] = request + self._request_blocks[request.req_id] = blocks + self._response_metadata[request.req_id] = { + "mooncake_key": ( + f"{self._key_prefix}_{_safe_key(request.req_id)}" + ), + "input_ids_list": token_ids, + "tensor_shapes": { + "hidden_states": ( + len(token_ids), + self._training_layers + 1, + self._hidden_size, + ), + "token_ids": (len(token_ids),), + }, + "tensor_dtypes": { + "hidden_states": "bfloat16", + "token_ids": "int64", + }, + } + + cached = scheduler_output.scheduled_cached_reqs + for index, request_id in enumerate(cached.req_ids): + if request_id not in self._active_requests: + continue + new_blocks = cached.new_block_ids[index] + if new_blocks is not None: + self._request_blocks[request_id].extend( + new_blocks[self._cache_group_id] + ) + request = self._active_requests[request_id] + metadata.add( + request_id, + request.prompt_token_ids or [], + self._request_blocks[request_id], + ) + return metadata + + def request_finished( + self, request: "Request", block_ids: list[int] + ) -> tuple[bool, dict[str, Any] | None]: + request_id = request.request_id + self._active_requests.pop(request_id, None) + self._request_blocks.pop(request_id, None) + return False, self._response_metadata.pop(request_id, None) + + def request_finished_all_groups( + self, request: "Request", block_ids: tuple[list[int], ...] + ) -> tuple[bool, dict[str, Any] | None]: + return self.request_finished(request, block_ids[0] if block_ids else []) + + @classmethod + def get_required_kvcache_layout(cls, vllm_config: VllmConfig) -> str | None: + return "NHD" diff --git a/verl_speco/trainer/draft_training_loop.py b/verl_speco/trainer/draft_training_loop.py index 9b706f55..53c75ceb 100644 --- a/verl_speco/trainer/draft_training_loop.py +++ b/verl_speco/trainer/draft_training_loop.py @@ -18,6 +18,7 @@ import asyncio import json import logging +import math import os import time from typing import Any, cast @@ -127,6 +128,7 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: last_saved_step = 0 store = None feature_replayer = None + feature_producer = None current_stage = "activate_training_model" try: stage_started = time.perf_counter() @@ -218,8 +220,58 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: bool(feature_store_cfg.get("shuffle", True)), bool(feature_store_cfg.get("repeat", True)), ) - for samples in loader: - if max_steps > 0 and successful_steps >= max_steps: + sample_source = loader + pipeline_cfg = training_cfg.get("target_feature_pipeline", {}) or {} + pipeline_enabled = bool(pipeline_cfg.get("enabled", False)) + if pipeline_enabled: + if feature_replayer is None or not feature_replayer.backend.startswith( + "vllm_" + ): + raise ValueError( + "target_feature_pipeline.enabled=true requires a vLLM replay " + "backend (vllm_file or vllm_mooncake)" + ) + from verl_speco.trainer.target_feature_pipeline import ( + TargetFeatureProducer, + ) + + feature_producer = TargetFeatureProducer( + loader, + feature_replayer, + rank=rank, + concurrency=max( + math.ceil( + int(pipeline_cfg.get("concurrency", 16) or 16) / world_size + ), + 1, + ), + transfer_concurrency=int( + max( + math.ceil( + int( + pipeline_cfg.get("transfer_concurrency", 8) or 8 + ) + / world_size + ), + 1, + ) + ), + producer_prefetch_depth=int( + pipeline_cfg.get("producer_prefetch_depth", 4) or 4 + ), + prefetch_depth=int(pipeline_cfg.get("prefetch_depth", 2) or 2), + queue_timeout=float(pipeline_cfg.get("queue_timeout", 300.0) or 300.0), + ) + sample_source = feature_producer + sample_iterator = iter(sample_source) + while max_steps <= 0 or successful_steps < max_steps: + current_stage = "load_next_batch" + samples = _next_batch_across_ranks( + sample_iterator, + rank=rank, + device=trainer.runtime_device, + ) + if samples is None: break step_started = time.perf_counter() attempted_batches += 1 @@ -252,7 +304,7 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: world_size=world_size, device=trainer.runtime_device, ) - if feature_replayer is not None: + if feature_replayer is not None and feature_producer is None: materialize_started = time.perf_counter() if log_batch_progress: logger.info( @@ -323,6 +375,8 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: ) if feature_replayer is not None: step_metrics.update(feature_replayer.metrics()) + if feature_producer is not None: + step_metrics.update(feature_producer.metrics()) _log_standalone_step_metrics(step_metrics, rank=rank) if save_interval > 0 and optimizer_step % save_interval == 0: current_stage = "save_checkpoint" @@ -358,10 +412,12 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: attempted_batches, successful_steps, ) - if store is not None: - store.close() + if feature_producer is not None: + feature_producer.close() if feature_replayer is not None: feature_replayer.close() + if store is not None: + store.close() logger.info("[standalone rank=%s] cleaning trainer resources", rank) await trainer.cleanup_training(clear_data=True) if dist.is_initialized(): @@ -934,13 +990,20 @@ def _log_standalone_step_metrics(metrics: dict[str, float], *, rank: int) -> Non ("perf/step_time", "step_time"), ("replay/cache_hit_ratio", "cache_hit"), ("replay/target_forward_time_total", "target_forward_total"), + ("replay/vllm_request_time_total", "vllm_request_total"), + ("replay/mooncake_get_time_total", "mooncake_get_total"), + ("producer/consumer_wait_time_total", "producer_wait_total"), + ("producer/ready_queue_size", "ready_batches"), ): if key not in metrics: continue value = float(metrics[key]) if key == "train/lr": fields.append(f"{label}={value:.3e}") - elif key in {"perf/step_time", "replay/target_forward_time_total"}: + elif key.endswith("_time_total") or key in { + "perf/step_time", + "replay/target_forward_time_total", + }: fields.append(f"{label}={value:.3f}s") else: fields.append(f"{label}={value:.4f}") @@ -991,6 +1054,64 @@ def _all_ranks_true(value: bool, device: torch.device) -> bool: return bool(ready.item()) +def _next_batch_across_ranks( + source, + *, + rank: int, + device: torch.device, +) -> list[Any] | None: + """Fetch one batch and make producer failures visible to every rank. + + Producer and Mooncake errors happen before the FSDP training step. Every + rank therefore reports its fetch result through the same collective before + any rank is allowed to enter model collectives. This prevents healthy + ranks from waiting in FSDP after another rank has already started cleanup. + """ + samples: list[Any] | None = None + local_error: BaseException | None = None + exhausted = False + try: + samples = next(source) + except StopIteration: + exhausted = True + except BaseException as exc: # noqa: BLE001 + local_error = exc + + state = torch.tensor( + [1 if local_error is not None else 0, 1 if exhausted else 0], + dtype=torch.int32, + device=device, + ) + if dist.is_initialized() and dist.get_world_size() > 1: + dist.all_reduce(state, op=dist.ReduceOp.MAX) + + any_failed = bool(state[0].item()) + any_exhausted = bool(state[1].item()) + if any_failed: + if local_error is not None: + raise RuntimeError( + f"Standalone target-feature producer failed on rank={rank}; " + "all ranks are stopping before the next training collective" + ) from local_error + raise RuntimeError( + "Standalone target-feature producer failed on another rank; " + f"rank={rank} is stopping before the next training collective" + ) + if any_exhausted: + if not exhausted: + logger.warning( + "[standalone rank=%s] discarding a prefetched batch because " + "another rank exhausted its data source", + rank, + ) + return None + if samples is None: + raise RuntimeError( + f"Rank={rank} reported a successful batch fetch without samples" + ) + return samples + + def _sync_any_rank_saved_checkpoint(saved: Any) -> bool: if not dist.is_initialized(): return bool(saved) diff --git a/verl_speco/trainer/mooncake_transfer.py b/verl_speco/trainer/mooncake_transfer.py new file mode 100644 index 00000000..8f494df1 --- /dev/null +++ b/verl_speco/trainer/mooncake_transfer.py @@ -0,0 +1,216 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +"""Small optional Mooncake client used by standalone target-feature replay. + +The payload is stored as one safetensors object. A single object makes the +producer response atomic and avoids the file creation/locking protocol used by +``vllm_file``. Mooncake is imported lazily so normal and online training do +not acquire a runtime dependency on it. +""" + +from __future__ import annotations + +import logging +import os +import socket +import time +from dataclasses import dataclass +from typing import Any + +import torch + +logger = logging.getLogger(__name__) + + +def _parse_size(value: str | int) -> int: + if isinstance(value, int): + return value + text = str(value).strip().upper() + multipliers = { + "TB": 1024**4, + "GB": 1024**3, + "MB": 1024**2, + "KB": 1024, + "T": 1024**4, + "G": 1024**3, + "M": 1024**2, + "K": 1024, + "B": 1, + } + for suffix in sorted(multipliers, key=len, reverse=True): + if text.endswith(suffix): + return int(float(text[: -len(suffix)]) * multipliers[suffix]) + return int(text) + + +@dataclass(frozen=True) +class MooncakeTransferConfig: + local_hostname: str + metadata_server: str + master_server_address: str + global_segment_size: int + local_buffer_size: int + protocol: str + device_name: str + get_timeout: float + get_poll_interval: float + + @classmethod + def from_mapping(cls, config: Any | None = None) -> "MooncakeTransferConfig": + config = config or {} + + def value(name: str, default: Any) -> Any: + getter = getattr(config, "get", None) + if callable(getter): + return getter(name, default) + return getattr(config, name, default) + + master = str( + value( + "master_server_address", + os.getenv("MOONCAKE_MASTER_SERVER", "127.0.0.1:50051"), + ) + ) + master_host = master.rsplit(":", 1)[0] + return cls( + local_hostname=str( + value( + "local_hostname", + os.getenv("MOONCAKE_LOCAL_HOSTNAME", socket.gethostname()), + ) + ), + metadata_server=str( + value( + "metadata_server", + os.getenv( + "MOONCAKE_METADATA_SERVER", + f"http://{master_host}:8090/metadata", + ), + ) + ), + master_server_address=master, + global_segment_size=_parse_size( + value( + "global_segment_size", + os.getenv("MOONCAKE_GLOBAL_SEGMENT_SIZE", "4GB"), + ) + ), + local_buffer_size=_parse_size( + value( + "local_buffer_size", + os.getenv("MOONCAKE_LOCAL_BUFFER_SIZE", "1GB"), + ) + ), + protocol=str( + value("protocol", os.getenv("MOONCAKE_PROTOCOL", "tcp")) + ), + device_name=str( + value("device_name", os.getenv("MOONCAKE_DEVICE_NAME", "")) + ), + get_timeout=float(value("get_timeout", 120.0)), + get_poll_interval=max(float(value("get_poll_interval", 0.02)), 0.001), + ) + + def export_environment(self) -> None: + os.environ["MOONCAKE_LOCAL_HOSTNAME"] = self.local_hostname + os.environ["MOONCAKE_METADATA_SERVER"] = self.metadata_server + os.environ["MOONCAKE_MASTER_SERVER"] = self.master_server_address + os.environ["MOONCAKE_GLOBAL_SEGMENT_SIZE"] = str(self.global_segment_size) + os.environ["MOONCAKE_LOCAL_BUFFER_SIZE"] = str(self.local_buffer_size) + os.environ["MOONCAKE_PROTOCOL"] = self.protocol + os.environ["MOONCAKE_DEVICE_NAME"] = self.device_name + if self.protocol.lower() == "tcp": + os.environ.setdefault("MC_STORE_MEMCPY", "0") + + +class MooncakeTensorStore: + """Store and retrieve a tensor dictionary as one Mooncake object.""" + + def __init__(self, config: MooncakeTransferConfig): + self.config = config + self._store: Any | None = None + + def setup(self) -> None: + if self._store is not None: + return + self.config.export_environment() + try: + from mooncake.store import MooncakeDistributedStore + except ImportError as exc: + raise RuntimeError( + "Mooncake replay requires mooncake-transfer-engine " + "(use mooncake-transfer-engine-npu on Ascend)" + ) from exc + store = MooncakeDistributedStore() + result = store.setup( + local_hostname=self.config.local_hostname, + metadata_server=self.config.metadata_server, + global_segment_size=self.config.global_segment_size, + local_buffer_size=self.config.local_buffer_size, + protocol=self.config.protocol, + rdma_devices=self.config.device_name, + master_server_addr=self.config.master_server_address, + ) + if result not in (None, 0): + raise RuntimeError(f"Mooncake client setup failed with code {result}") + self._store = store + + def put(self, key: str, tensors: dict[str, torch.Tensor]) -> dict[str, Any]: + self.setup() + assert self._store is not None + from safetensors.torch import save + + cpu_tensors = { + name: tensor.detach().to("cpu").contiguous() + for name, tensor in tensors.items() + } + payload = save(cpu_tensors) + result = self._store.put(key, payload) + if result not in (None, 0): + raise RuntimeError(f"Mooncake put failed for {key!r}: code={result}") + return { + "mooncake_key": key, + "tensor_shapes": { + name: tuple(tensor.shape) for name, tensor in cpu_tensors.items() + }, + "tensor_dtypes": { + name: str(tensor.dtype).removeprefix("torch.") + for name, tensor in cpu_tensors.items() + }, + "payload_bytes": len(payload), + } + + def get(self, key: str) -> dict[str, torch.Tensor]: + self.setup() + assert self._store is not None + from safetensors.torch import load + + deadline = time.monotonic() + self.config.get_timeout + while True: + payload = self._store.get(key) + if payload is not None: + return dict(load(bytes(payload))) + if time.monotonic() >= deadline: + raise TimeoutError( + f"Mooncake object {key!r} was unavailable for " + f"{self.config.get_timeout:.1f}s" + ) + time.sleep(self.config.get_poll_interval) + + def remove(self, key: str) -> None: + if self._store is None: + return + try: + remove = getattr(self._store, "remove", None) + if callable(remove): + remove(key, True) + else: + self._store.batch_remove([key], force=True) + except Exception: # noqa: BLE001 + logger.warning("Failed to remove Mooncake object %s", key, exc_info=True) + + def close(self) -> None: + if self._store is not None and hasattr(self._store, "close"): + self._store.close() + self._store = None diff --git a/verl_speco/trainer/target_feature_pipeline.py b/verl_speco/trainer/target_feature_pipeline.py new file mode 100644 index 00000000..916616f5 --- /dev/null +++ b/verl_speco/trainer/target_feature_pipeline.py @@ -0,0 +1,196 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +"""Bounded producer/prefetch pipeline for standalone target features.""" + +from __future__ import annotations + +import logging +import queue +import threading +import time +from collections import deque +from concurrent.futures import Future, ThreadPoolExecutor +from dataclasses import dataclass +from typing import Any, Iterable, Iterator + +from verl_speco.trainer.feature_store import DraftFeatureSample + +logger = logging.getLogger(__name__) + + +@dataclass +class _PipelineFailure: + error: BaseException + + +_END = object() + + +class TargetFeatureProducer: + """Materialize future batches while the current FSDP step is running. + + The coordinator owns the source iterator and puts complete batches into a + bounded ready queue. Sample requests inside each batch are concurrent. + Consequently a training rank never observes a partially materialized batch. + """ + + def __init__( + self, + source: Iterable[list[Any]], + replayer: Any, + *, + rank: int, + concurrency: int, + transfer_concurrency: int, + producer_prefetch_depth: int, + prefetch_depth: int, + queue_timeout: float, + ): + self.source = iter(source) + self.replayer = replayer + self.rank = int(rank) + self.concurrency = max(int(concurrency), 1) + self.transfer_concurrency = max(int(transfer_concurrency), 1) + self.producer_prefetch_depth = max(int(producer_prefetch_depth), 1) + self.prefetch_depth = max(int(prefetch_depth), 1) + self.queue_timeout = max(float(queue_timeout), 1.0) + self._ready: queue.Queue[Any] = queue.Queue(maxsize=self.prefetch_depth) + self._stop = threading.Event() + self._request_executor = ThreadPoolExecutor( + max_workers=self.concurrency, + thread_name_prefix=f"speco-request-r{self.rank}", + ) + self._transfer_executor = ThreadPoolExecutor( + max_workers=self.transfer_concurrency, + thread_name_prefix=f"speco-transfer-r{self.rank}", + ) + self._thread = threading.Thread( + target=self._run, + name=f"speco-target-producer-r{self.rank}", + daemon=True, + ) + self.produced_batches = 0 + self.produced_samples = 0 + self.producer_seconds = 0.0 + self.queue_wait_seconds = 0.0 + self.consumer_wait_seconds = 0.0 + self.transfer_seconds = 0.0 + self.failed_batches = 0 + self._thread.start() + logger.info( + "[target producer rank=%s] started request_concurrency=%s " + "transfer_concurrency=%s producer_prefetch_depth=%s prefetch_depth=%s", + self.rank, + self.concurrency, + self.transfer_concurrency, + self.producer_prefetch_depth, + self.prefetch_depth, + ) + + def _run(self) -> None: + try: + pending: deque[tuple[float, list[Future[Any]]]] = deque() + + def submit_next() -> bool: + if self._stop.is_set(): + return False + try: + samples = next(self.source) + except StopIteration: + return False + started = time.perf_counter() + if self.replayer.backend == "vllm_mooncake": + futures = [ + self._request_executor.submit( + self.replayer.produce_mooncake_descriptor, sample + ) + for sample in samples + ] + else: + futures = [ + self._request_executor.submit(self.replayer.materialize, [sample]) + for sample in samples + ] + pending.append((started, futures)) + return True + + for _ in range(self.producer_prefetch_depth): + if not submit_next(): + break + + while pending and not self._stop.is_set(): + started, request_futures = pending.popleft() + produced = [future.result() for future in request_futures] + self.producer_seconds += time.perf_counter() - started + submit_next() + transfer_started = time.perf_counter() + if self.replayer.backend == "vllm_mooncake": + transfer_futures = [ + self._transfer_executor.submit( + self.replayer.consume_pipeline_mooncake_descriptor, + descriptor, + ) + for descriptor in produced + ] + batch = [future.result() for future in transfer_futures] + else: + batch = [item for group in produced for item in group] + self.transfer_seconds += time.perf_counter() - transfer_started + self.produced_batches += 1 + self.produced_samples += len(batch) + if self.replayer.backend == "vllm_mooncake": + self.replayer.record_pipeline_materialized(len(batch)) + self._put(batch) + self._put(_END) + except BaseException as exc: # noqa: BLE001 + self.failed_batches += 1 + self._put(_PipelineFailure(exc)) + + def _put(self, value: Any) -> None: + started = time.perf_counter() + while not self._stop.is_set(): + try: + self._ready.put(value, timeout=0.2) + self.queue_wait_seconds += time.perf_counter() - started + return + except queue.Full: + continue + + def __iter__(self) -> Iterator[list[DraftFeatureSample]]: + return self + + def __next__(self) -> list[DraftFeatureSample]: + started = time.perf_counter() + try: + value = self._ready.get(timeout=self.queue_timeout) + except queue.Empty as exc: + raise TimeoutError( + "Timed out waiting for target-feature producer; inspect the vLLM " + "and Mooncake logs for a stalled request or missing object" + ) from exc + self.consumer_wait_seconds += time.perf_counter() - started + if value is _END: + raise StopIteration + if isinstance(value, _PipelineFailure): + raise RuntimeError("Target-feature producer failed") from value.error + return value + + def metrics(self) -> dict[str, float]: + return { + "producer/batches_total": float(self.produced_batches), + "producer/samples_total": float(self.produced_samples), + "producer/materialize_time_total": float(self.producer_seconds), + "producer/transfer_time_total": float(self.transfer_seconds), + "producer/queue_block_time_total": float(self.queue_wait_seconds), + "producer/consumer_wait_time_total": float(self.consumer_wait_seconds), + "producer/ready_queue_size": float(self._ready.qsize()), + "producer/ready_queue_capacity": float(self.prefetch_depth), + "producer/failed_batches_total": float(self.failed_batches), + } + + def close(self) -> None: + self._stop.set() + self._request_executor.shutdown(wait=True, cancel_futures=True) + self._transfer_executor.shutdown(wait=True, cancel_futures=True) + self._thread.join(timeout=5.0) diff --git a/verl_speco/trainer/target_feature_replay.py b/verl_speco/trainer/target_feature_replay.py index 0667fff1..4d590380 100644 --- a/verl_speco/trainer/target_feature_replay.py +++ b/verl_speco/trainer/target_feature_replay.py @@ -21,7 +21,9 @@ import logging import os import tempfile +import threading import time +from dataclasses import dataclass from pathlib import Path from typing import Any, Iterable, cast @@ -36,6 +38,27 @@ logger = logging.getLogger(__name__) +@dataclass(frozen=True) +class MooncakeReplayDescriptor: + sample: DraftReplaySample + prompt_ids: list[int] + key: str + + +@dataclass +class _VllmEndpointState: + index: int + url: str + client: Any | None = None + model: str | None = None + inflight: int = 0 + requests: int = 0 + failures: int = 0 + consecutive_failures: int = 0 + request_seconds: float = 0.0 + cooldown_until: float = 0.0 + + def _config_value(config: Any, key: str, default: Any = None) -> Any: if config is None: return default @@ -44,6 +67,26 @@ def _config_value(config: Any, key: str, default: Any = None) -> Any: return getattr(config, key, default) +def _normalize_vllm_endpoints(config: Any) -> list[str]: + configured = _config_value(config, "vllm_endpoints", None) + if configured is None: + configured = [ + _config_value(config, "vllm_endpoint", "http://localhost:8000/v1") + ] + elif isinstance(configured, str): + configured = [configured] + endpoints: list[str] = [] + for value in configured: + endpoint = str(value or "").strip().rstrip("/") + if endpoint and endpoint not in endpoints: + endpoints.append(endpoint) + if not endpoints: + raise ValueError( + "target_feature_replay.vllm_endpoints must contain at least one URL" + ) + return endpoints + + def _parse_dtype(value: Any) -> torch.dtype: normalized = str(value or "bfloat16").strip().lower() dtypes = { @@ -336,10 +379,12 @@ def __init__( .strip() .lower() ) - if self.backend not in {"torch", "vllm_file"}: + if self.backend == "mooncake": + self.backend = "vllm_mooncake" + if self.backend not in {"torch", "vllm_file", "vllm_mooncake"}: raise ValueError( f"Unsupported target_feature_replay.backend={self.backend!r}; " - "expected 'torch' or 'vllm_file'" + "expected 'torch', 'vllm_file', or 'vllm_mooncake'" ) configured_model_path = _config_value(self.replay_cfg, "model_path", None) model_path = configured_model_path or self.draft_config.model.path @@ -369,10 +414,8 @@ def __init__( self.logits_chunk_rows = max( int(_config_value(self.replay_cfg, "logits_chunk_rows", 32) or 32), 1 ) - self.vllm_endpoint = str( - _config_value(self.replay_cfg, "vllm_endpoint", "http://localhost:8000/v1") - or "http://localhost:8000/v1" - ) + self.vllm_endpoints = _normalize_vllm_endpoints(self.replay_cfg) + self.vllm_endpoint = self.vllm_endpoints[0] self.vllm_model = _config_value(self.replay_cfg, "vllm_model", None) self.vllm_timeout = float( _config_value(self.replay_cfg, "request_timeout", 120.0) or 120.0 @@ -380,6 +423,10 @@ def __init__( self.vllm_max_retries = max( int(_config_value(self.replay_cfg, "max_retries", 3) or 0), 0 ) + self.vllm_endpoint_cooldown = max( + float(_config_value(self.replay_cfg, "endpoint_cooldown", 5.0) or 0.0), + 0.0, + ) self.vllm_on_generate = ( str(_config_value(self.replay_cfg, "on_generate", "delete") or "delete") .strip() @@ -440,6 +487,7 @@ def __init__( self.target_config_fingerprint = hashlib.sha256(config_json).hexdigest() self.cache: BoundedReplayCache | None = None + self._cache_lock = threading.Lock() cache_cfg = _config_value(self.replay_cfg, "cache", {}) or {} if bool(_config_value(cache_cfg, "enabled", False)): cache_path = _config_value(cache_cfg, "path", None) @@ -460,25 +508,38 @@ def __init__( self.output_embedding: nn.Module | None = None self.vllm_client: Any | None = None self.vllm_resolved_model: str | None = None + self._vllm_endpoint_states = [ + _VllmEndpointState(index=index, url=endpoint) + for index, endpoint in enumerate(self.vllm_endpoints) + ] + self._vllm_clients_initialized = False + self.mooncake_store: Any | None = None + self._client_lock = threading.Lock() + self._endpoint_lock = threading.Lock() + self._metrics_lock = threading.Lock() + self._pending_keys_lock = threading.Lock() + self._pending_mooncake_keys: set[str] = set() self.cache_hits = 0 self.cache_misses = 0 self.materialized_samples = 0 self.target_forward_seconds = 0.0 self.vllm_request_seconds = 0.0 self.vllm_requests = 0 + self.mooncake_get_seconds = 0.0 + self.mooncake_gets = 0 self._warned_replay_algorithm_mismatch = False self._warned_replay_layer_mismatch = False self._warned_replay_layout_mismatch = False logger.info( "[target replay rank=%s] initialized backend=%s algorithm=%s " - "target_layers=%s hidden_layout=%s use_logits=%s endpoint=%s cache=%s", + "target_layers=%s hidden_layout=%s use_logits=%s endpoints=%s cache=%s", self.rank, self.backend, self.algorithm, self.target_layer_ids, self.hidden_layout, self.use_logits, - self.vllm_endpoint if self.backend == "vllm_file" else None, + self.vllm_endpoints if self.backend.startswith("vllm_") else None, self.cache is not None, ) @@ -498,15 +559,19 @@ def materialize( ) self._validate_target_path(sample) key = self._cache_key(sample) - cached = self.cache.get(key) if self.cache is not None else None + with self._cache_lock: + cached = self.cache.get(key) if self.cache is not None else None if cached is not None: - self.cache_hits += 1 + with self._metrics_lock: + self.cache_hits += 1 materialized.append(cached) continue - self.cache_misses += 1 + with self._metrics_lock: + self.cache_misses += 1 replayed = self._materialize_one(sample) if self.cache is not None: - self.cache.put(key, replayed) + with self._cache_lock: + self.cache.put(key, replayed) materialized.append(replayed) except Exception: metadata = getattr(sample, "metadata", {}) or {} @@ -520,7 +585,8 @@ def materialize( metadata.get("global_step"), ) raise - self.materialized_samples += len(materialized) + with self._metrics_lock: + self.materialized_samples += len(materialized) return materialized def _validate_target_path(self, sample: DraftReplaySample) -> None: @@ -638,6 +704,8 @@ def _ensure_model(self) -> None: def _materialize_one(self, sample: DraftReplaySample) -> DraftFeatureSample: if self.backend == "vllm_file": return self._materialize_one_vllm_file(sample) + if self.backend == "vllm_mooncake": + return self._materialize_one_vllm_mooncake(sample) return self._materialize_one_torch(sample) def _materialize_one_torch(self, sample: DraftReplaySample) -> DraftFeatureSample: @@ -806,6 +874,117 @@ def _materialize_one_vllm_file( logger.warning("Failed to delete vLLM hidden-states file %s", path) return feature + def _materialize_one_vllm_mooncake( + self, sample: DraftReplaySample + ) -> DraftFeatureSample: + return self.consume_mooncake_descriptor( + self._request_mooncake_descriptor(sample) + ) + + def produce_mooncake_descriptor( + self, sample: DraftReplaySample | DraftFeatureSample + ) -> MooncakeReplayDescriptor | DraftFeatureSample: + if isinstance(sample, DraftFeatureSample): + return sample + if not isinstance(sample, DraftReplaySample): + raise TypeError( + "Mooncake producer expected DraftReplaySample or " + f"DraftFeatureSample, got {type(sample)!r}" + ) + if self.use_logits: + raise NotImplementedError( + "target_feature_replay.backend=vllm_mooncake does not yet " + "support training.use_logits=true; use backend=torch." + ) + self._validate_target_path(sample) + cache_key = self._cache_key(sample) + with self._cache_lock: + cached = self.cache.get(cache_key) if self.cache is not None else None + if cached is not None: + with self._metrics_lock: + self.cache_hits += 1 + return cached + with self._metrics_lock: + self.cache_misses += 1 + return self._request_mooncake_descriptor(sample) + + def _request_mooncake_descriptor( + self, sample: DraftReplaySample + ) -> MooncakeReplayDescriptor: + self._validate_vllm_positions(sample) + feature_positions = sample.feature_positions.detach().cpu().long() + feature_end = int(feature_positions[-1].item()) + 1 + prompt_ids = sample.input_ids[:feature_end].detach().cpu().long().tolist() + response = self._request_vllm_response(prompt_ids) + params = getattr(response, "kv_transfer_params", None) + if params is None: + raise ValueError("vLLM response missing kv_transfer_params") + key = params.get("mooncake_key") + if not key: + raise ValueError( + "vLLM response missing mooncake_key; start vLLM with " + "SpeCoMooncakeHiddenStatesConnector" + ) + key = str(key) + with self._pending_keys_lock: + self._pending_mooncake_keys.add(key) + return MooncakeReplayDescriptor(sample, prompt_ids, key) + + def consume_mooncake_descriptor( + self, descriptor: MooncakeReplayDescriptor | DraftFeatureSample + ) -> DraftFeatureSample: + if isinstance(descriptor, DraftFeatureSample): + return descriptor + store = self._ensure_mooncake_store() + transfer_started = time.perf_counter() + payload = store.get(descriptor.key) + with self._metrics_lock: + self.mooncake_get_seconds += time.perf_counter() - transfer_started + self.mooncake_gets += 1 + try: + feature = self._feature_from_vllm_payload( + descriptor.sample, + payload, + prompt_ids=descriptor.prompt_ids, + source="token_replay_vllm_mooncake", + ) + return feature + finally: + if self.vllm_on_generate == "delete": + store.remove(descriptor.key) + with self._pending_keys_lock: + self._pending_mooncake_keys.discard(descriptor.key) + + def consume_pipeline_mooncake_descriptor( + self, descriptor: MooncakeReplayDescriptor | DraftFeatureSample + ) -> DraftFeatureSample: + feature = self.consume_mooncake_descriptor(descriptor) + if isinstance(descriptor, MooncakeReplayDescriptor) and self.cache is not None: + with self._cache_lock: + self.cache.put(self._cache_key(descriptor.sample), feature) + return feature + + def record_pipeline_materialized(self, count: int) -> None: + with self._metrics_lock: + self.materialized_samples += int(count) + + def _ensure_mooncake_store(self): + if self.mooncake_store is not None: + return self.mooncake_store + with self._client_lock: + if self.mooncake_store is None: + from verl_speco.trainer.mooncake_transfer import ( + MooncakeTensorStore, + MooncakeTransferConfig, + ) + + mooncake_cfg = _config_value(self.replay_cfg, "mooncake", {}) or {} + self.mooncake_store = MooncakeTensorStore( + MooncakeTransferConfig.from_mapping(mooncake_cfg) + ) + self.mooncake_store.setup() + return self.mooncake_store + def _validate_vllm_positions(self, sample: DraftReplaySample) -> None: if not self.vllm_require_arange_positions: return @@ -815,66 +994,219 @@ def _validate_vllm_positions(self, sample: DraftReplaySample) -> None: actual = sample.position_ids[:feature_end].detach().cpu().long() if not torch.equal(actual, expected): raise ValueError( - "target_feature_replay.backend=vllm_file currently requires " + "vLLM target replay currently requires " "position_ids to be contiguous arange positions for the replay prefix" ) - def _ensure_vllm_client(self) -> None: - if self.vllm_client is not None: + def _ensure_vllm_clients(self) -> None: + if self._vllm_clients_initialized: return + with self._client_lock: + if self._vllm_clients_initialized: + return + started = time.perf_counter() + logger.info( + "[target replay rank=%s] initializing vLLM endpoint pool=%s " + "configured_model=%s", + self.rank, + self.vllm_endpoints, + self.vllm_model, + ) + try: + import openai + except ImportError as exc: + raise RuntimeError( + "vLLM target replay requires the openai package" + ) from exc + resolved_models: set[str] = set() + for state in self._vllm_endpoint_states: + state.client = openai.OpenAI( + base_url=state.url, + api_key="EMPTY", + max_retries=0, + ) + state.model = ( + os.fspath(self.vllm_model) if self.vllm_model else self.model_path + ) + try: + models = state.client.models.list(timeout=self.vllm_timeout) + if not self.vllm_model and models.data: + state.model = str(models.data[0].id) + resolved_models.add(str(state.model)) + logger.info( + "[target replay rank=%s] vLLM endpoint[%s] ready url=%s model=%s", + self.rank, + state.index, + state.url, + state.model, + ) + except Exception as exc: # noqa: BLE001 + state.failures += 1 + state.consecutive_failures += 1 + state.cooldown_until = ( + time.monotonic() + self.vllm_endpoint_cooldown + ) + logger.warning( + "[target replay rank=%s] vLLM endpoint[%s] health check " + "failed; requests may retry it after cooldown url=%s error=%r", + self.rank, + state.index, + state.url, + exc, + ) + self._vllm_clients_initialized = True + self.vllm_client = self._vllm_endpoint_states[0].client + self.vllm_resolved_model = self._vllm_endpoint_states[0].model + if len(resolved_models) > 1: + logger.warning( + "[target replay rank=%s] vLLM endpoints advertise different " + "models=%s; set target_feature_replay.vllm_model explicitly " + "after confirming all servers use identical target weights", + self.rank, + sorted(resolved_models), + ) + logger.info( + "[target replay rank=%s] vLLM endpoint pool initialized " + "endpoints=%s elapsed=%.3fs", + self.rank, + len(self._vllm_endpoint_states), + time.perf_counter() - started, + ) + + def _acquire_vllm_endpoint( + self, excluded: set[int] | None = None + ) -> _VllmEndpointState: + self._ensure_vllm_clients() + excluded = excluded or set() + now = time.monotonic() + with self._endpoint_lock: + candidates = [ + state + for state in self._vllm_endpoint_states + if state.client is not None + and state.model is not None + and state.index not in excluded + and state.cooldown_until <= now + ] + if not candidates: + candidates = [ + state + for state in self._vllm_endpoint_states + if state.client is not None + and state.model is not None + and state.index not in excluded + ] + if not candidates: + candidates = [ + state + for state in self._vllm_endpoint_states + if state.client is not None and state.model is not None + ] + if not candidates: + raise RuntimeError("No configured vLLM endpoint has a usable client") + state = min( + candidates, + key=lambda item: ( + item.inflight, + item.consecutive_failures, + item.cooldown_until, + item.index, + ), + ) + state.inflight += 1 + return state + + def _release_vllm_endpoint( + self, + state: _VllmEndpointState, + *, + elapsed: float, + succeeded: bool, + ) -> None: + with self._endpoint_lock: + state.inflight = max(state.inflight - 1, 0) + state.request_seconds += float(elapsed) + if succeeded: + state.requests += 1 + state.consecutive_failures = 0 + state.cooldown_until = 0.0 + else: + state.failures += 1 + state.consecutive_failures += 1 + state.cooldown_until = ( + time.monotonic() + self.vllm_endpoint_cooldown + ) + + def _request_vllm_response(self, prompt_ids: list[int]) -> Any: + last_error: Exception | None = None started = time.perf_counter() - logger.info( - "[target replay rank=%s] connecting to vLLM endpoint=%s configured_model=%s", - self.rank, - self.vllm_endpoint, - self.vllm_model, - ) - try: - import openai - except ImportError as exc: - raise RuntimeError( - "target_feature_replay.backend=vllm_file requires the openai package" - ) from exc - self.vllm_client = openai.OpenAI( - base_url=self.vllm_endpoint, - api_key="EMPTY", - max_retries=0, - ) - if self.vllm_model: - self.vllm_resolved_model = os.fspath(self.vllm_model) - else: - models = self.vllm_client.models.list() - self.vllm_resolved_model = models.data[0].id - logger.info( - "[target replay rank=%s] connected to vLLM model=%s elapsed=%.3fs", - self.rank, - self.vllm_resolved_model, - time.perf_counter() - started, - ) + attempted_endpoints: set[int] = set() + for attempt in range(self.vllm_max_retries + 1): + state = self._acquire_vllm_endpoint(attempted_endpoints) + attempt_started = time.perf_counter() + try: + response = state.client.completions.create( + model=state.model, + prompt=prompt_ids, + max_tokens=1, + extra_body={"return_token_ids": True}, + timeout=self.vllm_timeout, + ) + choices = getattr(response, "choices", None) or [] + if choices: + actual = getattr(choices[0], "prompt_token_ids", None) + if actual is not None and list(actual) != prompt_ids: + raise ValueError("vLLM prompt_token_ids mismatch") + with self._metrics_lock: + self.vllm_requests += 1 + self.vllm_request_seconds += time.perf_counter() - started + self._release_vllm_endpoint( + state, + elapsed=time.perf_counter() - attempt_started, + succeeded=True, + ) + return response + except Exception as exc: # noqa: BLE001 + last_error = exc + attempted_endpoints.add(state.index) + self._release_vllm_endpoint( + state, + elapsed=time.perf_counter() - attempt_started, + succeeded=False, + ) + if attempt >= self.vllm_max_retries: + break + time.sleep(float(2**attempt)) + with self._metrics_lock: + self.vllm_request_seconds += time.perf_counter() - started + raise RuntimeError( + "Failed to request vLLM hidden states after " + f"{self.vllm_max_retries + 1} attempts: {last_error}" + ) from last_error def _request_vllm_hidden_states(self, prompt_ids: list[int]) -> dict[str, Any]: - self._ensure_vllm_client() - assert self.vllm_client is not None - assert self.vllm_resolved_model is not None last_error: Exception | None = None started = time.perf_counter() request_index = self.vllm_requests + 1 log_request = request_index <= 2 or request_index % 100 == 0 + attempted_endpoints: set[int] = set() for attempt in range(self.vllm_max_retries + 1): + state = self._acquire_vllm_endpoint(attempted_endpoints) try: attempt_started = time.perf_counter() if log_request: logger.info( "[target replay rank=%s] vLLM request starting request=%s " - "attempt=%s/%s prompt_tokens=%s", + "attempt=%s/%s endpoint=%s prompt_tokens=%s", self.rank, request_index, attempt + 1, self.vllm_max_retries + 1, + state.url, len(prompt_ids), ) - response = self.vllm_client.completions.create( - model=self.vllm_resolved_model, + response = state.client.completions.create( + model=state.model, prompt=prompt_ids, max_tokens=1, extra_body={"return_token_ids": True}, @@ -883,8 +1215,14 @@ def _request_vllm_hidden_states(self, prompt_ids: list[int]) -> dict[str, Any]: path = self._extract_hidden_states_path(response, prompt_ids) payload = self._load_vllm_hidden_states(path) payload["_path"] = path - self.vllm_requests += 1 - self.vllm_request_seconds += time.perf_counter() - started + with self._metrics_lock: + self.vllm_requests += 1 + self.vllm_request_seconds += time.perf_counter() - started + self._release_vllm_endpoint( + state, + elapsed=time.perf_counter() - attempt_started, + succeeded=True, + ) if log_request: hidden_states = payload.get("hidden_states") hidden_shape = ( @@ -894,10 +1232,11 @@ def _request_vllm_hidden_states(self, prompt_ids: list[int]) -> dict[str, Any]: ) logger.info( "[target replay rank=%s] vLLM request completed request=%s " - "attempt=%s path=%s hidden_shape=%s elapsed=%.3fs", + "attempt=%s endpoint=%s path=%s hidden_shape=%s elapsed=%.3fs", self.rank, request_index, attempt + 1, + state.url, path, hidden_shape, time.perf_counter() - attempt_started, @@ -905,13 +1244,20 @@ def _request_vllm_hidden_states(self, prompt_ids: list[int]) -> dict[str, Any]: return payload except Exception as exc: # noqa: BLE001 last_error = exc + attempted_endpoints.add(state.index) + self._release_vllm_endpoint( + state, + elapsed=time.perf_counter() - attempt_started, + succeeded=False, + ) logger.warning( "[target replay rank=%s] vLLM request failed request=%s " - "attempt=%s/%s prompt_tokens=%s elapsed=%.3fs error=%r", + "attempt=%s/%s endpoint=%s prompt_tokens=%s elapsed=%.3fs error=%r", self.rank, request_index, attempt + 1, self.vllm_max_retries + 1, + state.url, len(prompt_ids), time.perf_counter() - started, exc, @@ -919,7 +1265,8 @@ def _request_vllm_hidden_states(self, prompt_ids: list[int]) -> dict[str, Any]: if attempt >= self.vllm_max_retries: break time.sleep(float(2**attempt)) - self.vllm_request_seconds += time.perf_counter() - started + with self._metrics_lock: + self.vllm_request_seconds += time.perf_counter() - started raise RuntimeError( f"Failed to request vLLM hidden states after " f"{self.vllm_max_retries + 1} attempts: {last_error}" @@ -1096,11 +1443,28 @@ def metrics(self) -> dict[str, float]: "replay/materialized_samples_total": float(self.materialized_samples), "replay/target_forward_time_total": float(self.target_forward_seconds), } - if self.backend == "vllm_file": + if self.backend.startswith("vllm_"): metrics["replay/vllm_requests_total"] = float(self.vllm_requests) metrics["replay/vllm_request_time_total"] = float( self.vllm_request_seconds ) + with self._endpoint_lock: + metrics["replay/vllm_endpoints_total"] = float( + len(self._vllm_endpoint_states) + ) + for state in self._vllm_endpoint_states: + prefix = f"replay/vllm_endpoint_{state.index}" + metrics[f"{prefix}_inflight"] = float(state.inflight) + metrics[f"{prefix}_requests_total"] = float(state.requests) + metrics[f"{prefix}_failures_total"] = float(state.failures) + metrics[f"{prefix}_request_time_total"] = float( + state.request_seconds + ) + if self.backend == "vllm_mooncake": + metrics["replay/mooncake_gets_total"] = float(self.mooncake_gets) + metrics["replay/mooncake_get_time_total"] = float( + self.mooncake_get_seconds + ) total = self.cache_hits + self.cache_misses if total > 0: metrics["replay/cache_hit_ratio"] = self.cache_hits / float(total) @@ -1109,6 +1473,28 @@ def metrics(self) -> dict[str, float]: return metrics def close(self) -> None: + for state in self._vllm_endpoint_states: + client = state.client + close = getattr(client, "close", None) + if callable(close): + try: + close() + except Exception: # noqa: BLE001 + logger.warning( + "Failed to close vLLM endpoint client %s", + state.url, + exc_info=True, + ) + state.client = None + if self.mooncake_store is not None: + with self._pending_keys_lock: + pending_keys = tuple(self._pending_mooncake_keys) + self._pending_mooncake_keys.clear() + if self.vllm_on_generate == "delete": + for key in pending_keys: + self.mooncake_store.remove(key) + self.mooncake_store.close() + self.mooncake_store = None if self.model is None: return try: From 75ac479ce48752a732c776cd5fe1e70e88b681e4 Mon Sep 17 00:00:00 2001 From: xyr <1615681145@qq.com> Date: Wed, 19 Aug 2026 15:18:25 +0800 Subject: [PATCH 31/50] Add standalone TransferQueue foundation Implement the shared drafter sample protocol, standalone TQ owner and client bridge, smoke coverage, and development documentation. Assisted-by: OpenAI Codex --- ...sync_vllm_mooncake_dspark_training_plan.md | 2896 +++++++++++++++++ ...standalone_tq_foundation_implementation.md | 1037 ++++++ ...standalone_vllm_tq_dspark_training_plan.md | 1131 +++++++ examples/run_dspark_tq_owner.sh | 24 + examples/tq_connection_smoke.py | 205 ++ pyproject.toml | 4 + tests/unit/test_drafter_sample_protocol.py | 126 + tests/unit/test_transferqueue_bridge.py | 183 ++ verl_speco/config/speco_base.yaml | 24 + .../integration/transferqueue_bridge.py | 259 +- verl_speco/tq_owner.py | 115 + verl_speco/transport/__init__.py | 38 + .../transport/drafter_sample_protocol.py | 393 +++ 13 files changed, 6431 insertions(+), 4 deletions(-) create mode 100644 docs/async_vllm_mooncake_dspark_training_plan.md create mode 100644 docs/standalone_tq_foundation_implementation.md create mode 100644 docs/standalone_vllm_tq_dspark_training_plan.md create mode 100644 examples/run_dspark_tq_owner.sh create mode 100644 examples/tq_connection_smoke.py create mode 100644 tests/unit/test_drafter_sample_protocol.py create mode 100644 tests/unit/test_transferqueue_bridge.py create mode 100644 verl_speco/tq_owner.py create mode 100644 verl_speco/transport/__init__.py create mode 100644 verl_speco/transport/drafter_sample_protocol.py diff --git a/docs/async_vllm_mooncake_dspark_training_plan.md b/docs/async_vllm_mooncake_dspark_training_plan.md new file mode 100644 index 00000000..6cda0e40 --- /dev/null +++ b/docs/async_vllm_mooncake_dspark_training_plan.md @@ -0,0 +1,2896 @@ +# verl-SpeCo 异步 vLLM → TransferQueue(Mooncake 后端)→ DSpark 流式训练方案 + +## 1. 文档范围 + +本文只讨论当前 `verl-SpeCo-ls` 项目的独立草稿模型训练入口: + +```text +examples/run_qwen3-8b_drafter_separate_training.sh + → python -m verl_speco.draft_train_launcher + → torch.distributed.run + → python -m verl_speco.draft_train + → run_standalone_draft_training() +``` + +目标是训练 `mode=train/offline` 的草稿模型,不是 RL 与草稿模型一起训练,也不是先生成全量 hidden states 再长期保存到磁盘。 + +输入文件已经包含提前生成好的 response。新流水线需要: + +1. Producer 读取 prompt 和预生成 response,构造完整 token 序列; +2. Producer 并行请求一个或多个 vLLM endpoint 做 prefill; +3. hidden states 写入 TransferQueue(简称 TQ),TQ 的数据后端使用 Mooncake; +4. DSpark 的各训练 rank 从 TQ/Mooncake 并行读取自己的 local batch; +5. 所有 rank 完成同一个 optimizer step 后,清理这一批 TQ 数据; +6. 不使用自研 Stream Coordinator;第一版复用 PR #48 已验证的 TQ KV 传输模式,由 TQ 保存 payload 和 tags,rank 0 通过 `kv_list` 发现 ready key 并组成 global batch。 + +本文不照搬 AngelSpec 的进程组织。实现依据是旁边只读参考仓库 `verl-SpeCo` 的 PR #48 分支,尤其是: + +```text +verl_speco/integration/transferqueue_bridge.py +verl_speco/integration/sglang_runtime.py +verl_speco/integration/oldlogprob_runtime.py +verl_speco/workers/speco_worker.py +verl_speco/integration/task_runner.py +``` + +参考仓库只用于理解和复用设计,不修改其代码。真正实现仍放在本项目 `verl-SpeCo-ls`,服务于 standalone drafter training。 + +建议按下面顺序阅读: + +1. 第一部分先介绍 verl/SpeCo 中的进程、`TokenOutput`、`DataProto`、WorkerGroup、replica 和 owner rank; +2. 再完整解释 PR #48 的 SGLang TQ 路径; +3. 再解释 PR #48 的 old-logprob TQ 路径; +4. 再解释 hidden 恢复后如何进入 online buffer 和 optimizer step; +5. 第二部分才讨论如何把相同 TQ 传输能力改造成独立 Producer/DSpark Consumer。 + +## 第一部分:PR #48 原始 TQ 流程 + +这一部分只解释只读参考仓库 `../verl-SpeCo` 的 PR #48。这里的 Producer、Consumer、Ray driver、数据结构和生命周期全部是 PR #48 当前代码的行为,不是本文为 standalone 设计的行为。 + +### 1A. 阅读 PR #48 前必须知道的项目对象 + +#### SGLang server + +SGLang server 是 rollout 推理进程。它接收 prompt,执行 target model 推理并生成 response。SpeCo 的 SGLang patch 还会在推理过程中收集指定层 hidden states。 + +它不是 drafter trainer,也不执行 optimizer step。在 PR #48 中它是 hidden-state Producer。 + +#### TokenOutput + +`TokenOutput` 是一次 rollout request 的返回对象。核心字段可以概括为: + +```python +TokenOutput( + token_ids=list[int], + log_probs=..., + routed_experts=..., + extra_fields={ + "global_steps": int, + "drafter_sample": dict | None, + }, +) +``` + +`token_ids` 是生成结果;`extra_fields` 是 SpeCo 添加的旁路字段。`drafter_sample` 不影响正常 response 返回,它用于把草稿训练所需信息从 rollout 侧带回训练控制层。 + +#### DataProto 和 non_tensor_batch + +verl 将一批 rollout 结果整理成 `DataProto`。它通常分为: + +```python +DataProto( + batch=TensorDict(...), + non_tensor_batch={...}, + meta_info={...}, +) +``` + +- `batch`:规则的批量 tensor,例如 prompts、responses、attention mask; +- `non_tensor_batch`:不能直接组成规则 dense batch 的 Python 对象或 object array; +- `meta_info`:批次级配置和指标。 + +每个 request 的 `TokenOutput.extra_fields["drafter_sample"]` 最终会被 rollout/agent-loop 聚合到: + +```python +gen_batch_output.non_tensor_batch["drafter_sample"] +``` + +所以 SGLang 侧写的是单 request `TokenOutput`,driver 侧拿到的是批量 `DataProto`。 + +#### RayPPOTrainer driver + +`SpecoRayPPOTrainer` 所在进程是控制进程,简称 driver。它负责按训练 step 依次调用 rollout、old-logprob、actor update,以及向各 WorkerGroup 发 RPC。 + +driver 不执行 drafter 模型 forward。PR #48 之前,它会接触包含 hidden tensor 的 `drafter_sample`;PR #48 之后,它只处理中小 tensor、标量 metadata 和 TQ key。 + +#### WorkerGroup + +WorkerGroup 是 verl 对一组 Ray worker 的调用封装。代码: + +```python +self.drafter_wg.collect_rollout_features(buckets) +``` + +不是普通本地函数调用,而是根据注册的 dispatch rule,将参数分发到多张卡上的 `SpecoWorker.collect_rollout_features()`。 + +#### Rollout replica + +rollout replica 是一组共同承载一个 rollout 模型副本的进程/GPU。一个任务可能存在多个 data-parallel rollout replica。`replica_rank` 标识当前 sample 是哪个 rollout replica 生成的: + +```text +replica_rank = 0, 1, 2, ... +``` + +#### Drafter training replica、DP rank 和 SP rank + +drafter 训练也可能按 data parallel 和 sequence parallel 组织: + +```text +drafter replica / DP rank 0 + ├─ SP rank 0 + └─ SP rank 1 + +drafter replica / DP rank 1 + ├─ SP rank 0 + └─ SP rank 1 +``` + +同一个 drafter DP replica 内的 SP ranks 共同执行一个模型副本的训练。`replica_rank` 用来把 rollout replica 产生的数据路由到对应 drafter DP replica。 + +#### Owner rank + +`collect_rollout_features` 注册了: + +```python +@register( + dispatch_mode=make_nd_compute_dispatch_fn( + mesh_name="drafter_owner_route" + ) +) +``` + +每个 drafter DP replica 的 `SP rank 0` 被标记为 collect leader/owner。driver 传入的是按 replica 分好的 bucket,dispatch 层负责把 bucket 发到对应训练组。一个 replica 内可能有多个 rank 需要共同训练,但只有指定 leader 负责汇总 RPC 返回。 + +这里的 owner route 是 PR #48 为什么不能简单让任意 worker 随机拿 sample 的原因:sample 必须进入与当前 drafter device mesh 一致的训练组。 + +## 2. PR #48 改造前的 online 特征流程 + +PR #48 改造的是 SpeCo 的 online drafter feature transport。hidden states 有两条主要来源。 + +### 2.1 SGLang rollout hidden 路径 + +改造前: + +```text +SGLang server + → 生成 drafter_sample,其中直接包含 hidden_states CPU tensor + → TokenOutput.extra_fields + → RayPPOTrainer driver 收集 drafter_sample + → driver 按 drafter replica/owner 分桶 + → Ray dispatch / object store + → SpecoWorker.collect_rollout_features(samples) + → _store_rollout_sample() + → online drafter buffer/train +``` + +此时 `drafter_sample` 类似: + +```python +{ + "input_ids": int64[1, total_len], + "prompts": int64[1, prompt_len], + "responses": int64[1, response_len], + "hidden_states": bf16[1, hidden_rows, hidden_dim], + "hidden_positions": int64[1, hidden_rows], + "target_logprobs": tensor | None, + "global_step": 42, + "replica_rank": 1, +} +``` + +问题是整个字典经过 driver,而 `hidden_states` 是其中最大的字段。driver 的 host memory 和 Ray object store 都会承载这些 tensor。 + +### 2.2 old-logprob hook hidden 路径 + +另一条路径在 actor old-logprob forward 中捕获 hidden states。改造前,大 chunk 通过: + +```python +chunk_ref = ray.put(hidden_chunk) +``` + +sample 不直接带 tensor,而是带: + +```python +{ + "hidden_states_ref_chunks": [ + { + "ref": ray_object_ref, + "start": 0, + "length": 512, + }, + ], +} +``` + +drafter worker 再 `ray.get(ref)`,根据 `start/length` 切出每条 sample 所需行。 + +### 2.3 PR #48 要改变的边界 + +PR #48 没有改变: + +- rollout 什么时候产生 sample; +- driver 如何触发 drafter worker; +- drafter worker 如何调用 `_store_rollout_sample()`; +- drafter model 的训练逻辑; +- drafter 权重发布。 + +它只改变大 tensor 的跨进程介质: + +```text +改造前:Producer → Ray driver/object store → Consumer +改造后:Producer → TQ storage → Consumer + key 仍走原 Ray 控制路径 +``` + +## 3. PR #48 改造后的完整 TQ 流程 + +### 3.0 总览 + +PR #48 的目标不是让 TQ 自己产生训练 batch,而是把原来经过 Ray driver/Ray object store 的大 hidden tensor 攁到 TQ。原来的控制路径继续存在,只是控制路径上从“大 tensor”变成“小 key”。 + +整体结构是: + +```text + 原 Ray 控制路径 + drafter_sample / chunk ref +Producer ───────────────────── key ───────────────────▶ Consumer + │ │ + │ kv_put(large tensor) │ kv_batch_get(key) + ▼ ▼ +TransferQueue storage ─────────────────────────────────────┘ +``` + +因此 PR #48 同时保留两条通道: + +```text +控制通道:Producer → Ray driver → drafter worker +数据通道:Producer → TQ storage → drafter worker +``` + +控制通道负责告诉 Consumer “本次应该处理哪个样本、对应哪个 TQ key”;数据通道负责传输 hidden states 等大 tensor。 + +#### 3.0.1 配置放在哪里 + +PR #48 在 drafter training 配置下增加: + +```yaml +actor_rollout_ref: + rollout: + drafter: + training: + transfer_queue: + enable: false + backend: + storage_backend: SimpleStorage + SimpleStorage: + total_storage_size: 100000 + num_data_storage_units: 8 +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/config/speco_base.yaml +``` + +这里的 `transfer_queue` 是 SpeCo drafter 自己的配置,不是 standalone `feature_store.type`,也不是简单把上游 verl 的 `transfer_queue.enable` 打开。 + +#### 3.0.2 TaskRunner 创建整套 TQ + +RL 任务启动时,`SpecoTaskRunner.run()` 在创建 workers 之前调用: + +```python +from verl_speco.integration.transferqueue_bridge import ( + close_transfer_queue, + init_transfer_queue, +) + +transfer_queue_started = init_transfer_queue(config) +try: + trainer.init_workers() + trainer.fit() +finally: + if transfer_queue_started: + close_transfer_queue() +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/integration/task_runner.py:319 +``` + +`init_transfer_queue(config)` 内部读取: + +```python +config.actor_rollout_ref.rollout.drafter.training.transfer_queue +``` + +然后执行: + +```python +tq.init(_to_plain_dict(tq_cfg)) +``` + +并记录: + +```python +_state["initialized"] = True +_state["owner"] = True +``` + +这里的 owner 是“创建/拥有 TQ 生命周期的进程”。只有 owner 在任务结束时执行 `tq.close()`。 + +关键顺序是: + +```text +SpecoTaskRunner +→ tq.init(完整配置) +→ trainer.init_workers() +→ Ray workers 启动 +``` + +也就是说,PR #48 假设 TQ Controller/storage 已经由 TaskRunner 在 Ray 集群环境中建立,后启动的 worker 只需要连接。 + +#### 3.0.3 每个 Producer/Consumer 进程怎么连接 TQ + +bridge 中的 `_ensure_initialized()` 是进程级懒初始化: + +```python +def _ensure_initialized(): + if _state["initialized"]: + return + + with _state_lock: + if _state["initialized"]: + return + + tq.init() + _state["initialized"] = True +``` + +注意这里是: + +```python +tq.init() +``` + +不是: + +```python +tq.init(config) +``` + +无参初始化的含义是连接 TaskRunner 已创建的同一套 TQ。它依赖 PR #48 所处的 Ray 运行环境完成服务发现。 + +因此 PR #48 不是“每个 worker 各创建一套 TQ”,而是: + +```text +TaskRunner:tq.init(config),创建一次 +SGLang producer:tq.init(),连接 +actor producer:tq.init(),连接 +drafter consumer:tq.init(),连接 +``` + +#### 3.0.4 SGLang Producer 怎么写 hidden states + +SGLang 已经完成 rollout,并组装出 `drafter_sample` 后,PR #48 执行: + +```python +configure_transfer_queue(training_cfg) + +if is_transfer_queue_enabled(): + tq_key = make_sample_key( + collection_global_steps, + self.replica_rank, + request_id, + ) + + tq_payload = { + "hidden_states": hidden_states.unsqueeze(0).cpu(), + } + + if target_logprobs is not None: + tq_payload["target_logprobs"] = ( + target_logprobs.unsqueeze(0).cpu() + ) + + put_sample( + tq_key, + tq_payload, + tag={ + "global_step": collection_global_steps, + "replica_rank": self.replica_rank, + }, + ) + + drafter_sample["hidden_states_tq_key"] = tq_key + drafter_sample["hidden_states"] = None +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/integration/sglang_runtime.py:2020 +``` + +这里发生了两条不同的数据流: + +```text +大 tensor:SGLang → TQ +小字典/key:SGLang → 原 Ray/driver 路径 → drafter worker +``` + +写进 TQ 后将: + +```python +drafter_sample["hidden_states"] = None +``` + +是为了避免 hidden states 继续经过 driver/Ray object store。driver 仍收到 sample,但大 tensor 已替换成: + +```python +drafter_sample["hidden_states_tq_key"] +``` + +#### 3.0.5 `put_sample()` 实际怎么写 + +bridge 中: + +```python +def put_sample(key, tensor_dict, *, tag=None): + payload = { + k: v + for k, v in tensor_dict.items() + if torch.is_tensor(v) + } + + _ensure_initialized() + + tq.kv_put( + key=key, + partition_id="speco_drafter_features", + fields=payload, + tag=tag or {}, + ) +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/integration/transferqueue_bridge.py:206 +``` + +这里可以明确看到: + +- PR #48 使用 TQ 高层 KV API; +- 一个 key 对应一个 sample; +- `fields` 是 tensor 字典; +- `tag` 是小 metadata; +- partition 当前写死为 `speco_drafter_features`; +- 写入前 tensor 已 `.cpu()`; +- 写入失败直接抛异常,不静默回退。 + +key 的生成代码是: + +```python +def make_sample_key(global_step, replica_rank, request_id): + return f"speco:{global_step}:{replica_rank}:{request_id}" +``` + +这个 key 对 RL rollout 是合理的,因为它用 step、rollout replica 和 request ID 标识一次在线采样。 + +#### 3.0.5.1 PR #48 中“数据”和“元数据”分别长什么样 + +PR #48 一条 SGLang 样本在写 TQ 之前,`drafter_sample` 同时包含大 tensor 和控制信息。简化后类似: + +```python +drafter_sample = { + # 普通训练输入,仍走原 sample/Ray 控制路径 + "input_ids": int64[1, prompt_len + response_len], + "prompts": int64[1, prompt_len], + "responses": int64[1, response_len], + + # 大 tensor,开启 TQ 后从这个字典移除 + "hidden_states": bf16[1, hidden_rows, hidden_dim], + "target_logprobs": fp32[1, rows, topk_or_vocab] | None, + + # hidden 与 token 对齐所需的小字段 + "hidden_positions": int64[1, hidden_rows] | None, + "hidden_position_start": int, + "hidden_position_end": int, + "hidden_window_start": int, + "hidden_window_end": int, + + # 控制信息 + "global_step": int, + "replica_rank": int, +} +``` + +执行 `put_sample()` 时,并不是把整个 `drafter_sample` 放进 TQ。PR #48 只抽取占用大的 tensor: + +```python +tq_payload = { + "hidden_states": bf16[1, hidden_rows, hidden_dim], + "target_logprobs": fp32[1, rows, ...], # 可选 + "hidden_raw_target_logprobs": ..., # 可选 + "hidden_raw_target_logprobs_positions": ..., # 可选 +} +``` + +这就是 TQ 的 data payload。它被传给: + +```python +tq.kv_put(fields=tq_payload) +``` + +另外还有 TQ tag: + +```python +tag = { + "global_step": 42, + "replica_rank": 1, +} +``` + +tag 是 TQ 侧轻量 metadata,用于描述/检索对象,不承载 hidden tensor。 + +写入完成后,仍经 Ray 传递的轻量 `drafter_sample` 变成: + +```python +drafter_sample = { + "input_ids": int64[1, total_len], + "prompts": int64[1, prompt_len], + "responses": int64[1, response_len], + + "hidden_states": None, + "target_logprobs": None, + "hidden_states_tq_key": "speco:42:1:req-007", + + "hidden_positions": int64[1, hidden_rows] | None, + "hidden_position_start": 128, + "hidden_position_end": 640, + "global_step": 42, + "replica_rank": 1, +} +``` + +因此 PR #48 实际存在三类对象: + +| 对象 | 内容 | 传输路径 | 作用 | +|---|---|---|---| +| TQ fields/payload | `hidden_states` 等大 tensor | Producer → TQ storage → Consumer | 避免 driver 搬运大 tensor | +| TQ tag | `global_step`、`replica_rank` | TQ control/index metadata | 描述该 key | +| 轻量 `drafter_sample` | tokens、位置、标量 metadata、`hidden_states_tq_key` | 原 Ray driver 路径 | 告诉 Consumer 取哪个 key,以及如何解释取回的 tensor | + +代码实现解耦的关键不是“所有内容都进 TQ”,而是: + +```python +drafter_sample["hidden_states_tq_key"] = tq_key +drafter_sample["hidden_states"] = None +``` + +第一行给 Consumer 留下寻址信息;第二行阻止大 tensor 继续沿旧路径传输。 + +#### 3.0.5.2 Consumer 如何把两部分重新合成一个训练样本 + +Consumer 最初拿到的是轻量 sample: + +```python +sample["hidden_states"] is None +sample["hidden_states_tq_key"] == "speco:42:1:req-007" +``` + +它执行: + +```python +payload = get_sample(sample["hidden_states_tq_key"]) +sample["hidden_states"] = payload["hidden_states"] +``` + +合并后: + +```python +sample = { + "input_ids": ..., + "prompts": ..., + "responses": ..., + "hidden_positions": ..., + "hidden_states": bf16[1, hidden_rows, hidden_dim], + "hidden_states_tq_key": "speco:42:1:req-007", + ... +} +``` + +后面的 `_store_rollout_sample()` 看到的结构与关闭 TQ 时基本一致,所以训练主体不需要增加 TQ 分支。TQ bridge 只改变“大 tensor 从哪里恢复”,不改变 drafter trainer 的输入语义。 + +#### 3.0.6 old-logprob Producer 怎么写 chunk + +PR #48 的另一条 Producer 路径来自 actor old-logprob hidden hook。原来是: + +```python +chunk_ref = ray.put(hidden_chunk) +``` + +开启 TQ 后改成: + +```python +tq_key = make_sample_key( + global_step, + owner, + f"chunk{len(chunk_refs)}", +) + +put_sample( + tq_key, + {"hidden": hidden_chunk}, + tag={ + "global_step": global_step, + "owner": owner, + }, +) + +chunk_ref = tq_key +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/integration/oldlogprob_runtime.py:533 +``` + +后面的 driver 不需要知道 ref 是 Ray ObjectRef 还是 TQ string key,它只把 ref 当不透明 token 继续传递。 + +#### 3.0.7 Consumer 怎么根据 key 读取 + +drafter worker 收到原来的 sample 小字典后: + +```python +tq_key = sample.get("hidden_states_tq_key") + +if tq_key is not None and self._speco_tq_enabled: + payload = get_sample(tq_key) + + for field in ( + "hidden_states", + "target_logprobs", + "hidden_raw_target_logprobs", + "hidden_raw_target_logprobs_positions", + ): + if payload.get(field) is not None: + sample[field] = payload[field] + + if sample.get("hidden_states") is None: + raise RuntimeError( + "TQ key exists but hidden_states payload is missing" + ) +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/workers/speco_worker.py:848 +``` + +恢复 tensor 后,后面的逻辑仍使用原 `sample`/`batch`,drafter trainer 不需要知道 tensor 来自 Ray 还是 TQ。 + +#### 3.0.8 `get_sample()` 实际怎么读 + +```python +def get_sample(key): + _ensure_initialized() + + result = tq.kv_batch_get( + keys=[key], + partition_id="speco_drafter_features", + ) + + value = _extract_value(result, key) + return _tensordict_to_dict(value) +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/integration/transferqueue_bridge.py:237 +``` + +`_extract_value()` 兼容三种返回形态: + +```python +if isinstance(result, dict): + return result.get(key) +if isinstance(result, (list, tuple)): + return result[0] +return result +``` + +这是因为不同 TQ 版本/后端返回包装可能不同。 + +#### 3.0.9 为什么需要 `_densify_tq_tensor()` + +PR #48 后续修复发现,TQ 把单样本 tensor 放进 TensorDict 后,`kv_batch_get` 可能返回 NestedTensor,并额外带 batch 维。旧代码要执行: + +```python +tensor[start:start + length] +``` + +但 NestedTensor 不支持在 jagged dim 直接 slice。因此加入: + +```python +def _densify_tq_tensor(tensor): + if tensor.is_nested: + parts = [ + part + for part in tensor.unbind() + if part.numel() > 0 + ] + tensor = torch.cat(parts, dim=0) + + if tensor.dim() == 3: + tensor = tensor.squeeze(0) + elif tensor.dim() == 1: + tensor = tensor.unsqueeze(0) + + return tensor.contiguous() +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/workers/speco_worker.py:72 +``` + +standalone Consumer 同样必须做这个转换,不能假定 `kv_batch_get` 返回普通 dense `[seq, hidden]`。 + +#### 3.0.10 为什么需要 per-step cache + +old-logprob 路径里,一个 owner hidden chunk 可能被约 16 个 sample 共同引用。如果每个 sample 都: + +```python +get_sample(same_tq_key) +``` + +就会重复传输同一个数百 MB chunk。PR #48 在每次 `collect_rollout_features()` 开始时创建: + +```python +self._tq_chunk_cache = {} +``` + +解析 ref 时: + +```python +cache_key = ref if isinstance(ref, str) else id(ref) + +if cache_key not in cache: + cache[cache_key] = _resolve_tq_or_ray_ref(ref) + +tensor = cache[cache_key] +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/workers/speco_worker.py:98 +../verl-SpeCo/verl_speco/workers/speco_worker.py:854 +``` + +独立训练如果每个 sample 都是独立 TQ key,主要使用 `kv_batch_get(keys=[...])` 批量读取,不一定需要跨 sample chunk cache;但 prefetch 重试或共享 packed object 时仍应保留 key cache。 + +#### 3.0.11 PR #48 什么时候删除数据 + +PR #48 没有在 `get_sample()` 后删除。其注释明确说明:同一 drafter replica 的多个 TP/SP rank 可能读取同一 key,第一次读取后立刻删除会让剩余 rank 失败。 + +当前策略是任务结束时由 owner: + +```python +tq.close() +``` + +统一结束 TQ 生命周期。也就是说,PR #48 当前没有实现精细的逐 step `kv_clear`。 + +#### 3.0.12 PR #48 的完整时序 + +```text +SpecoTaskRunner + → tq.init(config) + → 启动 Ray workers + +SGLang/actor Producer process + → configure_transfer_queue() + → 第一次 put 时 tq.init() + → kv_put(key, tensor fields, tag) + → 把 key 塞回原 sample/ref + +Ray driver + → 只中转小 sample/key + +drafter worker Consumer process + → 第一次 get 时 tq.init() + → kv_batch_get([key]) + → 解包 TensorDict/NestedTensor + → 恢复 sample["hidden_states"] + → 原 drafter collect/train 逻辑 + +任务结束 + → TaskRunner owner tq.close() +``` + +### 3.1 已经实现的可复用能力 + +PR #48 新增 `verl_speco/integration/transferqueue_bridge.py`,锁定思路是把 TQ 当成独立传输库,不修改上游 verl。它提供: + +```python +configure_transfer_queue(training_cfg) +init_transfer_queue(config) +make_sample_key(global_step, replica_rank, request_id) +put_sample(key, tensor_dict, tag=...) +get_sample(key) +close_transfer_queue() +``` + +实际写入调用是: + +```python +tq.kv_put( + key=key, + partition_id="speco_drafter_features", + fields=payload, + tag=tag, +) +``` + +实际读取调用是: + +```python +result = tq.kv_batch_get( + keys=[key], + partition_id="speco_drafter_features", +) +``` + +另外,PR #48 已经处理了多项 standalone 方案也需要的问题: + +1. TQ 返回值可能是 direct value、`{key: value}` 或 list,需要统一解包; +2. TQ 返回的 tensor 可能是 NestedTensor,必须通过 `unbind + cat` 恢复为生产端写入的 dense tensor; +3. 多个样本引用同一 hidden chunk 时,需要 per-step cache,避免重复 `kv_batch_get` 同一个大对象; +4. TQ key 存在但 payload 缺失时 fail loud,不能静默丢样本; +5. `enable=false` 时保留原传输路径。 + +这些逻辑应直接作为本项目 TQ adapter 的参考。 + +### 3.2 PR #48 的数据流 + +PR #48 优化的是 RL online 路径: + +```text +SGLang/actor worker + → kv_put(hidden states) + → 把 hidden_states_tq_key 塞进原 drafter_sample + → 原 Ray driver 继续传递小 sample/key + → drafter worker collect_rollout_features() + → kv_batch_get(key) +``` + +它没有让 consumer 自己从 TQ 发现下一批 key;key 仍沿原来的 Ray 控制路径到达 drafter worker。 + +### 3.3 PR #48 没有提供的 standalone 能力 + +PR #48 当前没有实现: + +- 从预生成 response 文件读取数据的独立 Producer; +- Producer 并行请求外部 vLLM endpoint; +- standalone DSpark trainer 主动发现 ready key; +- global batch 到各 torchrun rank 的分片; +- 每个 optimizer step 后精确 `kv_clear`; +- EOS; +- standalone 无 Ray 的 TQ bootstrap; +- MooncakeStore 的实际运行验证。 + +PR #48 当前配置是: + +```yaml +transfer_queue: + enable: false + backend: + storage_backend: SimpleStorage + SimpleStorage: + total_storage_size: 100000 + num_data_storage_units: 8 +``` + +并且 bridge 注释针对 `TransferQueue==0.1.7`。旁边当前 verl 主线已使用 `0.1.8` 文案并包含 `MooncakeStore` 配置。当前机器没有安装 `transfer_queue` 包,因此正式实现前必须锁定版本并实机验证 API 签名,不能把 0.1.7 和 0.1.8 混用。 + +### 3.4 standalone 方案对 PR #48 的扩展 + +不再自行实现 `publish_ready/claim/lease/ack` HTTP 服务。第一版在 PR #48 KV 模式上补四个操作: + +```python +tq.kv_batch_put(...) # Producer 批量写 +tq.kv_list(...) # rank 0 列出 key + tag +tq.kv_batch_get(...) # 各 rank 并行读 +tq.kv_clear(...) # optimizer step 成功后删 +``` + +第一版不依赖 TQ `Sampler/StreamingDataLoader/get_meta`,因为 PR #48 并未使用或验证这些接口。等 KV 流程稳定后再升级为 TQ StreamingDataLoader。 + +### 3.5 PR #48 与 standalone 独立训练逐项映射 + +| PR #48 online RL | standalone drafter training | +|---|---| +| `SpecoTaskRunner` 调用 `tq.init(config)` | 新增独立 TQ owner/bootstrap 进程,或在确认安全后由 Producer owner 调用 `tq.init(config)` | +| SGLang/actor worker 是 Producer | `feature_producer.py` 是独立 Producer | +| rollout 过程中已经得到 hidden states | Producer 读取预生成 response,再请求外部 vLLM prefill | +| `make_sample_key(global_step, replica_rank, request_id)` | `make_replay_sample_key(dataset,row,tokens,target fingerprint)` | +| `put_sample()` / `kv_put()` | 继续复用同一写入模式,可扩展为 `kv_batch_put()` | +| key 通过 Ray driver/sample 传给 consumer | 没有 driver;rank 0 用 `kv_list()` 主动发现 ready keys | +| drafter worker `get_sample(key)` | 每个 torchrun rank `kv_batch_get(local_keys)` | +| `collect_rollout_features()` 恢复 sample | 转换为 `DraftFeatureSample` 后调用现有 `prepare_training_batch_from_samples()` | +| 多 TP/SP rank 可能读同一 key,所以不立即删除 | data-parallel rank 读取互不重叠的 key;全 rank step 成功后统一 `kv_clear(global_keys)` | +| TaskRunner 结束时 `tq.close()` | 每 step clear;输入 drain 完成后 owner 最后 `tq.close()` | + +standalone 需要新增的控制流是: + +```text +Producer DSpark rank 0 其他 ranks + │ │ │ + │ kv_put(sample key, fields, tag) │ │ + ├────────────────────────────────────▶│ │ + │ │ kv_list READY keys │ + │ │ │ + │ │ broadcast selected_keys ──▶│ + │ │ │ + │ │ kv_batch_get(local keys) │ kv_batch_get(local keys) + │ │ │ + │ ├──── DSpark synchronized step ────┤ + │ │ │ + │ │ kv_clear(global keys) │ +``` + +这个映射中,TQ 同时承担: + +- 大 tensor 存储/传输; +- key、tag 和 partition 的轻量索引。 + +但第一版 global batch 的选择仍由单个 DSpark job 的 rank 0 完成。这样最接近 PR #48 的 KV API,避免在同一次改造中再引入未经该 PR 验证的 Sampler/StreamingDataLoader。 + +### 3.6 PR #48 的 SGLang 路径:逐函数、逐对象完整流程 + +下面从一次生成请求开始,不省略中间层。 + +#### 阶段 1:SGLang完成生成并收集 hidden states + +执行进程:SGLang rollout server。 + +输入是一次 request 对应的 prompt 和生成配置。生成结束时,代码已经持有: + +```python +prompt_tensor: int64[prompt_len] +response_tensor: int64[response_len] +hidden_states: bf16[hidden_rows, hidden_dim] +hidden_positions: int64[hidden_rows] | None +target_logprobs: tensor | None +request_id: str +collection_global_steps: int +self.replica_rank: int +``` + +这些变量的语义: + +- `prompt_tensor`:输入 prompt token IDs; +- `response_tensor`:SGLang生成的 response token IDs; +- `hidden_states`:target model 指定层在部分 token positions 上的输出; +- `hidden_positions`:每个 hidden row 对应完整 `prompt+response` 序列中的哪个 token position; +- `target_logprobs`:可选的目标概率监督; +- `request_id`:当前 rollout request 标识; +- `replica_rank`:执行该 request 的 rollout replica。 + +SGLang 先构造完整 sample: + +```python +drafter_sample = { + "input_ids": torch.cat( + [prompt_tensor, response_tensor], dim=0 + ).unsqueeze(0), + "prompts": prompt_tensor.unsqueeze(0), + "responses": response_tensor.unsqueeze(0), + "hidden_states": hidden_states.unsqueeze(0).cpu(), + "hidden_positions": hidden_positions.unsqueeze(0).cpu(), + "target_logprobs": ( + target_logprobs.unsqueeze(0).cpu() + if target_logprobs is not None + else None + ), + "global_step": collection_global_steps, + "replica_rank": self.replica_rank, + # 还有 hidden window/alignment metadata +} +``` + +前导维 `1` 表示这是一个单样本 batch。`.cpu()` 表示跨进程传输前把大 tensor 放到 CPU 内存。 + +#### 阶段 2:PR #48 将大 fields 写入 TQ + +同一个 SGLang进程执行: + +```python +tq_key = make_sample_key( + collection_global_steps, + self.replica_rank, + request_id, +) + +tq_payload = { + "hidden_states": hidden_states.unsqueeze(0).cpu(), +} + +put_sample( + tq_key, + tq_payload, + tag={ + "global_step": collection_global_steps, + "replica_rank": self.replica_rank, + }, +) +``` + +调用展开后是: + +```python +tq.init() # 当前进程第一次使用时 +tq.kv_put( + key=tq_key, + partition_id="speco_drafter_features", + fields=tq_payload, + tag=tag, +) +``` + +效果是 TQ 中增加一行: + +```text +partition = speco_drafter_features +key = speco:42:1:req-007 +fields = {hidden_states: bf16[1, H, D], ...} +tag = {global_step: 42, replica_rank: 1} +``` + +`kv_put` 返回后,SGLang侧将旧 sample 改成: + +```python +drafter_sample["hidden_states_tq_key"] = tq_key +drafter_sample["hidden_states"] = None +``` + +此后该 Python 字典不再携带 hidden tensor,只携带定位它的 key。 + +#### 阶段 3:把 drafter_sample 放进 TokenOutput.extra_fields + +SGLang返回: + +```python +TokenOutput( + token_ids=token_ids, + log_probs=log_probs, + routed_experts=routed_experts, + extra_fields={ + "global_steps": collection_global_steps, + "drafter_sample": drafter_sample, + }, +) +``` + +此时 `TokenOutput` 中有两类输出: + +- 正常 rollout 输出:`token_ids/log_probs`; +- SpeCo 训练旁路输出:`extra_fields.drafter_sample`。 + +TQ payload 不在 `TokenOutput` 中,只有 `hidden_states_tq_key` 在其中。 + +#### 阶段 4:多个 TokenOutput 聚合成 gen_batch_output + +rollout/agent-loop 层将多个 request 的输出合并为批量 `DataProto`: + +```python +gen_batch_output.non_tensor_batch["drafter_sample"] +``` + +可能是 object array: + +```python +array([ + {"hidden_states_tq_key": "speco:42:0:req-A", ...}, + {"hidden_states_tq_key": "speco:42:1:req-B", ...}, +], dtype=object) +``` + +之所以进入 `non_tensor_batch`,是因为每条 sample 的 hidden window、Python metadata 和可选字段不一定具有统一 dense shape。 + +#### 阶段 5:driver 从 DataProto 取出 drafter samples + +`generate_sequences_with_speco()` 包装原 rollout 调用: + +```python +gen_batch_output = original_generate_sequences(...) +collected = self._speco_collect_generation_samples(gen_batch_output) +``` + +`_speco_collect_generation_samples()` 调用: + +```python +samples = pop_drafter_samples(gen_batch_output) +``` + +`pop_drafter_samples()` 实际执行: + +```python +non_tensor_batch = gen_batch_output.non_tensor_batch +samples_array = non_tensor_batch.pop("drafter_sample", None) +samples = normalize_drafter_samples(samples_array) +``` + +这里 `pop` 有两个作用: + +1. 取得 SpeCo drafter side-channel samples; +2. 从正常 PPO 的 `gen_batch_output` 中移除该旁路字段,避免后续 PPO batch 继续携带它。 + +`normalize_drafter_samples()` 将 dict、object array 或 list 统一成: + +```python +samples: list[dict] +``` + +#### 阶段 6:driver 按 replica_rank 分桶 + +假设有两个 rollout/drafter replicas,收到: + +```python +samples = [ + {"replica_rank": 1, "hidden_states_tq_key": "k1", ...}, + {"replica_rank": 0, "hidden_states_tq_key": "k2", ...}, + {"replica_rank": 1, "hidden_states_tq_key": "k3", ...}, +] +``` + +执行: + +```python +buckets = bucket_drafter_samples_by_replica( + samples, + num_replicas=2, +) +``` + +结果: + +```python +buckets = [ + [sample_k2], # bucket 0 + [sample_k1, sample_k3], # bucket 1 +] +``` + +分桶依据只有: + +```python +owner_rank = int(sample["replica_rank"]) +buckets[owner_rank].append(sample) +``` + +这一步没有读取 TQ,也没有处理 hidden tensor;只对小字典做路由。 + +#### 阶段 7:driver 通过 WorkerGroup RPC 分发 buckets + +driver 调用: + +```python +self._speco_set_drafter_global_step() +self._speco_collect_rollout_features_rpc( + "rollout", + buckets, +) +``` + +RPC 内部调用: + +```python +self.drafter_wg.collect_rollout_features(buckets) +``` + +因为 worker 方法注册了 `drafter_owner_route` dispatch,WorkerGroup 将 `buckets[0]` 发给 drafter DP replica 0,将 `buckets[1]` 发给 drafter DP replica 1。一个训练 replica 内的 SP ranks 根据 mesh dispatch 规则参与对应调用。 + +这里传输的对象仍是: + +```python +list[dict] +``` + +其中包含 tokens、position metadata 和 TQ key,不包含被置空的 hidden states。 + +#### 阶段 8:SpecoWorker 根据 key 从 TQ 恢复 tensor + +目标 worker 执行: + +```python +def collect_rollout_features(self, samples): + for sample in samples: + tq_key = sample.get("hidden_states_tq_key") + payload = get_sample(tq_key) + sample["hidden_states"] = payload["hidden_states"] +``` + +`get_sample()` 展开为: + +```python +tq.init() # 此 Consumer 进程第一次使用时 +result = tq.kv_batch_get( + keys=[tq_key], + partition_id="speco_drafter_features", +) +payload = _extract_value(result, tq_key) +payload = _tensordict_to_dict(payload) +``` + +现在 `sample` 再次包含: + +```python +{ + "input_ids": ..., + "hidden_positions": ..., + "hidden_states": bf16[1, hidden_rows, hidden_dim], + "hidden_states_tq_key": "...", +} +``` + +这与关闭 TQ 时 worker 收到的逻辑内容一致。 + +#### 阶段 9:worker 构造 DrafterBaseTrainer 所需 batch dict + +worker 先保留 token fields: + +```python +batch = { + "input_ids": sample["input_ids"], + "prompts": sample["prompts"], + "responses": sample["responses"], +} +``` + +再复制 hidden alignment metadata,例如: + +```python +batch["hidden_positions"] +batch["hidden_position_start"] +batch["hidden_position_end"] +batch["hidden_states_layout"] +batch["global_step"] +``` + +hidden tensor 单独作为参数: + +```python +self._store_rollout_sample( + batch=batch, + hidden_states=hidden, + target_logprobs=target_logprobs, +) +``` + +#### 阶段 10:样本进入在线 buffer 或落盘 + +`_store_rollout_sample()` 根据 training mode 分支: + +```python +if mode == "collect_only": + self._write_rollout_feature_sample( + batch, + hidden_states, + target_logprobs, + ) +else: + self.trainer.collect_online_data( + batch, + hidden_states, + target_logprobs, + ) +``` + +`collect_only` 会转换成 `DraftFeatureSample` 并写 `TorchShardFeatureStore`。online 模式则进入 `DrafterBaseTrainer.collect_online_data()`。 + +`collect_online_data()` 做: + +1. 将 `input_ids/hidden_states/positions/logprobs` 规范化到 CPU; +2. 按 batch 维拆成逐样本; +3. 根据 `hidden_positions` 校验 hidden row 与 token position; +4. 截取可训练窗口; +5. 构造内部 training item; +6. 保存到当前 step 的 `collected_data`,或在启用 data buffer 时保存到跨 step buffer。 + +因此 TQ get 完成不代表立即 optimizer step。它先恢复 online training sample,再进入现有数据准备逻辑。 + +#### 阶段 11:driver 在 actor update 周期触发 drafter 训练 + +driver 包装了 `update_actor()`: + +```python +should_train_drafter = ( + self._speco_should_attempt_drafter_train_this_step() +) + +actor_output = original_update_actor(...) + +if should_train_drafter: + drafter_trained, metrics = self._speco_train_drafter() +``` + +`_speco_train_drafter()` 再向 WorkerGroup 发: + +```python +self.drafter_wg.train_drafter() +``` + +每个 `SpecoWorker.train_drafter()`: + +1. 检查是否属于 drafter training group; +2. 检查 `training_interval_steps`; +3. 激活 drafter training model; +4. 循环 `train_steps_per_trigger` 次; +5. 每次调用 `self.trainer.training_step(global_step)`; +6. 成功时准备需要发布的 drafter state dict; +7. 清理训练期间临时状态。 + +`training_step()` 从刚才的 online `collected_data/DataBuffer` 组成 batch,执行 drafter forward、loss、backward 和 optimizer step。 + +所以 SGLang TQ 路径的最终效果是: + +```text +TQ 只替换 hidden tensor 跨进程传输 +→ sample 收集逻辑不变 +→ online buffer 不变 +→ drafter training trigger 不变 +→ loss/optimizer 不变 +``` + +### 3.7 PR #48 old-logprob 路径的完整差异 + +old-logprob 路径没有 `TokenOutput.extra_fields`。它从 PPO 的 actor old-logprob forward 开始。 + +#### 阶段 1:driver 构造 collect plan + +driver 根据 batch、collect interval 和 drafter owner 数量决定: + +```python +collect_mask: bool[batch] +hidden_positions: list/tensor per sample +owner_rank: int64[batch] +prompt_lens: int64[batch] +response_lens: int64[batch] +``` + +并把 `global_step` 等控制字段放入 old-logprob micro-batch。 + +#### 阶段 2:actor forward hook 选择 hidden rows + +actor worker 在 old-logprob forward 中捕获指定层 hidden states,根据 `collect_mask/position_mask` 只保留需要训练的样本和 token rows。 + +输出可以是 dense selected tensor,也可以是 sparse rows。随后 `_put_oldlogprob_hidden_refs()` 将同一个 owner 的多条 sample rows 拼成一个较大的 `hidden_chunk`。 + +#### 阶段 3:hidden chunk 写入 TQ + +改造前: + +```python +chunk_ref = ray.put(hidden_chunk) +``` + +PR #48: + +```python +tq_key = make_sample_key( + global_step, + owner, + f"chunk{chunk_index}", +) + +put_sample( + tq_key, + {"hidden": hidden_chunk}, + tag={"global_step": global_step, "owner": owner}, +) + +chunk_ref = tq_key +``` + +TQ fields: + +```python +{"hidden": bf16[total_owner_rows, hidden_dim]} +``` + +控制路径中的 chunk metadata: + +```python +chunk_info = { + "sample_indices": [0, 3, 5], + "starts": [0, 128, 384], + "lengths": [128, 256, 96], + "row_indices": [...], + "dtype": "bfloat16", + "shape": [480, hidden_dim], +} +``` + +`starts/lengths` 描述每条 sample 在共享 chunk 中对应的行区间。 + +#### 阶段 4:driver 将 chunk ref 还原为逐样本引用 + +driver 的 `_speco_collect_oldlogprob_features()` 读取: + +```python +chunk_refs = ["speco:42:0:chunk0", ...] +chunk_meta = [chunk_info, ...] +``` + +然后为每个 batch sample 构造: + +```python +sample["hidden_states_ref_chunks"] = [ + { + "ref": "speco:42:0:chunk0", + "chunk_start": 128, + "chunk_length": 256, + "chunk_row_indices": ..., + "dtype": "bfloat16", + "shape": [480, hidden_dim], + } +] +``` + +同时构造该 sample 的: + +```python +input_ids +prompts +responses +hidden_positions +hidden_states_layout +replica_rank=owner +``` + +再按 owner 放入 `buckets[owner]`,通过同一个 `collect_rollout_features()` RPC 发给 drafter worker。 + +#### 阶段 5:Consumer 获取共享 chunk 并切片 + +drafter worker 发现: + +```python +sample.get("hidden_states") is None +sample.get("hidden_states_ref_chunks") is not None +``` + +于是调用 `_resolve_hidden_state_chunks()`。对字符串 ref: + +```python +if ref.startswith("speco:"): + full_chunk = get_sample(ref)["hidden"] + full_chunk = _densify_tq_tensor(full_chunk) +``` + +然后按 sample metadata 取行: + +```python +sample_hidden = full_chunk[ + chunk_start : chunk_start + chunk_length +] +``` + +同一个 chunk 被多个 sample 复用,所以使用: + +```python +self._tq_chunk_cache[ref] = full_chunk +``` + +保证一次 `collect_rollout_features()` 中同一个 TQ key 只 get 一次。 + +得到逐样本 hidden 后,后续 `_store_rollout_sample → collect_online_data → train_drafter` 与 SGLang 路径相同。 + +### 3.8 PR #48 数据生命周期和清理 + +PR #48 的 TQ row 生命周期是: + +```text +TaskRunner tq.init(config) +→ Producer kv_put +→ key 经 Ray 控制路径传递 +→ 一个或多个 drafter TP/SP rank kv_batch_get +→ online drafter 收集/训练继续执行 +→ 整个 trainer.fit() 结束 +→ TaskRunner finally 调用 tq.close() +``` + +当前没有: + +```python +tq.kv_clear(key) +``` + +原因是同一 sample/key 可能被一个 drafter replica 的多个 rank 读取。若第一个 rank get 后删除,其余 rank 可能 get 失败。 + +因此 PR #48 采用任务级生命周期,而不是样本级确认和回收。这简化了并发正确性,但意味着长任务中的 TQ storage 占用可能持续增长;代码注释也把“leader 在 barrier 后精细 clear”留作后续工作。 + +### 3.9 PR #48 开启与关闭时的行为差异 + +`configure_transfer_queue()` 返回: + +```python +enabled_in_config and transfer_queue_importable +``` + +关闭时: + +```text +SGLang drafter_sample 继续内联 hidden_states +old-logprob 继续 ray.put(hidden_chunk) +Consumer 继续 ray.get/ref resolve +``` + +开启时: + +```text +SGLang hidden fields → TQ,sample 只带 key +old-logprob hidden chunk → TQ,ref 变成字符串 key +Consumer 根据 key 类型走 TQ get +``` + +如果配置开启但 `transfer_queue` 包未安装,bridge 会记录 warning,并让 `is_transfer_queue_enabled()` 返回 false,保留旧 Ray 路径。若已经进入 `put_sample/get_sample` 却发生 TQ 错误,则抛异常,不静默丢 hidden data。 + +## 第二部分:基于 PR #48 的 standalone drafter training 适配 + +从这一部分开始才讨论 `verl-SpeCo-ls` 的独立训练。下面的 `kv_list`、rank 0 选 global keys、逐 step `kv_clear`、独立 vLLM Producer 都是需要在本项目新增的逻辑,不是 PR #48 已有逻辑。 + +### 当前 standalone 基线 + +当前独立训练是: + +```text +draft_train_launcher +→ torch.distributed.run +→ 每个 rank 创建 DraftFeatureDataLoader +→ 每个 rank 在训练循环内同步 TargetFeatureReplayer.materialize() +→ vLLM/file hidden payload +→ prepare_training_batch_from_samples() +→ training_step_from_batch() +``` + +新方案要把 `TargetFeatureReplayer.materialize()` 从训练 rank 的同步取数路径移到独立 Producer,同时保持后两步训练接口不变。 + +## 4. TQ metadata 到底记录什么 + +### 4.1 Partition + +一次训练运行使用一个独立 partition: + +```python +partition_id = f"speco:{run_id}:dspark_train" +``` + +partition 用来隔离: + +- 不同训练 run; +- train 和 validation; +- 不同 target checkpoint 生成的 hidden states。 + +不能让两个 target 模型共用同一 partition,否则训练侧可能消费错误的 hidden states。 + +### 4.2 Sample key + +每条输入样本使用稳定 key: + +```python +sample_key = sha256( + dataset_id + + row_id + + prompt_token_ids + + response_token_ids + + tokenizer_fingerprint + + target_model_fingerprint + + target_layer_ids + + hidden_states_layout +).hexdigest() +``` + +稳定 key 用于: + +- vLLM HTTP 请求重试时不生成不同对象; +- Producer 重启后识别相同样本; +- 检查 hidden states 是否属于正确模型和正确层; +- TQ/Mooncake 清理时准确定位对象。 + +### 4.3 Fields 与 READY 约定 + +每个样本包含固定字段: + +```python +{ + "input_ids": int64[seq], + "loss_mask": float32[seq], + "position_ids": int64[seq], + "hidden_states": bf16[seq, aux_hidden_dim or aux_hidden_dim+hidden_size], +} +``` + +这里必须保持当前 `DraftFeatureSample` 的契约:`TargetFeatureReplayer._feature_from_vllm_payload()` 会把多个 aux layers flatten;当 layout 是 `dflash_aux_plus_last` 时,还会把 final hidden 拼到同一个 `hidden_states` tensor 尾部。`DSparkTrainerBackend.preprocess_individual_items()` 再根据 metadata 中的 `hidden_states_layout` 将 final hidden 切出来,生成训练 batch 的 `target_last_hidden_states`。 + +因此 TQ 不需要新增一个独立 `target_last_hidden_states` field。开启 DSpark L1 时,要求: + +```text +metadata.hidden_states_layout = dflash_aux_plus_last +hidden_states.shape[-1] = num_context_layers * hidden_size + hidden_size +``` + +完整 TQ native metadata API 可以做字段级 ready 判定,但 PR #48 当前走的是高层 KV API:一次 `kv_put` 把一个样本的多个 tensor fields 一起写入。因此第一版采用更直接的约定: + +```python +required_fields = [ + "input_ids", + "loss_mask", + "position_ids", + "hidden_states", +] +``` + +Producer 先在内存中验证所有必需字段,再进行一次 `kv_put`。只有 `kv_put` 成功返回,key 才会出现在 `kv_list` 结果中,并带有: + +```python +tag={ + "status": "ready", + "run_id": run_id, + "sample_id": sample_key, +} +``` + +Consumer 只选择 `status=ready` 且 `run_id` 匹配的 key。不要先 put `input_ids`、再单独 put `hidden_states`,否则 Consumer 可能观察到半成品。 + +### 4.4 Tags + +tags 是轻量 metadata,不放大 tensor: + +```python +tags = { + "sample_id": sample_key, + "source_row": row_id, + "seq_len": seq_len, + "payload_bytes": payload_bytes, + "target_model_fp": target_model_fingerprint, + "target_layers": "8,16,24", + "hidden_layout": "dflash_aux_plus_last", + "producer_status": "success", +} +``` + +tags 可用于过滤、监控、背压统计和错误排查。它不能替代 tensor shape/dtype 校验。 + +### 4.5 Run ID,而不是先依赖 task_name + +PR #48 的 `kv_put/kv_batch_get` 路径没有使用 `task_name` 或 Sampler consumption history。第一版按它的已验证接口,在 partition 和 tags 中放 `run_id`: + +```python +partition_id = "speco_drafter_features" +tag = { + "run_id": run_id, + "status": "ready", +} +``` + +不同 run 最好直接使用不同 partition: + +```python +partition_id = f"speco_drafter_features_{run_id}" +``` + +这样 job 重启和清理更简单。`task_name=dspark_train` 留到后续迁移 TQ native metadata/Sampler 时再使用。 + +### 4.6 standalone 中一条样本的完整对象形态 + +standalone 没有 PR #48 的轻量 `drafter_sample → Ray driver` 路径,因此需要让 TQ tag 承担“如何找到和解释 payload”的 metadata 作用。 + +#### Producer 读到的原始记录 + +```python +source_record = { + "dataset_id": "math-train", + "row_id": 12345, + "prompt": "...", + "response": "已经提前生成的 response", +} +``` + +#### Token replay 样本 + +分词和对齐后: + +```python +replay_sample = DraftReplaySample( + input_ids=int64[full_seq], + loss_mask=float32[full_seq], + position_ids=int64[full_seq], + feature_positions=int64[feature_rows], + draft_position_ids=int64[feature_rows], + metadata={ + "dataset_id": "math-train", + "row_id": 12345, + }, +) +``` + +这里的 `full_seq` 是 prompt 与预生成 response 拼接后的长度;`feature_positions` 指明哪些 token 位置最终进入 drafter 监督样本。 + +#### vLLM 返回的原始 hidden payload + +当前文件协议要求 safetensors 至少包含: + +```python +vllm_payload = { + "token_ids": int64[prefill_rows], + "hidden_states": bf16[prefill_rows, returned_layers, hidden_size], +} +``` + +这还不能直接给 DSpark。Producer 应复用当前: + +```python +TargetFeatureReplayer._feature_from_vllm_payload(...) +``` + +完成 token 校验、position 对齐、选层和 flatten。 + +#### Producer 最终得到的 DraftFeatureSample + +```python +feature = DraftFeatureSample( + algorithm="DSpark", + input_ids=int64[feature_rows], + loss_mask=float32[feature_rows], + position_ids=int64[feature_rows], + hidden_states=bf16[feature_rows, feature_hidden_dim], + metadata={ + "hidden_states_layout": "dflash_aux_plus_last", + "target_layer_ids": [8, 16, 24], + "target_model_path": "...", + "target_config_fingerprint": "...", + "feature_start": 128, + "feature_end": 640, + "sequence_length": 512, + }, +) +``` + +若有 3 个 aux layer、target hidden size 为 4096,并包含 final hidden: + +```text +feature_hidden_dim = 3 * 4096 + 4096 = 16384 +hidden_states.shape = [feature_rows, 16384] +``` + +前 `12288` 维是 aux/context hidden,最后 `4096` 维是 DSpark L1 所需 final hidden。现有 DSpark backend 会根据 `dflash_aux_plus_last` 自动拆分。 + +#### 写入 TQ 的 data fields + +第一版建议一个 key 对应一个已经规范化完成的 `DraftFeatureSample`: + +```python +tq_fields = { + "input_ids": feature.input_ids.cpu(), + "loss_mask": feature.loss_mask.cpu(), + "position_ids": feature.position_ids.cpu(), + "hidden_states": feature.hidden_states.cpu(), +} +``` + +这里只放 tensor,因为 PR #48 的 `put_sample()` 会过滤非 tensor: + +```python +payload = { + key: value + for key, value in tensor_dict.items() + if torch.is_tensor(value) +} +``` + +#### 写入 TQ 的 tag metadata + +```python +tq_tag = { + "run_id": "run-20260818-001", + "status": "ready", + "sample_id": sample_key, + "sequence_no": 12345, + "algorithm": "DSpark", + "hidden_states_layout": "dflash_aux_plus_last", + "target_model_fingerprint": "sha256:...", + "target_layer_ids": "8,16,24", + "feature_rows": 512, + "hidden_dim": 16384, + "payload_bytes": 16777216, +} +``` + +tag 中只放 TQ 版本支持序列化的小标量/字符串。列表等复杂对象可以编码为稳定字符串或 JSON。`payload_bytes` 用于背压统计。 + +#### TQ 中逻辑上保存的 row + +```text +partition: speco_drafter_features_run-20260818-001 +key: 86a4...ef2 + +fields: + input_ids → int64[512] + loss_mask → float32[512] + position_ids → int64[512] + hidden_states → bf16[512, 16384] + +tag: + status → ready + sequence_no → 12345 + hidden_layout → dflash_aux_plus_last + target_model_fp → sha256:... +``` + +#### Consumer 恢复出的对象 + +rank 0 通过 `kv_list` 同时取得 key 和 tag;每个 rank 用 local keys 调用 `kv_batch_get` 取得 fields,然后组合: + +```python +feature = DraftFeatureSample( + algorithm=tag["algorithm"], + input_ids=densify(fields["input_ids"]).reshape(-1), + loss_mask=densify(fields["loss_mask"]).reshape(-1), + position_ids=densify(fields["position_ids"]).reshape(-1), + hidden_states=densify(fields["hidden_states"]), + metadata={ + "hidden_states_layout": tag["hidden_states_layout"], + "target_model_fingerprint": tag["target_model_fingerprint"], + }, +) + +feature.validate(strict=True) +``` + +这样传给: + +```python +trainer.prepare_training_batch_from_samples([feature, ...]) +``` + +的数据结构,与现有文件 feature store 读出的 `DraftFeatureSample` 一致。也就是说,TQ 替换的是存取介质和调度方式,不改变 DSpark backend 的样本契约。 + +## 5. 新的整体架构 + +```text + 小 metadata + ┌──────────────────────────┐ + │ TransferQueueController │ + │ KV metadata / key / tags │ + │ partition / storage map │ + └────────────┬─────────────┘ + │ +JSONL/token replay │ + │ │ + ▼ │ +Feature Producer │ + ├─ tokenizer/window │ + ├─ asyncio bounded concurrency │ + ├─ vLLM endpoint pool │ + ├─ validate/pack │ + └─ TQ put ─────────────────────┤ + ▼ + TQ Mooncake backend + hidden-state tensors + │ + ┌───────────────────┼───────────────────┐ + ▼ ▼ ▼ + DSpark rank 0 DSpark rank 1 DSpark rank N + TQ get TQ get TQ get + └───────────────────┼───────────────────┘ + ▼ + synchronized optimizer step + │ + ▼ + TQ clear after success +``` + +大 tensor 的路径是: + +```text +vLLM/Producer memory → TQ Mooncake backend → each training rank +``` + +不会走: + +```text +Mooncake → rank 0 → rank 1/2/3 +``` + +rank 0 最多只广播 `kv_list` 得到的 key 字符串列表;tags 只在选 batch 时由 rank 0 使用。 + +### 5.1 standalone 每一步为什么能实现推理和训练异步 + +#### 步骤 A:Producer 独立推进输入 cursor + +Producer 自己维护: + +```python +reader_cursor = 12346 +``` + +它不等待 Trainer 请求样本。只要 TQ ready bytes 没超过背压上限,就继续读取文件并创建 vLLM task。 + +效果是 Producer 的执行进度与 `optimizer_step` 解耦: + +```text +Producer sequence_no: 1200,1201,1202,... +Trainer optimizer_step: 87 +``` + +两者通过 TQ 中的 ready rows 衔接,不互相直接调用。 + +#### 步骤 B:并发 vLLM task 完成顺序可以乱序 + +例如 Producer 同时提交: + +```text +sequence_no 100 → endpoint 0 +sequence_no 101 → endpoint 1 +sequence_no 102 → endpoint 0 +``` + +完成顺序可能是: + +```text +101 → 100 → 102 +``` + +每个 task 完成后独立执行 `kv_put`,所以慢请求不会阻塞已经完成的请求写入。tag 中的 `sequence_no` 保留原数据顺序。 + +#### 步骤 C:`kv_put` 成功是 READY 可见性的边界 + +Producer 在调用前已经得到完整 `DraftFeatureSample`。一次 `kv_put` 写入该 sample 的全部 tensor fields,并在 tag 中标记 `status=ready`。 + +因此 Consumer 的判断规则是: + +```text +kv_list 能列出该 key +且 tag.run_id 匹配 +且 tag.status == ready +→ 可以尝试 kv_batch_get +``` + +Consumer 仍需对取回 fields 做完整性校验;tag 是调度 metadata,不是正确性证明。 + +#### 步骤 D:rank 0 只负责选 key + +rank 0 执行: + +```python +entries = list_ready_keys() +selected = sorted(entries, key=sequence_no)[:global_batch_size] +``` + +这一步处理的数据只是: + +```python +[ + {"key": "k100", "sequence_no": 100, ...}, + {"key": "k101", "sequence_no": 101, ...}, +] +``` + +不包含 `[seq, hidden_dim]` hidden tensor,所以 rank 0 不成为大数据中转瓶颈。 + +#### 步骤 E:广播保证所有 rank 对同一个 global step 达成一致 + +所有 rank 调用同一次: + +```python +dist.broadcast_object_list(holder, src=0) +``` + +广播结束后,每个 rank 看到完全相同的 global key list。然后按确定性区间切分: + +```text +rank 0: keys[0:per_rank] +rank 1: keys[per_rank:2*per_rank] +... +``` + +这样不会出现 rank 0 训练 batch A、rank 1 训练 batch B 数量不同,或者某个 rank 没进入 backward 的情况。 + +#### 步骤 F:各 rank 直接读取 Mooncake 后端 + +每个 rank 执行: + +```python +tq.kv_batch_get(keys=local_keys, partition_id=partition_id) +``` + +TQ 根据 key 找到 storage backend 中的数据。使用 MooncakeStore 时,大 tensor 数据路径是 Mooncake → 本 rank;rank 0 不读取其他 rank 的 local payload。 + +因此: + +```text +控制面:rank 0 → broadcast small keys +数据面:Mooncake → each rank directly +``` + +#### 步骤 G:恢复现有 DraftFeatureSample 契约 + +每个 rank 将 TQ fields 与 tag 合并、densify、校验,得到 `list[DraftFeatureSample]`。从这里开始继续执行项目现有代码: + +```python +batch = trainer.prepare_training_batch_from_samples( + materialized_samples, + step=optimizer_step, +) + +ok = await trainer.training_step_from_batch( + batch, + optimizer_step, +) +``` + +所以 TQ 不进入 DSpark model/backend 内部,训练数学逻辑不变。 + +#### 步骤 H:全 rank 成功以后才能清理 + +每个 rank 的 `ok` 通过现有 `_all_ranks_true()` 聚合: + +```text +rank 0 ok = true +rank 1 ok = true +rank 2 ok = true +rank 3 ok = true +→ global_ok = true +``` + +只有此时 rank 0 执行: + +```python +tq.kv_clear(keys=global_keys, partition_id=partition_id) +``` + +这样保证被删除的数据已经参与完成的 optimizer step。若任何 rank get/OOM/backward 失败,不执行 clear,便于作业失败后的诊断或恢复。 + +#### 步骤 I:异步重叠如何形成 + +时间线上: + +```text +时间 ─────────────────────────────────────────▶ + +Producer: vLLM(batch N+1) ─ put ─ vLLM(batch N+2) ─ put +Trainer: get(batch N) ─ train(batch N) ─ get/train(batch N+1) +``` + +Producer 和 Trainer 是不同进程,互相不调用;TQ ready rows 是缓冲区。因此 vLLM prefill、网络传输和 DSpark GPU 训练可以重叠。背压只在缓冲区达到容量上限时暂停 Producer。 + +## 6. Producer:读取预生成 response 并并行请求 vLLM + +### 6.1 输入处理 + +Producer 从现有 JSONL/token replay 数据源读取: + +```python +sample = { + "row_id": "12345", + "prompt": "...", + "response": "提前生成好的文本", +} +``` + +构造: + +```python +prompt_ids = tokenizer.encode(sample["prompt"]) +response_ids = tokenizer.encode(sample["response"]) +input_ids = prompt_ids + response_ids +``` + +同时产生: + +```python +loss_mask +position_ids +feature_positions +sample_key +``` + +### 6.2 有界并发 + +不能按样本串行请求: + +```python +for sample in samples: + result = request_vllm(sample) +``` + +改成: + +```python +async def run_producer(samples): + semaphore = asyncio.Semaphore(max_inflight_requests) + + async def run_one(sample): + async with semaphore: + result = await vllm_pool.prefill(sample) + feature = validate_and_pack(sample, result) + await tq_transport.put(feature) + + async with asyncio.TaskGroup() as group: + for sample in samples: + group.create_task(run_one(sample)) +``` + +`max_inflight_requests` 是 Producer 同时未完成的请求数量,不是 vLLM batch size。vLLM 服务端仍会对同时到达的请求做自己的 continuous batching。 + +### 6.3 多 endpoint + +多个 endpoint 例如: + +```yaml +vllm_endpoints: + - http://node0:8000/v1 + - http://node1:8000/v1 + - http://node2:8000/v1 +``` + +调度器维护每个 endpoint 的 inflight 数: + +```python +endpoint = min( + endpoints, + key=lambda item: item.inflight, +) +``` + +请求前 `inflight += 1`,在 `finally` 中 `inflight -= 1`。失败只重试对应样本,不阻塞全部 Producer。 + +### 6.4 当前 vLLM 文件桥接 + +当前客户端协议期望: + +```python +response.kv_transfer_params["hidden_states_path"] +``` + +所以第一阶段仍然是: + +```text +vLLM 写临时 safetensors +→ Producer load_file +→ 校验 token_ids/hidden_states +→ TQ put 到 Mooncake backend +→ TQ put 成功后删除临时文件 +``` + +删除必须发生在 TQ put 成功之后: + +```python +path = request_vllm_hidden_file(sample) +try: + feature = load_and_validate(path) + await tq_transport.put(feature) +finally: + if put_succeeded: + Path(path).unlink(missing_ok=True) +``` + +### 6.5 目标版本:vLLM 直接写 TQ/Mooncake + +目标响应可改成: + +```json +{ + "kv_transfer_params": { + "backend": "transfer_queue", + "partition_id": "speco:run-1:dspark_train", + "sample_key": "abc123" + } +} +``` + +服务端顺序必须是: + +```text +prefill +→ 捕获指定层 hidden states +→ TQ/Mooncake put 完成 +→ 返回 HTTP success 和 sample key +``` + +这需要定制 vLLM exporter;当前 `verl-SpeCo-ls` 中没有服务端 writer 实现。 + +## 7. 按 PR #48 扩展 TQ bridge + +不要重新发明一套 transport。将 PR #48 的 bridge 设计移植到本项目并增加 standalone 所需方法: + +```python +class StandaloneTQTransport: + def put_sample(self, key, tensor_dict, tag): ... + def list_ready_keys(self, run_id): ... + def get_samples(self, keys, fields=None): ... + def clear_samples(self, keys): ... + def put_control(self, key, tag): ... + def close(self): ... +``` + +写入延续 PR #48 的真实形式: + +```python +tq.kv_put( + key=key, + partition_id=partition_id, + fields={ + "input_ids": feature.input_ids.cpu(), + "loss_mask": feature.loss_mask.cpu(), + "position_ids": feature.position_ids.cpu(), + "hidden_states": feature.hidden_states.cpu(), + }, + tag={ + "run_id": run_id, + "status": "ready", + "sequence_no": sequence_no, + "payload_bytes": payload_bytes, + }, +) +``` + +批量读取延续 PR #48 的 `kv_batch_get`: + +```python +result = tq.kv_batch_get( + keys=keys, + partition_id=partition_id, + fields=required_fields, # 0.1.7 是否支持该参数需实机确认 +) +``` + +新增发现和清理: + +```python +items = tq.kv_list(partition_id=partition_id) +tq.kv_clear(keys=keys, partition_id=partition_id) +``` + +这里的 `kv_list/kv_clear` 参数名必须根据锁定的 TQ 版本验证。当前参考仓库只实际调用了 `kv_put/kv_batch_get`,没有为这两个接口提供运行证据。 + +读取结果继续复用 PR #48 的两个适配函数: + +```python +value = _extract_value(result, key) +row = _tensordict_to_dict(value) +row["hidden_states"] = _densify_tq_tensor(row["hidden_states"]) +``` + +## 8. DSpark 多 rank 如何消费 + +### 8.1 第一版:rank 0 用 kv_list 发现 READY keys + +PR #48 中 key 由 Ray driver 传给 drafter worker;standalone 没有这条控制路径,所以 rank 0 需要主动列出 key: + +```python +rank = dist.get_rank() +world_size = dist.get_world_size() +global_batch_size = batch_size_per_gpu * world_size + +if rank == 0: + entries = tq_transport.list_ready_keys(run_id=run_id) + entries.sort(key=lambda x: (x.tag["sequence_no"], x.key)) + selected_keys = [x.key for x in entries[:global_batch_size]] +else: + selected_keys = None + +holder = [selected_keys] +dist.broadcast_object_list(holder, src=0) +selected_keys = holder[0] +``` + +`sequence_no` 是输入文件顺序。它让多个 Producer 并发完成顺序不同的情况下,Trainer 仍能确定性地组成 batch。 + +rank 0 广播的是字符串 key 列表,不是 hidden-state tensor。 + +### 8.2 各 rank 切自己的 keys + +例如 global batch keys: + +```text +[s0, s1, s2, s3, s4, s5, s6, s7] +``` + +world size 为 4、每卡 batch size 为 2: + +```text +rank 0 → [s0, s1] +rank 1 → [s2, s3] +rank 2 → [s4, s5] +rank 3 → [s6, s7] +``` + +代码: + +```python +def shard_keys(keys, rank, world_size): + assert len(keys) % world_size == 0 + per_rank = len(keys) // world_size + start = rank * per_rank + end = start + per_rank + return keys[start:end] +``` + +### 8.3 每个 rank 并行 get + +所有进程执行: + +```python +local_keys = shard_keys( + selected_keys, + rank=rank, + world_size=world_size, +) + +local_payloads = tq_transport.get_samples(local_keys) +``` + +数据路径: + +```text +rank 0 ← Mooncake(s0,s1) +rank 1 ← Mooncake(s2,s3) +rank 2 ← Mooncake(s4,s5) +rank 3 ← Mooncake(s6,s7) +``` + +不是 rank 0 get 全部后再 scatter。 + +### 8.4 转成当前训练格式 + +TQ 返回的数据需要先按 PR #48 的规则解包、densify,再转换成现有 `DraftFeatureSample`: + +```python +def tq_row_to_feature(row, tag): + return DraftFeatureSample( + algorithm="DSpark", + input_ids=row["input_ids"], + loss_mask=row["loss_mask"], + position_ids=row["position_ids"], + hidden_states=row["hidden_states"], + metadata={ + "hidden_states_layout": tag["hidden_states_layout"], + **row.get("metadata", {}), + }, + ) +``` + +`hidden_states_layout` 不能只放在无法恢复的临时 Python 对象里;应随 TQ tag 保存。rank 0 从 `kv_list` 取得 key/tag 后,需要把每个 local key 对应的 tag 一起交给本 rank 的转换逻辑。Consumer 必须把 layout 放回 `DraftFeatureSample.metadata`。DSpark L1 所需的 `target_last_hidden_states` 随后由现有 backend 从拼接的 `hidden_states` 中切出,transport 层不计算 loss,也不拆 layout。 + +## 9. 修改当前训练循环 + +在 `run_standalone_draft_training()` 中增加数据源分支: + +```python +feature_store_type = str(feature_store_cfg.get("type", "torch_shard")) + +if feature_store_type == "transfer_queue": + tq_stream = build_transfer_queue_stream( + config=config, + rank=rank, + world_size=world_size, + ) + store = None + loader = None + feature_replayer = None +else: + store = build_feature_store_from_config( + feature_store_cfg, + read_only=True, + ) + loader = DraftFeatureDataLoader(...) +``` + +流式训练循环: + +```python +while successful_steps < max_steps: + global_keys, materialized_samples = tq_stream.next_local_batch() + + batch = trainer.prepare_training_batch_from_samples( + materialized_samples, + step=optimizer_step, + ) + + has_batch = batch is not None + if not _all_ranks_true(has_batch, trainer.runtime_device): + raise RuntimeError("at least one rank failed to fetch its TQ batch") + + ok = await trainer.training_step_from_batch( + batch, + optimizer_step, + ) + + if not _all_ranks_true(ok, trainer.runtime_device): + raise RuntimeError("DSpark step failed on at least one rank") + + dist.barrier() + if rank == 0: + tq_stream.clear_global_batch(global_keys) + dist.barrier() +``` + +TQ 接在数据输入层,不放进 `DSparkTrainerBackend`。Backend 只负责模型、forward、loss、backward 和 optimizer。 + +## 10. READY key、inflight key 和训练提交 + +### 10.1 Ready + +在本方案中 ready 表示: + +> `dspark_train` 所需的所有 fields 已经成功写入 TQ storage,训练侧可以读取。 + +第一版不是通过 PR #48 尚未验证的 Sampler 自动判定,而是 Producer 完整校验 payload 后一次 `kv_put`,并设置 `tag.status=ready`。rank 0 的 `kv_list` 只选择这些 key。 + +### 10.2 Inflight key + +rank 0 选出一个 global batch 后,在本地保存: + +```python +inflight_global_keys = selected_keys +``` + +其他 step 不应再次选中这批 key。因为 standalone 只有一个同步 DSpark job,第一版由 rank 0 进程内 `set` 排除 inflight key 即可: + +```python +ready = [x for x in listed if x.key not in inflight_keys] +``` + +若以后允许多个独立 Trainer job 同时消费同一 partition,进程内 set 就不够,届时必须使用 TQ Controller/Sampler 的原子消费分配能力。 + +### 10.3 Optimizer committed + +optimizer committed 表示所有 DSpark rank 已经完成: + +```text +forward → backward → gradient synchronization → optimizer.step +``` + +它比 `kv_list` 发现 key、甚至 `kv_batch_get` 取出 tensor 都更晚。 + +第一版推荐简单语义: + +```text +TQ 负责 key/tag 和 tensor 传输 +rank 0 负责单 Trainer job 的 batch 选择和 inflight set +训练失败 → 整个作业 fail-fast +训练成功 → kv_clear payload,并从 inflight set 移除 +恢复 → 从最近 checkpoint + 输入 cursor 重新启动 +``` + +这样不需要重新实现 lease/ack 状态机,也没有夸大 PR #48 当前尚未使用的 Sampler 能力。 + +## 11. 为什么训练完一个 step 才清理 + +不能在 `kv_batch_get()` 后立即 clear: + +```text +get 成功 +→ clear +→ forward OOM +→ 数据已不存在,无法重试 +``` + +正确顺序: + +```text +rank 0..N get +→ 所有 rank 确认 batch 有效 +→ training_step_from_batch +→ _all_ranks_true(ok) +→ rank 0 kv_clear global keys +``` + +当前 `_all_ranks_true()` 已经是项目现有的跨 rank 同步工具,可以继续复用。 + +## 12. 背压 + +背压表示 Producer 生成速度高于 Trainer 消费速度时,Producer 必须暂停继续提交,以免 Mooncake 内存无限增长。 + +建议限制: + +```yaml +max_vllm_inflight_requests: 32 +max_pending_put_bytes: 8589934592 +max_tq_ready_samples: 256 +max_tq_ready_bytes: 68719476736 +``` + +Producer 在 tags 中写: + +```python +{"payload_bytes": payload_bytes} +``` + +周期性通过 TQ list/metadata 统计当前 partition 尚未消费的数据量: + +```python +while ready_bytes >= max_tq_ready_bytes: + await asyncio.sleep(backpressure_poll_interval) +``` + +如果锁定的 TQ 版本提供容量/ready 统计接口,应直接使用,避免全量扫描 keys。 + +## 13. Stable ID、幂等和孤儿数据 + +### 13.1 幂等 + +幂等表示同一操作重复执行,最终逻辑结果仍只有一份。 + +Producer 对同一样本重试时必须使用相同 `sample_key`: + +```python +await tq.put(key="abc123", ...) +await tq.put(key="abc123", ...) +``` + +不能每次生成随机 key: + +```text +abc123-retry-1 +abc123-retry-2 +``` + +否则一个输入可能训练多次并持续占用 Mooncake。 + +### 13.2 孤儿数据 + +孤儿数据表示 tensor 已经写进 storage,但由于 Producer 崩溃或 metadata 更新失败,没有进入正常消费路径。 + +使用 TQ 后,metadata 和 storage 由同一套系统管理,可以减少“Mooncake有对象、自研 Coordinator 没记录”的双系统窗口,但仍要配置 TTL/partition cleanup: + +```text +训练正常结束 → clear partition +训练异常退出 → 下次启动检查旧 partition +超过 TTL → 清理未消费数据 +``` + +## 14. EOS 和 drop-last + +EOS 表示 Producer 已经读完输入,并且所有 vLLM/TQ put 都已完成。 + +TQ 中需要一种结束条件,具体使用 Controller API、特殊 metadata 或单独的运行状态记录取决于锁定版本。不能把“当前暂时没有 ready sample”当成 EOS,因为 Producer 可能仍在请求 vLLM。 + +最后不足一个 global batch 时: + +```python +global_batch_size = batch_size_per_gpu * world_size +``` + +第一版使用 `drop_last=true`,避免某些 rank 有数据、某些 rank 没数据导致分布式训练不同步。 + +结束条件: + +```text +producer_done == true +and ready_samples < global_batch_size +and inflight_requests == 0 +and pending_puts == 0 +``` + +## 15. 双缓冲预取 + +训练 batch N 时,CPU 后台线程预取 batch N+1: + +```python +next_future = executor.submit(tq_stream.next_local_batch) + +current_batch = first_batch +while current_batch is not None: + next_batch = next_future.result() + next_future = executor.submit(tq_stream.next_local_batch) + + train(current_batch) + current_batch = next_batch +``` + +实际顺序应调整为避免等待 future 后才训练。推荐: + +```python +current = tq_stream.next_local_batch() + +while current is not None: + future = executor.submit(tq_stream.next_local_batch) + train_and_clear(current) + current = future.result() +``` + +第一版只预取一个 global batch,避免 Trainer 崩溃时大量数据已被采样但未训练。 + +如果 Mooncake/TQ Python get 是阻塞函数,使用专用 `ThreadPoolExecutor`,不要阻塞 Producer 的 asyncio event loop。 + +## 16. 建议代码结构 + +```text +verl_speco/ + trainer/ + tq_transport.py # TQ client、put/get/meta/clear 封装 + tq_feature_stream.py # kv_list、global keys、rank shard、decode/prefetch + feature_producer.py # JSONL → 并发 vLLM → TQ + draft_training_loop.py # 增加 transfer_queue 数据源分支 + target_feature_replay.py # 复用 vLLM payload 校验/转换逻辑 +``` + +不要新增: + +```text +coordinator.py +coordinator_client.py +``` + +建议抽象: + +```python +class StreamingFeatureSource(Protocol): + def next_local_batch(self) -> tuple[list[str], list[DraftFeatureSample]]: ... + def clear_global_batch(self, keys: list[str]) -> None: ... + def close(self) -> None: ... +``` + +这样训练循环不依赖 TQ 的具体类型。 + +## 17. 配置草案 + +```yaml +actor_rollout_ref: + rollout: + drafter: + training: + mode: offline + backend: dspark + batch_size_per_gpu: 2 + max_steps: 1000 + + feature_store: + type: transfer_queue + partition_id: speco_drafter_features_${run_id} + drop_last: true + prefetch_steps: 1 + + transfer_queue: + # 与 PR #48 的配置层级和 init 方式保持一致。 + enable: true + package_version: 0.1.8 # 最终以实测版本为准 + backend: + storage_backend: MooncakeStore + MooncakeStore: + auto_init: false + metadata_server: localhost:50123 + master_server_address: localhost:50124 + local_hostname: localhost + protocol: tcp + global_segment_size: 4294967296 + local_buffer_size: 1073741824 + device_name: "" + + required_fields: + - input_ids + - loss_mask + - position_ids + - hidden_states + + producer: + input_path: /path/to/generated_responses.jsonl + vllm_endpoints: + - http://node0:8000/v1 + - http://node1:8000/v1 + max_inflight_requests: 32 + max_pending_put_bytes: 8589934592 + max_ready_samples: 256 + max_ready_bytes: 68719476736 +``` + +当前 examples 中的: + +```bash +transfer_queue.enable=False +``` + +属于 verl RL 主入口配置,当前 standalone `draft_train_launcher` 不读取它。不能只改成 `True`;必须实现上述 `feature_store.type=transfer_queue` 分支。 + +## 18. 启动顺序 + +逻辑顺序: + +```text +1. 启动 Mooncake metadata/master 服务; +2. 启动 standalone TQ owner 进程,调用一次 `tq.init(tq_config)` 创建/连接 TQ Controller 和 MooncakeStore backend; +3. 启动一个或多个定制 vLLM server +4. 启动 Feature Producer +5. Producer 和所有 Trainer rank 调用无参 `tq.init()`,连接 owner 创建的同一套 TQ; +6. 启动 verl_speco.draft_train_launcher +7. torchrun 启动所有 DSpark rank +8. 各 rank 连接 TQ +9. rank 0 用 `kv_list` 选择 global keys,各 rank 并行 `kv_batch_get`/train; +10. 输入耗尽后 Producer 发布 done 状态 +11. Trainer drain 完整 global batches 后退出 +12. 清理 partition,停止 Producer、vLLM、TQ、Mooncake +``` + +PR #48 的 `init_transfer_queue()` 在 Ray `SpecoTaskRunner` 中调用 `tq.init(config)`,worker 的无参 `tq.init()`依靠同一个 Ray 集群发现 named Controller;这部分不能原样复制到 standalone torchrun。 + +本项目要求一开始不使用 Ray,因此 Phase 0 必须先证明锁定的 TQ 版本支持独立 owner/controller 进程,以及 Producer/torchrun rank 如何获得连接信息。若 TQ 0.1.7/0.1.8 实际只能通过 Ray named actor 完成发现,那么有两个选择: + +1. 接受仅用 Ray 承载 TQ 控制面的最小方案; +2. 给 TQ 增加或使用其已有的独立 ZMQ/server-info bootstrap。 + +在这项验证完成前,文档不能声称 PR #48 已经提供“无 Ray TQ standalone 启动”。 + +## 19. 故障处理 + +### vLLM 请求失败 + +- 对单个 sample 按稳定 key 重试; +- 指数退避; +- 超过次数记录失败,并根据配置 fail-fast 或跳过; +- 不写不完整 TQ fields。 + +### vLLM 文件读取成功,但 TQ put 失败 + +- 暂时保留临时文件; +- 重试 TQ put; +- put 成功后再删除; +- 不把样本视为 ready。 + +### 某个训练 rank get 失败 + +- 该 rank 报告 `local_ok=false`; +- `_all_ranks_true()` 使全部 rank 得到一致失败结果; +- 第一版整个训练 fail-fast; +- 不 clear global batch。 + +### OOM/optimizer step 失败 + +- 不 clear; +- 所有 rank 一致退出; +- 从最近训练 checkpoint 恢复; +- 根据 TQ 消费提交语义决定是否重放当前 batch。 + +### clear 失败 + +- optimizer 已成功,不能再次训练这批; +- 将 batch keys 写入本地小型 `gc_pending` 日志; +- 后台重试 clear; +- checkpoint 保存最近 committed sample IDs,避免恢复时重复消费。 + +## 20. 观测指标 + +Producer: + +```text +producer/vllm_inflight +producer/vllm_requests_per_sec +producer/vllm_prefill_tokens_per_sec +producer/vllm_p50_latency +producer/vllm_p95_latency +producer/tq_put_bytes_per_sec +producer/tq_put_failures +producer/pending_put_bytes +``` + +TQ/Mooncake: + +```text +tq/ready_samples +tq/ready_bytes +tq/consumed_samples +tq/storage_bytes +tq/clear_failures +mooncake/put_bandwidth +mooncake/get_bandwidth +``` + +Trainer: + +```text +trainer/tq_wait_seconds +trainer/tq_get_seconds +trainer/tq_get_bytes_per_sec +trainer/decode_seconds +trainer/h2d_seconds +trainer/step_seconds +trainer/data_stall_ratio +trainer/successful_steps +``` + +## 21. 实施阶段 + +### Phase 0:锁定依赖和契约 + +- 从 PR #48 的 `TransferQueue==0.1.7` 起验证,同时对比当前 verl 文档使用的 0.1.8; +- 实测 `tq.init(config)`、worker `tq.init()`、`kv_put`、`kv_batch_get`、`kv_list`、`kv_clear`; +- 验证该版本是否支持无 Ray owner/controller 部署以及连接信息传递; +- 先用 PR #48 已配置的 `SimpleStorage` 做最小闭环; +- 再把 backend 换成 `MooncakeStore`,验证 tcp,最后再验证 rdma; +- 固定 fields、tags、partition 和 `run_id/sequence_no/status`; +- 写 fake TQ 单元测试。 + +### Phase 1:文件桥接 + TQ KV 模式 + +- 新增独立 Producer; +- 32 个有界并发 vLLM 请求; +- 读取 vLLM 临时 safetensors; +- TQ put 成功后删除文件; +- standalone trainer 由 rank 0 `kv_list` 并广播 global keys; +- 各 rank 并行 `kv_batch_get`; +- 复用 PR #48 的返回值解包、NestedTensor densify 和 per-step cache; +- optimizer 成功后 `kv_clear`。 + +验收:连续训练 1000 step,临时文件数量、TQ ready bytes 和 Mooncake占用均保持有界。 + +### Phase 2:双缓冲与多 endpoint + +- 增加多 endpoint 最少 inflight 调度; +- 增加一个 global batch 预取; +- 动态背压; +- 注入单 rank get 失败,确认所有 rank 一致退出而非死锁。 + +### Phase 3:vLLM 直接写 TQ/Mooncake + +- 修改外部定制 vLLM exporter; +- 去掉 `hidden_states_path` 临时文件; +- HTTP 响应返回 partition/sample key; +- 验证 HTTP 重试的幂等性。 + +### Phase 4:可选升级到 TQ StreamingDataLoader + +- 在当前保守方案稳定后再引入 RankAwareSampler; +- 让每个 rank 自动取得 local micro-batch; +- 去掉 rank 0 手工 key-list 广播; +- 验证与 torchrun/DSpark 的 global step 对齐。 + +## 22. 最终推荐 + +针对当前 `verl-SpeCo-ls`,推荐的第一版不是 AngelSpec 的完整架构,也不是单独写 Coordinator,而是: + +```text +当前预生成 response 文件 +→ 独立 asyncio Producer +→ 并行访问多个 vLLM endpoint +→ 读取并校验临时 hidden-state 文件 +→ TransferQueue put +→ Mooncake storage backend +→ rank 0 kv_list 获取 READY global keys +→ broadcast key list +→ 各 DSpark rank 并行 kv_batch_get +→ 现有 prepare_training_batch_from_samples() +→ 现有 training_step_from_batch() +→ 全 rank 成功 +→ TQ clear +``` + +这套方案保留当前 standalone DSpark 训练主体,只替换 `DraftFeatureDataLoader + TargetFeatureReplayer.materialize()` 所在的数据输入路径。它真正建立在 PR #48 已实现的 KV transport 之上,而不是假设 PR #48 已经实现了 standalone Sampler/StreamingDataLoader。 + +## 23. 参考 + +- verl TransferQueue: +- TransferQueue: +- Mooncake Store: +- 本地只读参考:`../verl-SpeCo/verl_speco/integration/transferqueue_bridge.py`(PR #48) +- 本地只读参考:`../verl-SpeCo/docs/transferqueue_integration_plan.md`(PR #48) diff --git a/docs/standalone_tq_foundation_implementation.md b/docs/standalone_tq_foundation_implementation.md new file mode 100644 index 00000000..5bb7b48e --- /dev/null +++ b/docs/standalone_tq_foundation_implementation.md @@ -0,0 +1,1037 @@ +# Standalone TQ 公共基础层实现说明 + +## 1. 文档范围和已验证结论 + +本文解释当前仓库已经实现并测试通过的 TQ 公共基础层: + +```text +verl_speco/transport/drafter_sample_protocol.py +verl_speco/integration/transferqueue_bridge.py +verl_speco/config/speco_base.yaml +verl_speco/tq_owner.py +examples/run_dspark_tq_owner.sh +examples/tq_connection_smoke.py +tests/unit/test_drafter_sample_protocol.py +tests/unit/test_transferqueue_bridge.py +pyproject.toml +``` + +当前已经实现: + +1. Producer 和 Consumer 共用的样本 key、tag、fields 和 metadata 协议; +2. 普通进程连接 Ray 集群; +3. TQ Owner 创建 named `TransferQueueController`; +4. 独立 Client 发现并连接同一个 Controller; +5. 单样本 put、元数据 list、批量 get 和批量 clear; +6. Owner 与 Client 不同的关闭边界; +7. 独立 Owner 入口和共享 Hydra 配置; +8. mock 单元测试和真实双进程 TQ 0.1.7 smoke test。 + +当前还没有实现: + +1. Producer 读取 JSONL、并发访问多个 vLLM endpoint 的完整 pipeline; +2. `feature_store.type=tq` 工厂分支; +3. `TQFeatureStore` 和 `TQFeatureDataLoader`; +4. rank 0 选择 global keys、各 rank 读取 local keys; +5. TQ batch 接入 DSpark optimizer step; +6. optimizer step 成功后的 rank 0 clear。 + +因此当前代码已经证明“两个独立进程能通过 Ray 连接同一个 TQ,并按共享协议批量传输多个样本”,但尚未接到正式 Producer 和 DSpark Consumer 主循环。 + +## 2. 运行时角色和术语 + +### 2.1 Ray head + +Ray head 是 Ray 集群的控制节点。TQ 0.1.7 使用 Ray 管理 named Controller 和 storage actors。 + +Ray 不传输 standalone hidden-state payload。当前代码没有对这些样本调用: + +```python +ray.put(hidden_states) +``` + +### 2.2 TQ Owner + +TQ Owner 是普通 Python OS 进程,入口为: + +```text +python -m verl_speco.tq_owner +``` + +它不是 Ray actor。它先调用 `ray.init(address=...)` 加入 Ray,再调用带完整配置的 `tq.init(config)` 创建 TQ Controller 和 storage。 + +Owner 是唯一允许调用全局 `tq.close()` 的进程。 + +### 2.3 Named TransferQueueController + +TQ 0.1.7 内部创建: + +```python +TransferQueueController.options( + name="TransferQueueController" +).remote(...) +``` + +`name` 将 Controller actor 注册到 Ray actor registry。其他进程加入同一 Ray address 和 namespace 后,通过: + +```python +ray.get_actor("TransferQueueController") +``` + +取得 actor handle,再读取 TQ backend 配置。 + +Controller 保存控制信息和数据位置;使用 MooncakeStore 时,大 tensor 本身保存在 MooncakeStore。 + +### 2.4 TQ Client + +Owner、Producer、每个 Consumer rank 都在各自 OS 进程内拥有独立 TQ Client。 + +普通 Client 通过无参: + +```python +tq.init() +``` + +发现 named Controller并初始化本进程的 storage manager。Client 不是 Ray actor,Producer 和 torchrun rank 也不需要改成 Ray actor。 + +### 2.5 Partition、key、tag 和 fields + +当前固定 partition: + +```text +speco_drafter_features +``` + +TQ 中一条记录逻辑上是: + +```text +partition_id +└── key + ├── tag:轻量 dict,由 kv_list 发现 + └── fields:Tensor payload,由 kv_batch_get 读取 +``` + +## 3. 共享配置如何工作 + +共享配置定义在 `verl_speco/config/speco_base.yaml`: + +```yaml +transfer_queue: + enable: false + package_version: "0.1.7" + ray: + address: null + namespace: speco-drafter + partition_id: speco_drafter_features + run_id: null + schema_version: 1 + connect_timeout_seconds: 120 + poll_interval_seconds: 0.5 + drop_last: true + controller: + polling_mode: true + backend: + storage_backend: SimpleStorage + SimpleStorage: + total_storage_size: 100000 + num_data_storage_units: 8 + MooncakeStore: + auto_init: false + metadata_server: localhost:50050 + master_server_address: localhost:50051 + local_hostname: "" + protocol: tcp + global_segment_size: 4294967296 + local_buffer_size: 1073741824 + device_name: "" +``` + +### 3.1 Ray 连接字段 + +```yaml +ray: + address: 10.0.0.1:6379 + namespace: speco-drafter +``` + +它们决定当前进程连接哪个 Ray 集群,以及在哪个 namespace 查找 `TransferQueueController`。Owner、Producer 和所有 Consumer ranks 必须使用相同值。 + +### 3.2 SPECO 协议字段 + +```yaml +partition_id: speco_drafter_features +run_id: dspark-20260819-a +schema_version: 1 +``` + +这些字段用于路由和校验,不属于 TQ 原生配置。`run_id` 用于隔离不同 pipeline 的记录。 + +### 3.3 TQ 原生字段 + +```yaml +controller: ... +backend: ... +``` + +只有这些字段应传给 `tq.init(full_config)`。bridge 的 `_native_tq_config()` 会删除: + +```text +enable +package_version +ray +partition_id +run_id +schema_version +connect_timeout_seconds +poll_interval_seconds +drop_last +``` + +对象变化为: + +```text +完整 SPECO transfer_queue dict +→ _native_tq_config() +→ controller/backend等TQ字段 +→ OmegaConf DictConfig +→ tq.init() +``` + +## 4. Bridge 的进程内状态 + +`verl_speco/integration/transferqueue_bridge.py` 在每个 OS 进程内分别维护: + +```python +_state = { + "enabled": False, + "configured": False, + "initialized": False, + "config": None, + "owner": False, + "ray_initialized_here": False, + "ray_address": None, + "ray_namespace": None, +} +``` + +该 dict 不跨进程共享。Owner、Producer、每个 rank 分别拥有自己的 `_state`。 + +| 字段 | 含义 | +|---|---| +| `enabled` | 当前进程配置是否开启 TQ | +| `configured` | 是否调用过 `configure_transfer_queue()` | +| `initialized` | 当前进程是否执行过 `tq.init()` | +| `config` | 当前进程保存的普通 dict 配置 | +| `owner` | 当前进程是否创建了全局 Controller/Storage | +| `ray_initialized_here` | bridge 是否负责调用了本进程的 `ray.init()` | +| `ray_address/namespace` | 本进程的 Ray 连接信息 | + +`_state_lock` 只保护同一进程内多个线程同时初始化,不是分布式锁。 + +## 5. Owner 的完整启动数据流 + +Owner 入口是 `verl_speco/tq_owner.py`,直接加载 Hydra 主配置 +`verl_speco/config/speco_base.yaml`。`run_owner()` 会复制其中的 +`transfer_queue` 子配置,并仅在 Owner 进程的副本中强制设置 +`enable=true`;普通训练任务看到的共享默认值仍是 `false`。 + +### 阶段 1:读取配置 + +执行者:Owner OS 进程。 + +入口: + +```python +run_owner(config) +``` + +取得: + +```python +training_cfg = config.actor_rollout_ref.rollout.drafter.training +tq_cfg = training_cfg.transfer_queue +``` + +然后调用: + +```python +configure_transfer_queue(training_cfg) +``` + +该函数只把 OmegaConf 转成普通 dict并更新当前进程 `_state`,不会连接 Ray,也不会创建 TQ。 + +### 阶段 2:连接 Ray + +Owner 调用: + +```python +connect_ray_cluster(ray_address, namespace) +``` + +内部执行: + +```python +if not ray.is_initialized(): + ray.init(address=ray_address, namespace=namespace) +``` + +边界类型是 Ray control-plane connection。此时还没有传输训练 tensor。 + +### 阶段 3:创建 Controller 和 Storage + +Owner 调用: + +```python +start_transfer_queue_owner(tq_cfg) +``` + +执行顺序: + +1. `_extract_tq_config()` 得到普通 dict; +2. 检查 `enable=true`; +3. 检查 `TransferQueue` 包可用; +4. 防止本进程重复初始化; +5. `_native_tq_config()` 删除 SPECO 字段; +6. `_as_tq_config()` 转 OmegaConf; +7. `tq.init(native_config)` 创建 Controller、Storage 和 Owner Client; +8. 设置 `_state.owner=True`、`initialized=True`。 + +Ray 中形成: + +```text +Ray cluster / namespace +├── named actor: TransferQueueController +└── storage backend + ├── SimpleStorage actors + └── 或 MooncakeStore connection/process +``` + +### 阶段 4:发布 owner-ready + +调用: + +```python +publish_owner_ready(run_id, schema_version) +``` + +生成: + +```python +key = "control:v1::owner-ready" +fields = {"marker": torch.tensor([1], dtype=torch.uint8)} +tag = { + "record_type": "control", + "status": "owner_ready", + "schema_version": 1, + "run_id": run_id, +} +``` + +这是一条控制记录,不进入训练 batch。 + +### 阶段 5:常驻和关闭 + +Owner 安装 `SIGINT/SIGTERM` handler,并等待: + +```python +stop_event.wait() +``` + +收到信号后调用: + +```python +close_transfer_queue_owner() +``` + +执行: + +```text +验证owner身份 +→ tq.close() +→ 清理Controller/Storage +→ ray.shutdown() +``` + +Owner 必须在 Producer 和 Consumer 退出后才能关闭。 + +## 6. 普通 Client 如何连接同一个 TQ + +Producer 和每个 Consumer rank 后续使用相同顺序: + +```python +configure_transfer_queue(training_cfg) +connect_ray_cluster(ray_address, namespace) +connect_transfer_queue_client() +``` + +`connect_transfer_queue_client()` 最终调用无参: + +```python +tq.init() +``` + +TQ 0.1.7 内部通过: + +```python +ray.get_actor("TransferQueueController") +``` + +找到 Owner 创建的 Controller,读取 backend 配置,然后创建当前进程的 TQ Client。 + +对象和边界变化: + +```text +actor名称字符串 +→ Ray actor registry +→ Controller actor handle +→ Controller.get_config.remote() +→ TQ DictConfig +→ 当前进程TransferQueueClient +→ 同一个SimpleStorage/MooncakeStore +``` + +## 7. 一条具体样本的初始对象 + +真实 smoke test使用: + +```python +sample = DraftFeatureSample( + algorithm="DSPARK", + input_ids=torch.tensor([1, 2, 3]), # int64[3], CPU + loss_mask=torch.tensor([0.0, 1.0, 1.0]), # float32[3], CPU + position_ids=torch.tensor([0, 1, 2]), # int64[3], CPU + hidden_states=torch.arange( + 12, dtype=torch.float32 + ).reshape(3, 4), # float32[3,4], CPU +) +``` + +同时构造: + +```python +meta = SampleMetadata( + schema_version=1, + run_id="codex-batch-smoke", + sample_id="smoke-0000", + sequence_no=0, + algorithm="DSPARK", + target_model_id="smoke-target", + target_model_revision="smoke-revision", + tokenizer_fingerprint="smoke-tokenizer", + target_layer_ids=[0], + hidden_states_layout="dflash_aux", + hidden_dtype="float32", + hidden_shape=[3, 4], + feature_length=3, + full_sequence_length=3, + feature_start=0, + feature_end=3, + use_logits=False, +) +``` + +`DraftFeatureSample` 和 `SampleMetadata` 都是进程内 Python 对象,不直接经过 TQ。 + +## 8. Key 的生成和两个同名函数 + +共享协议调用: + +```python +make_sample_key(meta) +``` + +输出: + +```text +drafter:v1:codex-batch-smoke:000000000000:smoke-0000 +``` + +字段顺序: + +```text +drafter / schema version / run_id / 12位sequence_no / sample_id +``` + +`sequence_no` 是输入文件 record 序号,不是训练 batch 编号。 + +bridge 为兼容 PR #48 还保留另一个: + +```python +transferqueue_bridge.make_sample_key( + global_step, + replica_rank, + request_id, +) +``` + +它生成: + +```text +speco::: +``` + +standalone Producer 必须从 `verl_speco.transport.drafter_sample_protocol` import `make_sample_key`,不能使用 bridge 中的 PR #48 旧函数。 + +## 9. Tag 如何生成 + +```python +tag = make_ready_tag(meta) +``` + +输出: + +```python +{ + "record_type": "sample", + "status": "ready", + "schema_version": 1, + "run_id": "codex-batch-smoke", + "sequence_no": 0, + "sample_id": "smoke-0000", + "algorithm": "DSPARK", +} +``` + +tag 只包含发现、过滤和排序需要的标量/字符串。`list_samples()` 只返回 key/tag,不读取 hidden states。 + +## 10. `encode_sample()` 如何生成 fields + +调用: + +```python +fields = encode_sample(sample, meta) +``` + +### 10.1 校验 + +执行: + +```text +SampleMetadata.validate() +DraftFeatureSample.validate(strict=True) +``` + +随后检查: + +1. hidden states 是一个 dense tensor; +2. ids/mask/position 长度等于 `feature_length`; +3. hidden 第一维等于 `feature_length`; +4. hidden shape 等于 metadata; +5. hidden dtype 等于 metadata; +6. feature window 长度正确。 + +### 10.2 Tensor 规范化 + +```text +input_ids → CPU contiguous int64[L] +loss_mask → CPU contiguous float32[L] +position_ids → CPU contiguous int64[L] +hidden_states → CPU contiguous,保持模型dtype +``` + +没有 `position_ids` 时生成 `torch.arange(L, dtype=int64)`。 + +### 10.3 Metadata JSON 编码 + +```text +SampleMetadata dataclass +→ dict +→ JSON UTF-8 bytes +→ torch.uint8[M] +``` + +实现等价于: + +```python +raw = json.dumps(metadata).encode("utf-8") +metadata_json = torch.tensor(list(raw), dtype=torch.uint8) +``` + +### 10.4 最终 fields + +```python +fields = { + "input_ids": int64[3], + "loss_mask": float32[3], + "position_ids": int64[3], + "hidden_states": float32[3,4], + "metadata_json": uint8[M], +} +``` + +如果存在,协议也保留 `last_hidden_states`、`target` 和 `target_logprobs`。 + +## 11. Bridge 如何写入 TQ + +调用: + +```python +put_sample(key, fields, tag=tag) +``` + +bridge 执行: + +1. 检查 TQ 已启用; +2. 丢弃 fields 中非 tensor 值; +3. 确保本进程已经 `tq.init()`; +4. 取得配置中的 partition; +5. 调用: + +```python +tq.kv_put( + key=key, + partition_id="speco_drafter_features", + fields=fields, + tag=tag, +) +``` + +使用 MooncakeStore 时,大 tensor 路径是: + +```text +Producer CPU tensor +→ Producer TQ Client +→ MooncakeStore +``` + +不是 Ray `ObjectRef`。 + +## 12. Consumer 如何发现 key + +调用: + +```python +records = list_samples() +``` + +内部调用: + +```python +tq.kv_list(partition_id="speco_drafter_features") +``` + +标准化返回类型: + +```python +dict[str, dict[str, Any]] +``` + +示例: + +```python +{ + "drafter:v1:...:smoke-0000": { + "record_type": "sample", + "status": "ready", + "run_id": "codex-batch-smoke", + "sequence_no": 0, + ... + } +} +``` + +bridge 兼容 `key → tag` 和 `partition → key → tag` 两种 wrapper。这一步不读取 fields。 + +## 13. Consumer 如何批量取样本 + +输入: + +```python +keys = [key0, key1] +``` + +调用: + +```python +records = get_samples(keys) +``` + +bridge 只调用一次: + +```python +result = tq.kv_batch_get( + keys=keys, + partition_id="speco_drafter_features", +) +``` + +TQ 0.1.7 返回带 batch 维的 TensorDict。bridge 检查 `result.batch_size`,再执行: + +```python +rows = [result[index] for index in range(len(keys))] +``` + +每行转成普通 dict,最终返回: + +```python +[ + (key0, fields0), + (key1, fields1), +] +``` + +返回顺序与输入 keys 一致。重复 key 会提前报错。 + +## 14. `decode_sample()` 如何恢复训练对象 + +调用: + +```python +sample = decode_sample( + key, + tag, + fields, + expected_config, +) +``` + +### 14.1 Metadata 解码 + +```text +metadata_json uint8[M] +→ bytes +→ UTF-8 +→ json.loads +→ dict +→ SampleMetadata.from_dict +``` + +### 14.2 身份一致性 + +代码根据 metadata 重新生成 key,要求输入 key 完全相等;然后逐项校验 tag 的: + +```text +record_type/status/schema_version/run_id/sequence_no/sample_id/algorithm +``` + +所以 key、tag 和 payload metadata 不能来自不同样本。 + +### 14.3 Consumer 合同 + +Consumer 提供: + +```python +ExpectedFeatureConfig( + run_id="codex-batch-smoke", + schema_version=1, + algorithm="DSPARK", + target_model_id="smoke-target", + target_model_revision=None, + tokenizer_fingerprint=None, + target_layer_ids=None, + hidden_states_layout="dflash_aux", + hidden_dtype="float32", +) +``` + +值为 `None` 的字段不检查;其他字段必须完全一致。正式 Consumer 应填写 target checkpoint、tokenizer、layers、layout 和 dtype,避免使用错误 target 特征。 + +### 14.4 输出 + +完成 tensor 类型、长度、shape、dtype 校验后,构造: + +```python +DraftFeatureSample.from_dict(payload, strict=True) +``` + +输出可以交给现有 `trainer.prepare_training_batch_from_samples()`。当前 smoke test验证到这里,正式 `TQFeatureDataLoader` 尚未实现。 + +## 15. EOS 控制记录 + +调用: + +```python +key, fields, tag = make_eos_record(run_id, total_samples) +``` + +输出: + +```python +key = "control:v1::eos" +fields = {"marker": torch.tensor([1], dtype=torch.uint8)} +tag = { + "record_type": "control", + "status": "eos", + "schema_version": 1, + "run_id": run_id, + "total_samples": total_samples, +} +``` + +EOS 表示不会再发布新样本。Consumer 应先 drain ready samples 再退出。协议函数已实现,正式 Producer/Consumer 尚未调用。 + +## 16. Clear 和数据生命周期 + +bridge 提供: + +```python +clear_samples(keys) +``` + +内部调用: + +```python +tq.kv_clear(keys=keys, partition_id="speco_drafter_features") +``` + +基础层只执行删除,不决定删除时机。正式 Consumer 必须遵守: + +```text +rank 0选择global keys +→ 各rank读取local keys +→ 所有rank完成同一optimizer step +→ 汇总global success +→ rank 0 clear global keys +``` + +不能在 `get_samples()` 后立即 clear,因为 get 成功不代表训练 step 成功。 + +## 17. Client close 和 Owner close + +### 17.1 Client close + +Producer/rank 调用: + +```python +close_transfer_queue_client() +``` + +执行: + +```text +tq.get_client() +→ 当前进程client.close() +→ 如果bridge负责ray.init,则ray.shutdown() +``` + +它不调用全局 `tq.close()`,不会 kill Controller。调用后本进程不能继续使用 TQ。 + +### 17.2 Owner close + +Owner 调用: + +```python +close_transfer_queue_owner() +``` + +执行: + +```text +tq.close() +→ Controller/Storage全局清理 +→ ray.shutdown() +``` + +Owner 若误用 Client close,bridge 会抛 `RuntimeError`。 + +## 18. PR #48 兼容边界 + +bridge 继续保留: + +```python +init_transfer_queue(config) +get_sample(key) +close_transfer_queue() +``` + +PR #48 的 `SpecoTaskRunner` 已是 Ray actor,所以不调用 standalone 的 `connect_ray_cluster()`。旧 worker 继续使用单 key `get_sample()`。 + +standalone 后续使用新增的: + +```python +list_samples() +get_samples(keys) +clear_samples(keys) +``` + +因此没有修改 PR #48 现有调用点的函数签名。 + +## 19. 依赖和命令入口 + +`pyproject.toml` 新增: + +```toml +[project.optional-dependencies] +transfer-queue = ["TransferQueue==0.1.7"] +``` + +安装: + +```bash +pip install -e ".[transfer-queue]" +``` + +TQ 0.1.7 会安装 Ray,并要求 `numpy<2.0.0`。若现有包要求 NumPy 2,需要隔离环境或重新解决依赖。 + +Owner 命令: + +```text +verl-speco-tq-owner +``` + +也可以使用 `examples/run_dspark_tq_owner.sh`。 + +## 20. 单元测试 + +协议测试 `tests/unit/test_drafter_sample_protocol.py` 覆盖: + +1. encode/decode round trip; +2. key 格式; +3. tag 身份冲突; +4. Consumer contract 冲突; +5. hidden shape 冲突; +6. EOS 格式。 + +bridge 测试 `tests/unit/test_transferqueue_bridge.py` 覆盖: + +1. Ray address/namespace 参数; +2. Owner 只向 TQ 传原生配置; +3. Client 无参 `tq.init()`; +4. put/list/get-many/clear; +5. batch 返回顺序; +6. Client close 不调用全局 close; +7. Owner 不能误用 Client close。 + +运行: + +```bash +python -m pytest \ + tests/unit/test_drafter_sample_protocol.py \ + tests/unit/test_transferqueue_bridge.py \ + -q +``` + +## 21. 真实双进程 smoke test + +程序: + +```text +examples/tq_connection_smoke.py +``` + +它使用真实 `TransferQueue==0.1.7`、Ray、SimpleStorage、两个独立 Python进程和两个 batch samples。 + +Owner 路径: + +```text +连接Ray +→ tq.init(full config) +→ 写sample 0和sample 1 +→ 等待client-done +→ clear done marker +→ 全局关闭 +``` + +Client 路径: + +```text +连接同一个Ray +→ tq.init() +→ kv_list发现两个key +→ 一次kv_batch_get([k0,k1]) +→ 拆成两个fields dict +→ 分别decode_sample +→ clear两个sample keys +→ 写client-done +→ 只关闭本地client +``` + +已验证输出: + +```text +OWNER_READY keys=[k0, k1] +CLIENT_OK samples=2 shape=(3, 4) +CLIENT_CLOSED_LOCAL_ONLY +OWNER_OBSERVED_CLIENT_DONE +OWNER_CLOSED +``` + +这证明: + +1. 两个普通进程能连接同一个 TQ; +2. named Controller 发现有效; +3. 0.1.7 的 `kv_list/kv_batch_get/kv_clear` 参数有效; +4. TensorDict batch 能按 key 顺序拆开; +5. 共享协议能恢复 `DraftFeatureSample`; +6. Client close 不会杀掉 Owner; +7. Owner 能最终统一关闭。 + +## 22. 当前完整路径总结 + +```text +Owner +→ ray.init(address, namespace) +→ tq.init(native config) +→ named TransferQueueController + +普通Client +→ ray.init(same address, same namespace) +→ tq.init() +→ 找到同一个Controller + +DraftFeatureSample + SampleMetadata +→ make_sample_key +→ make_ready_tag +→ encode_sample +→ fields + metadata_json tensor +→ bridge.put_sample +→ tq.kv_put +→ SimpleStorage/MooncakeStore + +Consumer/测试Client +→ bridge.list_samples +→ key + tag +→ bridge.get_samples(keys) +→ tq.kv_batch_get +→ TensorDict batch +→ 每个key对应一个fields dict +→ decode_sample +→ DraftFeatureSample + +正式训练成功后(待实现) +→ bridge.clear_samples(global_keys) + +Client退出 +→ close_transfer_queue_client + +所有业务进程退出 +→ Owner close_transfer_queue_owner +→ tq.close +→ ray.shutdown +``` + +## 23. 下一阶段接入约束 + +后续代码不能重新定义协议或直接访问 TQ 私有对象。 + +Producer 应复用: + +```text +SampleMetadata +make_sample_key +make_ready_tag +encode_sample +bridge.put_sample +make_eos_record +``` + +Consumer 应复用: + +```text +bridge.list_samples +bridge.get_samples +decode_sample +bridge.clear_samples +``` + +下一阶段需要新增: + +```text +verl_speco/trainer/tq_feature_store.py +verl_speco/trainer/tq_sample_source.py +feature_store.py 的 type=tq 分支 +draft_training_loop.py 的流式训练分支 +Producer入口、输入读取和并发vLLM文件 +``` + +这些文件应建立在本文已经实现和真实验证过的连接、协议、KV 和关闭接口之上。 diff --git a/docs/standalone_vllm_tq_dspark_training_plan.md b/docs/standalone_vllm_tq_dspark_training_plan.md new file mode 100644 index 00000000..6644c6de --- /dev/null +++ b/docs/standalone_vllm_tq_dspark_training_plan.md @@ -0,0 +1,1131 @@ +# 独立 vLLM Producer + TQ + DSpark Consumer 第一版方案 + +## 1. 第一版要实现什么 + +只实现下面这条主链路: + +```text +包含 prompt + 预生成 response 的输入文件 +→ Producer 并发请求 vLLM prefill +→ Producer 将每条训练样本写入 TQ +→ Consumer 从同一个 TQ 取样本 +→ 独立 torchrun/FSDP DSpark 训练 +→ 一个 optimizer step 成功后删除该 step 的 TQ 样本 +``` + +第一版允许 **TQ 基础设施内部使用 Ray**,因为实测 `TransferQueue==0.1.7` 通过 Ray named actor 发现 `TransferQueueController`。Producer 和 Consumer 仍是普通 OS 进程,不改成 Ray actor;二者不使用 Ray RPC、Ray `ObjectRef` 或 Ray object store 传输训练样本。hidden-state payload 仍通过 TQ 的 `kv_put/kv_batch_get` 和配置的 MooncakeStore backend 传输。第一版不使用 DataProto/WorkerGroup,不生成长期 hidden-state feature store,也暂不实现复杂重试、自动重启和严格 checkpoint 数据恢复。 + +需要运行的组件: + +| 组件 | 数量 | 作用 | +|---|---:|---| +| vLLM server | 一个或多个 | 加载 target model,执行并行 prefill | +| Ray head | 1 个集群 | 保存 TQ named Controller/Storage actors;不承载 Producer/Consumer 业务 RPC | +| TQ owner | 1 个普通进程 | 连接 Ray,创建并保持 TQ controller/storage,最后统一关闭 TQ | +| Producer | 1 个进程 | 读文件、并发请求 vLLM、写 TQ | +| Consumer | 1 个 torchrun 任务 | 多个 rank 从 TQ 取数并训练 DSpark | + +Producer 和 Consumer 不互相调用,也不通过 HTTP 或 Ray 传样本。二者先连接同一个 Ray 集群,再通过无参 `tq.init()` 找到同一个 TQ,最后使用 TQ KV API 读写样本。 + +## 2. 共同的数据约定 + +这部分由两位开发者共同完成并先合入。建议文件: + +```text +verl_speco/transport/drafter_sample_protocol.py +tests/unit/test_drafter_sample_protocol.py +``` + +### 2.1 一个 key 对应一条样本 + +第一版固定: + +```text +一个输入文件 record +→ 一个 sequence_no +→ 一个 sample_id +→ 一个 TQ sample_key +→ 一个单样本 payload +``` + +`sequence_no` 是输入文件中的样本序号,不是 batch 编号。多个 sample keys 到 Consumer 后才组成训练 batch。 + +例如: + +```python +run_id = "dspark-20260818-a" +sequence_no = 17 +sample_id = "train-000017" + +partition_id = "speco_drafter_features" # 第一版沿用 PR 固定值 +sample_key = ( + "drafter:v1:dspark-20260818-a:" + "000000000017:train-000017" +) +``` + +### 2.2 Partition、key、tag 和 payload 的关系 + +TQ 中逻辑上是: + +```text +TQ 实例 +└── partition_id + └── sample_key + ├── tag + └── fields/payload +``` + +- TQ 实例:Producer 和 Consumer 共同连接的 controller/storage; +- `partition_id`:第一版沿用 PR 固定的 `speco_drafter_features`;每次启动新的 TQ owner,保证该 TQ 实例初始为空; +- `sample_key`:该分区中一条训练样本的地址; +- `tag`:轻量索引,Consumer 用 `kv_list()` 看见; +- `fields/payload`:真正的 Tensor 数据,Consumer 用 `kv_batch_get()` 读取。 + +Producer 写入: + +```python +tq.kv_put( + partition_id=partition_id, + key=sample_key, + fields=fields, + tag=tag, +) +``` + +Consumer 先发现 key: + +```python +all_records = tq.kv_list() +tags_by_key = all_records[partition_id] +``` + +这一步只拿 key 和 tag,不搬运 hidden states。 + +Consumer 再取数据: + +```python +result = tq.kv_batch_get( + partition_id=partition_id, + keys=selected_keys, +) +``` + +`kv_batch_get` 中的 batch 表示“一次读取多个独立 sample keys”,不是这些样本在 Producer 写入时就属于同一个对象。 + +### 2.3 Payload 字段 + +一个 sample key 对应的 `fields`: + +```python +fields = { + "input_ids": input_ids, # CPU int64[L] + "loss_mask": loss_mask, # CPU float32[L] + "position_ids": position_ids, # CPU int64[L] + "hidden_states": hidden_states, # CPU bf16[L,D] + "metadata_json": metadata_bytes, # CPU uint8[M] +} +``` + +| field | 含义 | Consumer 中的用途 | +|---|---|---| +| `input_ids` | 经过 feature window 选择后的 token IDs | DSpark token 输入 | +| `loss_mask` | 每个 token 是否参与 loss | 排除 prompt/padding/无效位置 | +| `position_ids` | 每个 row 对应的序列位置 | 位置编码和对齐校验 | +| `hidden_states` | target 指定层 hidden 沿最后一维拼接 | DSpark target/context 特征 | +| `metadata_json` | 模型、layers、layout、shape、样本身份 | Consumer 校验并恢复 metadata | + +符号: + +- `L`:这条训练 feature 保留的 token row 数; +- `H`:target model hidden size; +- `C`:DSpark context layer 数; +- L1 关闭:`D=C*H`,layout=`dflash_aux`; +- L1 开启:`D=C*H+H`,layout=`dflash_aux_plus_last`。 + +示例:`H=4096,C=5,L=1536`,开启 L1: + +```python +input_ids.shape == [1536] +loss_mask.shape == [1536] +position_ids.shape == [1536] +hidden_states.shape == [1536, 24576] +``` + +`metadata_json` 使用 JSON UTF-8 编码为 Tensor,因为参考 PR 的 bridge 只把 Tensor 放入 TQ fields: + +```python +raw = json.dumps(metadata, sort_keys=True).encode("utf-8") +metadata_bytes = torch.tensor(list(raw), dtype=torch.uint8) +``` + +### 2.4 Tag 字段 + +```python +tag = { + "record_type": "sample", + "status": "ready", + "schema_version": 1, + "run_id": "dspark-20260818-a", + "sequence_no": 17, + "sample_id": "train-000017", + "algorithm": "DSPARK", +} +``` + +tag 只存用于发现和筛选的 flat scalar/string。Consumer 只选择: + +```text +record_type=sample +status=ready +schema_version=1 +run_id=当前 run +algorithm=DSPARK +``` + +### 2.5 Metadata 字段 + +`metadata_json` 解码后至少包含: + +```python +metadata = { + "schema_version": 1, + "run_id": "dspark-20260818-a", + "sample_id": "train-000017", + "sequence_no": 17, + "algorithm": "DSPARK", + "target_model_id": "/models/Qwen3-8B", + "target_model_revision": "revision-or-checksum", + "tokenizer_fingerprint": "sha256:...", + "target_layer_ids": [2, 8, 14, 20, 26, -1], + "hidden_states_layout": "dflash_aux_plus_last", + "hidden_dtype": "bfloat16", + "hidden_shape": [1536, 24576], + "feature_length": 1536, + "full_sequence_length": 1800, + "feature_start": 264, + "feature_end": 1800, + "use_logits": False, +} +``` + +其中: + +- `target_model_id/revision`:hidden states 来自哪个 target checkpoint; +- `tokenizer_fingerprint`:Producer 使用的 tokenizer/template 版本; +- `target_layer_ids`:vLLM 返回和参与拼接的层; +- `hidden_states_layout`:Consumer 应如何拆 hidden 最后一维; +- `feature_length`:payload 中四个主要 Tensor 的第一维; +- `full_sequence_length`:完整 prompt+response 的 token 数; +- `[feature_start,feature_end)`:feature 在完整序列中的范围。 + +### 2.6 共享协议接口 + +Producer 和 Consumer 都 import 同一个模块,不各自手写字段名: + +```python +@dataclass(frozen=True) +class SampleMetadata: + schema_version: int + run_id: str + sample_id: str + sequence_no: int + algorithm: str + target_model_id: str + target_model_revision: str + tokenizer_fingerprint: str + target_layer_ids: list[int] + hidden_states_layout: str + hidden_dtype: str + hidden_shape: list[int] + feature_length: int + full_sequence_length: int + feature_start: int + feature_end: int + use_logits: bool + +def make_sample_key(meta: SampleMetadata) -> str: ... +def make_ready_tag(meta: SampleMetadata) -> dict: ... +def encode_sample(sample, meta: SampleMetadata) -> dict[str, Tensor]: ... +def decode_sample(key, tag, fields, expected_config) -> DraftFeatureSample: ... +def make_eos_record(run_id: str, total_samples: int): ... +``` + +Producer 使用 `SampleMetadata/make_sample_key/make_ready_tag/encode_sample`;Consumer 使用 `decode_sample`。 + +`SampleMetadata` Python 对象本身不经过 TQ: + +```text +Producer SampleMetadata +→ JSON +→ uint8 Tensor +→ TQ metadata_json +→ uint8 Tensor +→ JSON +→ Consumer metadata dict +``` + +`decode_sample()` 负责: + +1. 解码 `metadata_json`; +2. 校验 key、tag、metadata 中的 sample 身份一致; +3. 校验模型、tokenizer、layers 和 layout 与 Consumer 配置一致; +4. 校验 Tensor 必需字段、dtype 和 shape; +5. 返回现有 `DraftFeatureSample`。 + +### 2.7 EOS + +Producer 完成全部输入后写一个控制 record: + +```python +eos_key = f"control:v1:{run_id}:eos" +eos_fields = {"marker": torch.tensor([1], dtype=torch.uint8)} +eos_tag = { + "record_type": "control", + "status": "eos", + "schema_version": 1, + "run_id": run_id, + "total_samples": total_samples, +} +``` + +EOS 不进入训练 batch。Consumer 看到 EOS 后继续处理剩余 ready samples;`EOS 已出现且 ready 为空` 时结束。 + +## 3. TQ 怎么启动,Producer 和 Consumer 怎么连接同一个 TQ + +### 3.1 已验证的 TQ 0.1.7 连接机制 + +`TransferQueue==0.1.7` 没有提供“把 Controller 地址直接传给第二个进程”的高层连接接口。它的无参 `tq.init()` 内部执行: + +```python +_TQ_CONTROLLER = ray.get_actor("TransferQueueController") +conf = ray.get(_TQ_CONTROLLER.get_config.remote()) +_maybe_create_tq_client(conf) +``` + +因此所有进程必须先加入同一个 Ray 集群。Ray 在本方案中只承担 TQ 控制面:保存 named `TransferQueueController`、返回 TQ 配置、管理 TQ 创建的 actor。业务进程不通过 Ray 发送 sample key 或 hidden-state tensor。 + +实际连接链路是: + +```text +TQ owner:ray.init(address) → tq.init(full_tq_config) → 创建 named Controller +Producer:ray.init(address) → tq.init() → ray.get_actor() → 创建本地 TQ client +Consumer rank 0..N:ray.init(address) → tq.init() → ray.get_actor() → 创建各自 TQ client +``` + +### 3.2 直接移植并扩展 PR #48 的 bridge + +参考文件: + +```text +C:/Users/xxyyrr/Desktop/上班/verl-SpeCo/ + verl_speco/integration/transferqueue_bridge.py +``` + +第一版不新增 `standalone_tq.py`,直接把该文件移植到目标项目并扩展。保留 PR 已有的 import 检查、进程内幂等状态、`kv_put`、`kv_batch_get`、TensorDict 解包和 owner-only shutdown;新增 Ray 显式连接、批量 list/get/clear 和本地 client 关闭接口。 + +目标文件必须实现以下函数,而不只是提供一个笼统的 transport class: + +```python +def configure_transfer_queue(config: Mapping[str, Any]) -> bool: ... +def connect_ray_cluster(ray_address: str, namespace: str | None = None) -> None: ... +def start_transfer_queue_owner(tq_config: Mapping[str, Any]) -> None: ... +def connect_transfer_queue_client() -> None: ... +def make_sample_key(run_id: str, sequence_no: int, sample_id: str) -> str: ... +def put_sample(key: str, fields: dict[str, Tensor], tag: dict[str, Any]) -> None: ... +def list_samples() -> dict[str, dict[str, Any]]: ... +def get_samples(keys: list[str]) -> list[tuple[str, dict[str, Tensor]]]: ... +def clear_samples(keys: list[str]) -> None: ... +def close_transfer_queue_client() -> None: ... +def close_transfer_queue_owner() -> None: ... +``` + +逐个函数的责任如下。 + +#### `configure_transfer_queue(config)` + +- 从 Hydra/OmegaConf 中读取 Ray address、可选 namespace、固定 partition、run ID、schema version 和 TQ native backend 配置; +- 转成普通 Python dict,保存在进程内 `_state`; +- 校验 `TransferQueue==0.1.7` 可 import; +- 不连接 Ray,不创建 TQ,不产生跨进程副作用; +- 返回该进程是否启用了 TQ。 + +#### `connect_ray_cluster(ray_address, namespace)` + +- 若 `ray.is_initialized()` 为 false,调用 `ray.init(address=ray_address, namespace=namespace)`; +- 若已经初始化,校验现有 Ray context 指向期望集群,不能悄悄连到另一个本地 Ray; +- Owner、Producer 和所有 torchrun ranks 都调用它; +- 此函数只建立当前 OS 进程到 Ray control plane 的连接。 + +#### `start_transfer_queue_owner(tq_config)` + +- 仅由 `tq_owner.py` 调用; +- 前置条件是 `connect_ray_cluster()` 已成功; +- 调用一次 `tq.init(OmegaConf.create(tq_config))`; +- 将 `_state.owner=True`、`_state.initialized=True`; +- TQ 0.1.7 会创建名为 `TransferQueueController` 的 Ray actor,并创建所选 storage backend; +- 重复调用必须报错,不能启动第二套同名 Controller。 + +#### `connect_transfer_queue_client()` + +- 由 Producer 和每个 Consumer rank 调用; +- 前置条件是当前进程已经连接 Ray; +- 调用无参 `tq.init()`,通过 `ray.get_actor("TransferQueueController")` 发现 owner; +- 只创建当前进程的 TQ client,不创建新的 Controller; +- 成功后设置 `_state.initialized=True`;重复调用直接返回。 + +#### `put_sample/list_samples/get_samples/clear_samples` + +- 全部固定使用 `_SPECO_TQ_PARTITION = "speco_drafter_features"`; +- `put_sample()` 调用单样本 `tq.kv_put()`; +- `list_samples()` 调用 `tq.kv_list(partition_id=...)`,只返回 key/tag 元数据; +- `get_samples()` 一次调用 `tq.kv_batch_get(keys=...)`,再按输入 key 顺序解包为普通 dict; +- `clear_samples()` 调用 `tq.kv_clear(keys=..., partition_id=...)`; +- 这些函数不调用 Ray RPC,不把 tensor 放入 Ray object store。 + +#### `close_transfer_queue_client()` 与 `close_transfer_queue_owner()` + +TQ 0.1.7 的公共 `tq.close()` 会 kill 共享 Controller,所以两者不能写成同一个实现: + +- `close_transfer_queue_client()`:通过 0.1.7 已公开的 `tq.get_client()` 取得当前进程 client并调用它的 `close()`,然后 `ray.shutdown()`;绝不能调用会 kill Controller 的全局 `tq.close()`。这个 client close 只用于进程退出阶段,关闭后本进程不能再次调用 TQ; +- `close_transfer_queue_owner()`:仅当 `_state.owner=True` 时调用 `tq.close()`,清理 Controller/Storage,最后 `ray.shutdown()`; +- Producer 或任一训练 rank 提前退出都不能关闭全局 TQ。 + +### 3.3 共享配置 + +Owner、Producer 和 Consumer 必须使用相同的 Ray address/namespace,并通过同一个 named Controller 获得 backend 配置。第一版 partition 固定,不按任务动态创建: + +```yaml +transfer_queue: + enable: true + package_version: "0.1.7" + ray: + address: "ray-head-node:6379" + namespace: "speco-drafter" + partition_id: "speco_drafter_features" + run_id: "dspark-20260819-a" + schema_version: 1 + backend: + storage_backend: MooncakeStore + MooncakeStore: + auto_init: false + metadata_server: "node0:50050" + master_server_address: "node0:50051" + local_hostname: "" + protocol: tcp + global_segment_size: 4294967296 + local_buffer_size: 1073741824 + device_name: "" +``` + +`partition_id/run_id/schema_version` 用于过滤和校验样本;它们不能帮助进程发现 TQ。真正让三个任务连接到同一 TQ 的是“连接同一个 Ray 集群和 namespace,然后找到同名 Controller”。 + +依赖也必须锁定并单独验证。实机安装 `TransferQueue==0.1.7` 会安装 Ray,并要求 `numpy<2.0.0`;当前测试环境中的 `twinkle-kit` 要求 `numpy>=2.0.0`,两者冲突。开发时应使用专门的 TQ/训练环境或重新确认整套依赖约束,不能直接把 0.1.7 安装进已有生产环境后忽略 resolver warning。 + +### 3.4 `tq_owner.py` 要实现的入口和函数 + +新增: + +```text +verl_speco/tq_owner.py +examples/run_dspark_tq_owner.sh +``` + +`tq_owner.py` 建议明确实现: + +```python +def install_signal_handlers(stop_event: threading.Event) -> None: ... +def publish_owner_ready(run_id: str, schema_version: int) -> None: ... +def wait_until_stopped(stop_event: threading.Event) -> None: ... +def run_owner(config: DictConfig) -> int: ... +def main() -> None: ... +``` + +`run_owner()` 的执行顺序必须是: + +```text +configure_transfer_queue(config) +→ connect_ray_cluster(ray.address, ray.namespace) +→ start_transfer_queue_owner(full TQ native config) +→ put owner_ready 控制 record +→ 安装 SIGINT/SIGTERM handler +→ 保持 owner 进程存活 +→ 收到停止信号 +→ close_transfer_queue_owner() +``` + +Owner 必须常驻。TQ 0.1.7 创建 Controller 时没有设置 `lifetime="detached"`,不能在初始化后立即退出。 + +### 3.5 启动和关闭顺序 + +第一版由外部脚本管理全生命周期: + +```text +1. ray start --head,记录 Ray address +2. 启动 Mooncake metadata/master(若 auto_init=false) +3. 启动 TQ owner;owner 连接 Ray并调用 tq.init(full config) +4. 等待 owner_ready +5. 启动一个或多个 vLLM servers +6. 启动 Consumer;每个 torchrun rank 连接 Ray,然后 tq.init() +7. 启动 Producer;连接 Ray,然后 tq.init() +8. Producer 写 EOS,关闭本地 client并退出 +9. Consumer drain、保存 final checkpoint,所有 ranks 关闭本地 client并退出 +10. 给 TQ owner 发送 SIGTERM;仅 owner 执行 tq.close() +11. 等 owner 退出后执行 ray stop +12. 停止 Mooncake 服务 +``` + +外部脚本要用 `trap` 保证异常退出也按“Producer/Consumer → owner → Ray → Mooncake”的顺序清理。不能在 Consumer rank 的 `finally` 中调用全局 `tq.close()`。 + +## 4. Producer 要实现什么 + +### 4.1 Producer 完整顺序 + +```text +读取共享配置 +→ 连接 TQ并校验 owner_ready +→ 初始化 tokenizer +→ 初始化多个 vLLM endpoint clients +→ 流式读取输入文件 +→ 为每条输入分配 sequence_no/sample_id +→ 拼接 prompt+预生成 response,得到 input_ids/loss_mask +→ 并发请求 vLLM prefill +→ 读取 vLLM hidden-state 临时结果 +→ 转换成 DSpark DraftFeatureSample +→ 构造 SampleMetadata +→ encode_sample 得到 fields/tag/key +→ TQ kv_put 一条 sample +→ 删除该请求临时文件 +→ 所有输入完成后写 EOS +→ close_transfer_queue_client()并退出 +``` + +### 4.2 并发模型 + +Producer 是一个进程,内部并发请求多个 endpoint: + +```text +InputReader +→ bounded asyncio input_queue +→ N 个 RequestWorker +→ bounded publish_queue +→ TQ Publisher +``` + +- `vllm_endpoints` 是列表; +- 每个 endpoint 有独立 semaphore; +- 总并发由 `max_inflight_requests` 限制; +- input/publish queue 必须有上限; +- TQ `kv_put` 如果是同步 API,用 `asyncio.to_thread()` 调用; +- TQ ready 数达到 `max_pending_samples` 时暂停继续请求,形成简单背压。 + +`sequence_no` 在 InputReader 中分配,不按 vLLM 完成顺序分配。因此并发乱序不会改变 sample key。 + +### 4.3 vLLM 结果转换 + +复用现有 `TargetFeatureReplayer` 的: + +- OpenAI-compatible vLLM 请求; +- `prompt_token_ids` 校验; +- `kv_transfer_params.hidden_states_path`; +- safetensors 加载; +- `[seq,layers,hidden]` 校验; +- feature positions 选择; +- aux layers flatten; +- DSpark L1 时拼 final hidden。 + +不要复制 `_feature_from_vllm_payload()`;把纯转换逻辑提取成公共函数,让旧 file backend 和新 Producer 共同调用。 + +临时文件顺序: + +```text +加载 +→ 校验/转换 +→ TQ put 成功 +→ 删除 +``` + +第一版仍允许每个并发请求产生短期临时文件,但不生成全量 hidden-state 数据集。 + +### 4.4 Producer 文件分工 + +| 文件 | 实现内容 | +|---|---| +| `verl_speco/standalone_tq_producer.py` | CLI/Hydra 入口,连接 TQ,启动 asyncio pipeline,写 EOS和汇总指标 | +| `verl_speco/producer/input_reader.py` | 流式读输入、分配 `sequence_no/sample_id`、tokenize、构造 loss mask | +| `verl_speco/producer/vllm_feature_client.py` | endpoint pool、并发控制、HTTP 请求、临时结果读取和删除 | +| `verl_speco/trainer/target_feature_replay.py` | 提取可复用的 vLLM payload → `DraftFeatureSample` 纯转换函数 | +| `examples/run_dspark_tq_producer.sh` | Producer 配置和启动命令 | +| `tests/unit/test_drafter_sample_protocol.py` | 协议编码、字段和 shape 测试(两人共同) | +| `tests/integration/test_tq_producer_smoke.py` | fake/短 vLLM → TQ sample + EOS | + +Producer 开发者同时负责移植/扩展 `integration/transferqueue_bridge.py` 和实现 `tq_owner.py`,因为这部分与 TQ 写入和连接直接相关。 + +### 4.5 Producer 各文件的函数级实现规格 + +#### `verl_speco/standalone_tq_producer.py` + +需要实现: + +```python +@dataclass +class ProducerStats: + input_count: int + published_count: int + failed_count: int + pending_bytes: int + +async def publish_one(result: PreparedFeature, transport) -> str: ... +async def run_producer(config: DictConfig) -> ProducerStats: ... +def validate_producer_config(config: DictConfig) -> None: ... +def main() -> None: ... +``` + +`main()` 只负责 Hydra/日志/退出码。`run_producer()` 是可测试的业务入口,执行:连接 Ray → 连接 TQ client → 创建 InputReader 和 vLLM client pool → 启动有界 asyncio pipeline → 等待所有请求及 `kv_put` 完成 → 发布 EOS → 关闭本地 client。它不能启动 Ray head、不能创建 TQ owner、不能调用全局 `tq.close()`。 + +`publish_one()` 接收已经完成转换的一条 `PreparedFeature`,调用共享协议的 `encode_sample()` 得到 `(key, fields, tag)`,再通过 bridge 的 `put_sample()` 发布。只有 `put_sample()` 成功返回,`published_count` 才增加,vLLM 临时文件才允许删除。 + +#### `verl_speco/producer/input_reader.py` + +需要实现: + +```python +@dataclass(frozen=True) +class InputRecord: + sequence_no: int + sample_id: str + prompt: str + response: str + source_metadata: dict[str, Any] + +def iter_input_records(path: str) -> Iterator[InputRecord]: ... +def tokenize_record(record: InputRecord, tokenizer, config) -> TokenizedRequest: ... +def build_loss_mask(input_ids: Tensor, prompt_length: int) -> Tensor: ... +``` + +`iter_input_records()` 流式读取,不把全文件载入内存;在这里按文件顺序分配稳定的 `sequence_no`。`tokenize_record()` 拼接已经存在的 prompt/response,不调用模型生成 response;输出至少包含 `input_ids:int64[L]`、`position_ids:int64[L]`、`loss_mask:float32[L]` 和请求 vLLM 所需字段。 + +#### `verl_speco/producer/vllm_feature_client.py` + +需要实现: + +```python +@dataclass(frozen=True) +class VllmEndpoint: + base_url: str + max_concurrency: int + +class VllmFeatureClientPool: + async def start(self) -> None: ... + async def prefill(self, request: TokenizedRequest) -> RawVllmFeature: ... + async def close(self) -> None: ... + +async def request_prefill(endpoint, request) -> VllmResponse: ... +def choose_endpoint(endpoints, state) -> VllmEndpoint: ... +def load_hidden_state_result(response) -> RawVllmFeature: ... +def delete_temporary_result(raw: RawVllmFeature) -> None: ... +``` + +`prefill()` 必须允许多个 coroutine 同时运行;全局 semaphore 限制总并发,每个 endpoint 另有独立 semaphore。`request_prefill()` 只负责 HTTP 请求和响应校验;`load_hidden_state_result()` 负责读取 vLLM 返回的临时 safetensors/path。临时结果的删除不放在 `load_hidden_state_result()`,而由 `publish_one()` 成功后触发。 + +#### `verl_speco/trainer/target_feature_replay.py` + +把当前类内部的纯转换部分抽成: + +```python +def feature_from_vllm_payload( + payload: RawVllmFeature, + request: TokenizedRequest, + feature_config: FeatureContract, +) -> DraftFeatureSample: ... +``` + +它不进行 HTTP、不访问 TQ、不删除文件,只完成 shape/layout/layer 校验、feature-position 选择、aux layer flatten 和可选 L1 final hidden 拼接。现有 file replay backend 和新 Producer 都调用这一函数,避免两套转换规则。 + +#### `examples/run_dspark_tq_producer.sh` + +负责提供同一套: + +```text +RAY_ADDRESS / Ray namespace +run_id / schema_version / 固定 partition +Mooncake/TQ backend 配置 +输入文件和 tokenizer/model 配置 +vLLM endpoint 列表 +max_inflight_requests / per_endpoint_concurrency +``` + +脚本只启动 Producer,不启动 TQ owner 或 Consumer,便于两位开发者独立调试。 + +## 5. Consumer 要实现什么 + +### 5.1 不新写另一套训练器 + +继续使用现有入口: + +```text +draft_train_launcher.py +→ draft_train.py +→ trainer/draft_training_loop.py +→ DrafterBaseTrainer +→ DSparkTrainerBackend +``` + +训练模式继续使用现有 standalone `training.mode=offline`,用 `feature_store.type=tq` 选择 TQ 数据源。不复制 FSDP、loss、optimizer 或 checkpoint 代码。 + +当前 `feature_store.type` 的作用是告诉 `build_feature_store_from_config()` 应创建哪一种训练数据来源: + +| 当前 type | 对象 | 数据来源 | +|---|---|---| +| `torch_shard` | `TorchShardFeatureStore` | 本地 `.pt` shard,当前 separate-training 示例使用它 | +| `token_replay` | `TokenReplayFeatureStore` | 保存 token replay 的 shard | +| `vllm_safetensors` / `safetensors` | `VllmSafetensorsFeatureStore` | 已生成的 safetensors feature | +| `jsonl_token_replay` / `jsonl` | `JsonlTokenReplayFeatureStore` | JSONL token/text replay | + +第一版新增: + +```yaml +actor_rollout_ref: + rollout: + drafter: + training: + mode: offline + feature_store: + type: tq + path: null + shuffle: false + repeat: false +``` + +这里 `offline` 的含义是“独立于 RL 的 standalone drafter training”,不等于数据必须来自磁盘。 + +不能只给现有工厂增加一个 `type=tq` 然后继续使用通用 `DraftFeatureDataLoader`。当前 loader 会: + +```python +keys = list(store.iter_keys(...)) +``` + +它假设数据集 keys 是一个静态快照;空 store 会直接结束,并且各 rank 独立枚举时可能在 Producer 持续写入的过程中看到不同快照。TQ 是流式、会新增并删除 keys 的数据源,所以 `type=tq` 必须选择专用的 `TQFeatureDataLoader`,由 rank 0 统一发现和分配 keys。 + +#### 5.1.1 当前磁盘 feature store 为什么每个 rank 都会取 keys + +`draft_train_launcher.py` 通过 `torchrun` 启动多个训练进程。每个进程都是一个 rank,并且每个 rank 都会独立进入 `run_standalone_draft_training()`、创建 `DraftFeatureDataLoader`、执行它的 `__iter__()`。当前 loader 的核心逻辑是: + +```python +keys = list( + self.store.iter_keys( + shuffle=self.shuffle, + seed=self.seed + epoch, + ) +) +rank_keys = keys[rank::world_size] + +for key in rank_keys: + batch.append(self.store.read(key)) +``` + +因此当前并不是 rank 0 枚举 keys 后再发送给其他 rank,而是所有 rank 都访问相同的静态 feature-store 路径: + +```text +rank 0:iter_keys() 得到完整静态列表 → keys[0::world_size] → read 自己的 samples +rank 1:iter_keys() 得到完整静态列表 → keys[1::world_size] → read 自己的 samples +... +``` + +例如 store 中固定存在: + +```python +keys = ["k0", "k1", "k2", "k3", "k4", "k5", "k6", "k7"] +``` + +当 `world_size=2` 时,两个 rank 都先得到上述完整列表,然后分别计算: + +```python +# rank 0 +rank_keys = keys[0::2] # ["k0", "k2", "k4", "k6"] + +# rank 1 +rank_keys = keys[1::2] # ["k1", "k3", "k5", "k7"] +``` + +这种实现能够成立,是因为 `.pt`、safetensors 或 replay 文件在训练期间是静态数据集。只要所有 rank 使用同一路径和相同 shuffle seed,`iter_keys()` 就会产生相同顺序的 key 快照,各 rank 可以无通信地算出互不重叠的子集。 + +#### 5.1.2 TQ 为什么必须改成 rank 0 发现 keys + +TQ 中的 keys 会在训练期间动态变化:Producer 持续 `kv_put`,训练成功后 Consumer 又执行 `kv_clear`。如果所有 rank 仍然各自调用 `kv_list`,不同调用时刻可能看到不同快照: + +```python +# rank 0 较早调用 +rank0_keys = ["k0", "k1", "k2", "k3"] + +# Producer 随后写入 k4、k5,rank 1 较晚调用 +rank1_keys = ["k0", "k1", "k2", "k3", "k4", "k5"] +``` + +各 rank 再独立切片后,可能得到不同数量的 local samples,进而无法保证它们以相同顺序进入 forward、backward 和梯度 collective。为此,TQ 专用 loader 必须把控制面和数据面分开: + +```text +控制面:rank 0 执行 kv_list,固定本 step 的 global_keys,并向各 rank 分发 local_keys +数据面:每个 rank 使用自己的 local_keys 直接执行 kv_batch_get,从 TQ/Mooncake 读取 hidden-state payload +``` + +rank 0 只发送较小的 key/tag 元数据,不读取并转发其他 rank 的 hidden-state tensor。所有 rank 仍然都会连接同一个 TQ,也都会调用批量读取接口;只有动态 key 的发现和本 step 的 global batch 决策集中在 rank 0。 + +因此 `feature_store.type=tq` 不是只替换底层 `read()`:它同时改变了 key 的发现、分配、等待和删除语义,需要专用 `TQFeatureDataLoader`,并要求训练循环在所有 rank 成功完成 optimizer step 后,由 rank 0 对本 step 的 `global_keys` 执行 `kv_clear`。 + +### 5.2 Consumer 完整顺序 + +```text +torchrun 启动多个 ranks +→ 每个 rank 初始化 torch.distributed +→ 每个 rank 连接同一个 TQ +→ rank 0 校验 owner_ready,并 broadcast 结果 +→ 初始化现有 DSpark trainer +→ rank 0 kv_list 查找 ready sample keys +→ rank 0 选一个 global batch并分给各 rank +→ 每个 rank kv_batch_get 自己的 local keys +→ decode_sample 得到 list[DraftFeatureSample] +→ prepare_training_batch_from_samples() +→ training_step_from_batch() +→ 所有 rank 汇总 success +→ 成功后 rank 0 kv_clear 这个 global batch 的 keys +→ 继续下一批 +→ 看到 EOS 且 ready 为空 +→ 保存 final checkpoint +→ 所有 ranks close_transfer_queue_client()并退出 +``` + +### 5.3 多 rank 如何分 key + +例如: + +```text +world_size=2 +batch_size_per_gpu=2 +global batch size=4 +``` + +rank 0 选出: + +```python +global_keys = ["k10", "k11", "k12", "k13"] +assignments = [ + ["k10", "k11"], # rank 0 + ["k12", "k13"], # rank 1 +] +``` + +通过 `dist.scatter_object_list` 或 broadcast 发送短字符串列表。每个 rank 直接从 TQ 读取自己的 payload,不通过 rank 0 转发 hidden states。 + +### 5.4 从 TQ 到训练 batch + +每个 rank: + +```python +records = tq_transport.get_samples(local_keys) + +samples = [ + decode_sample( + key=key, + tag=tags_by_key[key], + fields=fields, + expected_config=expected_contract, + ) + for key, fields in records +] + +batch = trainer.prepare_training_batch_from_samples( + samples, + step=optimizer_step, +) + +ok = await trainer.training_step_from_batch( + batch, + optimizer_step, +) +``` + +`records` 是多个独立单样本 payload;`samples` 是变长 `DraftFeatureSample` 列表。Consumer transport 层不要直接 `torch.stack()`,现有 trainer/backend 负责对齐和组 batch。 + +### 5.5 删除与结束 + +第一版采用简单逻辑: + +```text +所有 rank get/decode/train 都成功 +→ all_reduce(global_success)=True +→ rank 0 kv_clear(global_batch_keys) +``` + +任何 rank 失败都不 clear。程序报错退出,第一版不自动恢复。 + +EOS 后不足一个 global batch 的尾部,第一版使用 `drop_last=True`:记录数量,rank 0 clear 这些尾部 keys,然后正常结束。 + +checkpoint 仍使用现有 `save_interval_steps` 和 final checkpoint;第一版不保证崩溃后已 clear 数据能够严格重放。 + +### 5.6 Consumer 文件分工 + +| 文件 | 实现内容 | +|---|---| +| `verl_speco/trainer/feature_store.py` | 工厂增加 `type=tq`,返回 `TQFeatureStore`;TQ 类型不要求 `path` | +| `verl_speco/trainer/tq_feature_store.py` | 实现固定 partition 上的 list/get-many/clear/EOS,不伪装静态 `iter_keys()` | +| `verl_speco/trainer/tq_sample_source.py` | 实现专用 `TQFeatureDataLoader`:rank 0 动态 list/filter/sort、key 分配、各 rank get/decode、EOS/tail | +| `verl_speco/trainer/draft_training_loop.py` | 保持 `mode=offline`;`type=tq` 时跳过 `feature_store.path` 必填检查并选择专用 loader;调用现有 prepare/train,global success 后 clear,final checkpoint | +| `verl_speco/draft_train_launcher.py` | `feature_store.type=tq` 时不要求 feature-store path,透传 torchrun/TQ 配置 | +| `verl_speco/config/speco_base.yaml` | 新增 TQ connection、run、poll、batch 等默认配置 | +| `examples/run_dspark_tq_consumer.sh` | Consumer GPU、batch、checkpoint 和共享 TQ 配置 | +| `tests/unit/test_tq_sample_source.py` | key 过滤/排序/分 rank、EOS、decode 调用测试 | +| `tests/integration/test_dspark_tq_consumer_smoke.py` | 两 rank 取不同样本,完成训练并 clear | + +### 5.7 Consumer 各文件的函数级实现规格 + +#### `verl_speco/trainer/feature_store.py` + +修改现有工厂: + +```python +def build_feature_store_from_config(feature_store_cfg, read_only=False): + store_type = str(feature_store_cfg.get("type", "torch_shard")).lower() + if store_type == "tq": + return TQFeatureStore.from_config(feature_store_cfg) + ... +``` + +要求: + +- 保留 `torch_shard/token_replay/vllm_safetensors/jsonl_token_replay` 的当前行为; +- `type=tq` 时不读取 `path`; +- 工厂只创建对象,不在 import 阶段连接 Ray/TQ; +- `TQFeatureStore` 是流式数据源适配器,不强行实现有误导性的静态 `iter_keys()` 和逐条 `read()`。 + +#### `verl_speco/trainer/tq_feature_store.py` + +需要实现: + +```python +@dataclass(frozen=True) +class ReadyEntry: + key: str + tag: dict[str, Any] + +class TQFeatureStore: + @classmethod + def from_config(cls, cfg) -> "TQFeatureStore": ... + def connect(self) -> None: ... + def list_ready(self, run_id: str) -> list[ReadyEntry]: ... + def get_many(self, entries: list[ReadyEntry]) -> list[DraftFeatureSample]: ... + def clear_many(self, keys: list[str]) -> None: ... + def read_eos(self, run_id: str) -> EosMetadata | None: ... + def close_local(self) -> None: ... +``` + +`connect()` 调用共享 bridge 的 `connect_ray_cluster()` 和 `connect_transfer_queue_client()`。`list_ready()` 调用 `list_samples()` 后只保留 `tag.status=ready`、匹配 `run_id/schema_version` 的数据 key,并按 `(sequence_no,key)` 排序。`get_many()` 一次批量读取 fields,然后逐条调用共享 `decode_sample()`;返回顺序必须和 entries 相同。`clear_many()` 只允许 rank 0 在全局 step 成功后调用。`close_local()` 只关闭本 rank client,不能关闭 Controller。 + +#### `verl_speco/trainer/tq_sample_source.py` + +需要实现: + +```python +@dataclass +class TQLocalBatch: + local_keys: list[str] + local_samples: list[DraftFeatureSample] + global_keys: list[str] | None + +class TQFeatureDataLoader: + def __iter__(self) -> Iterator[TQLocalBatch]: ... + def _select_global_batch(self) -> tuple[list[str], dict[str, ReadyEntry]]: ... + def _broadcast_assignments(self, assignments) -> list[ReadyEntry]: ... + def _handle_eos_and_tail(self) -> bool: ... + def clear_completed_batch(self, global_keys: list[str] | None) -> None: ... +``` + +执行责任必须明确: + +- 所有 rank 创建 loader 并调用 `store.connect()`; +- 只有 rank 0 执行 `_select_global_batch()` 和 `kv_list`; +- rank 0 生成 `list[list[ReadyEntry]]` assignments,使用 `torch.distributed.scatter_object_list` 或 broadcast 分发小型 key/tag; +- 每个 rank 对自己的 local entries 调用 `store.get_many()`,hidden states 从 TQ/Mooncake 直接进入该 rank; +- loader yield `TQLocalBatch`,不能只 yield samples,因为训练后清理还需要 `global_keys`; +- 暂时为空时轮询等待,不能像静态 loader 一样结束;只有 EOS 已出现并且 ready 数据 drain 完才停止; +- `drop_last=True` 时由 rank 0 记录并清理不足 global batch 的尾部 key。 + +#### `verl_speco/trainer/draft_training_loop.py` + +需要新增或调整: + +```python +def build_training_source(config, rank, world_size): ... +def all_ranks_succeeded(local_ok: bool, device) -> bool: ... +async def run_tq_training_loop(trainer, loader, config) -> dict[str, Any]: ... +``` + +`build_training_source()` 根据 `feature_store.type` 分支:静态类型继续创建 `DraftFeatureDataLoader`;`tq` 创建 `TQFeatureDataLoader`,跳过 `feature_store.path` 必填检查,并禁止 `shuffle/repeat`。`run_tq_training_loop()` 对每个 `TQLocalBatch` 调用现有 `prepare_training_batch_from_samples()` 和 `training_step_from_batch()`;所有 rank 通过 collective 得到 global success 后,才让 rank 0 调用 `clear_completed_batch(global_keys)`。异常路径不 clear,finally 只执行 `store.close_local()`。 + +#### `verl_speco/draft_train_launcher.py` + +保留当前 `torch.distributed.run` 启动方式,只增加配置校验和环境透传: + +```python +def validate_tq_launch_config(overrides, launch_config) -> None: ... +def build_child_env(config) -> dict[str, str]: ... +``` + +它不启动 Ray head、不调用 `tq.init()`。职责是确认 `type=tq` 时提供了 Ray address/namespace,且不要求 feature-store path;随后把同一 Ray connection 配置传给每个 torchrun 子进程。每个子 rank 自己建立 Ray/TQ client,不能在 launcher 父进程建一个 client后期待 fork 继承。 + +#### `verl_speco/config/speco_base.yaml` + +增加默认字段: + +```yaml +feature_store: + type: torch_shard + path: null + shuffle: true + repeat: true + tq: + ray_address: null + ray_namespace: speco-drafter + partition_id: speco_drafter_features + run_id: null + schema_version: 1 + poll_interval_seconds: 0.5 + connect_timeout_seconds: 120 + drop_last: true +``` + +当 `type=tq` 时运行期覆盖为 `shuffle=false/repeat=false`。backend 的完整 owner 配置只需要 TQ owner 使用;Producer/Consumer client 从 named Controller 获取它,不各自重新创建 backend。 + +#### Consumer 测试必须覆盖的函数边界 + +- `test_feature_store_factory_builds_tq_without_path()`; +- `test_rank0_filters_and_sorts_ready_entries()`; +- `test_nonzero_rank_never_calls_kv_list()`; +- `test_assignments_are_disjoint_and_global_batch_complete()`; +- `test_each_rank_gets_only_local_keys()`; +- `test_decode_preserves_hidden_states_layout()`; +- `test_clear_only_after_all_ranks_success()`; +- `test_failure_does_not_clear()`; +- `test_eos_drains_ready_then_stops()`; +- `test_client_close_does_not_kill_owner()`。 + +## 6. 两个人怎么分工 + +### 共同先完成 + +1. `drafter_sample_protocol.py`; +2. Ray/TQ connection 配置字段; +3. 一个小型 golden sample; +4. 启动一个测试 Ray head 和 TQ owner,独立进程 A put、独立进程 B list/get/clear 的 smoke test; +5. 验证 client 退出不会 kill owner,只有 owner shutdown 才销毁 Controller。 + +### 开发者 A:Producer/TQ + +负责: + +```text +integration/transferqueue_bridge.py +tq_owner.py +standalone_tq_producer.py +producer/input_reader.py +producer/vllm_feature_client.py +target_feature_replay.py 的公共转换函数 +owner/producer 启动脚本 +Producer/TQ 测试 +``` + +开发者 A 的可交付接口不是“提供一个 TQ 类”,而是: + +```text +bridge:connect_ray_cluster/start_owner/connect_client/put/list/get_many/clear/client_close/owner_close +owner:run_owner/main/signal handler/owner_ready +producer:run_producer/publish_one/统计与 EOS +input reader:iter_input_records/tokenize_record/build_loss_mask +vLLM client:endpoint pool/request_prefill/load/delete +feature conversion:feature_from_vllm_payload +``` + +开发者 B 可以先针对这些接口写 fake transport,不需要等待真实 vLLM 和 Mooncake 联通。 + +### 开发者 B:Consumer/训练 + +负责: + +```text +feature_store.py 的 type=tq 工厂分支 +tq_feature_store.py +tq_sample_source.py / TQFeatureDataLoader +draft_training_loop.py 的 offline + type=tq 分支 +draft_train_launcher.py 配置适配 +speco_base.yaml Consumer 配置 +Consumer 启动脚本 +Consumer/DSpark 测试 +``` + +开发者 B 的可交付接口是: + +```text +feature-store factory:type=tq 分支 +TQFeatureStore:connect/list_ready/get_many/clear_many/read_eos/close_local +TQFeatureDataLoader:select global batch/distribute local keys/get/yield/EOS-tail +training loop:build source/train/global success/clear/final checkpoint +launcher:TQ 配置校验和 torchrun 子进程环境透传 +``` + +### 联调入口 + +建议再提供: + +```text +examples/run_dspark_tq_pipeline_local.sh +``` + +只用于单机联调,顺序启动: + +```text +ray start --head +→ Mooncake metadata/master +→ TQ owner(ray.init + tq.init(full config)) +→ owner_ready +→ vLLM health check +→ Consumer +→ Producer +→ 等 Producer/Consumer 退出 +→ SIGTERM TQ owner(owner 执行 tq.close) +→ ray stop +→ 停止 Mooncake +``` + +最小联调:Producer 发布 8 条样本,2 个 Consumer ranks、每 rank batch size 2,完成 2 个 optimizer steps,8 个 sample keys 被清理,EOS 后保存 final checkpoint。 + +## 7. 第一版验收标准 + +1. TQ owner、Producer 和所有 Consumer ranks 日志显示相同 Ray address/namespace、TQ Controller、partition 和 run ID。 +2. 在同一 Ray 集群中,独立 owner 创建 TQ 后,独立进程 A put,独立进程 B 能 list/get/clear。 +3. Producer 对多个 vLLM endpoints 并发请求,不串行访问。 +4. 一个输入 record 只生成一个 sample key 和一个 payload。 +5. `kv_list` 只拿 tag;`kv_batch_get` 才拿 hidden states。 +6. Consumer 各 rank 读取不同 local keys,不通过 rank 0 搬运 Tensor。 +7. `decode_sample` 能拒绝模型、layer、layout、dtype 或 shape 不匹配的数据。 +8. 所有 rank 训练成功后才 clear 当前 global batch。 +9. Producer 先完成时,Consumer 能 drain 后再退出。 +10. 不产生长期 hidden-state feature store。 +11. Producer 或任一 Consumer rank 退出不会销毁 TQ Controller;只有 owner 调用全局 `tq.close()`。 +12. Ray object store 中不承载 hidden-state payload,训练 tensor 通过 TQ/Mooncake 路径读取。 + +## 8. 后续建议:第一版跑通后再做 + +以下内容不进入第一版开发: + +- Producer HTTP/TQ 复杂重试和 endpoint 熔断; +- Producer 发布 journal,避免重启后重复生成已 clear 样本; +- 自动重启;第一版失败后人工停止整条 pipeline并使用新 `run_id` 重跑; +- Consumer 从最新 checkpoint 自动恢复; +- checkpoint 成功后再 clear 的严格提交窗口; +- TQ owner/storage 整体丢失后的数据重建; +- 多个独立 Consumer 竞争同一 partition; +- lease、ack、超时回收和 exactly-once; +- 动态扩缩容; +- vLLM server 直接写 TQ。 + +第一版先保证:同一个 TQ 能连通、Producer 能并发生产、Consumer 能正确取数训练、每步成功后能及时清理。 diff --git a/examples/run_dspark_tq_owner.sh b/examples/run_dspark_tq_owner.sh new file mode 100644 index 00000000..b94d38b7 --- /dev/null +++ b/examples/run_dspark_tq_owner.sh @@ -0,0 +1,24 @@ +#!/usr/bin/env bash +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +set -euo pipefail + +: "${RAY_ADDRESS:?Set RAY_ADDRESS to the running Ray head, for example 10.0.0.1:6379}" +: "${SPECO_TQ_RUN_ID:?Set SPECO_TQ_RUN_ID to a unique pipeline run id}" + +python -m verl_speco.tq_owner \ + actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.address="${RAY_ADDRESS}" \ + actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.namespace=speco-drafter \ + actor_rollout_ref.rollout.drafter.training.transfer_queue.run_id="${SPECO_TQ_RUN_ID}" \ + actor_rollout_ref.rollout.drafter.training.transfer_queue.backend.storage_backend=SimpleStorage diff --git a/examples/tq_connection_smoke.py b/examples/tq_connection_smoke.py new file mode 100644 index 00000000..df873e76 --- /dev/null +++ b/examples/tq_connection_smoke.py @@ -0,0 +1,205 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Real two-process smoke test for Ray + TransferQueue 0.1.7. + +Start a Ray head first, then run ``owner`` and ``client`` in separate shells. +The owner publishes one protocol-valid sample; the client list/get/decodes and +clears it, publishes a done marker, and exits without killing the owner. +""" + +from __future__ import annotations + +import argparse +import time + +import torch + +from verl_speco.integration.transferqueue_bridge import ( + clear_samples, + close_transfer_queue_client, + close_transfer_queue_owner, + configure_transfer_queue, + connect_ray_cluster, + connect_transfer_queue_client, + get_samples, + list_samples, + put_sample, + start_transfer_queue_owner, +) +from verl_speco.trainer.feature_store import DraftFeatureSample +from verl_speco.transport.drafter_sample_protocol import ( + ExpectedFeatureConfig, + SampleMetadata, + decode_sample, + encode_sample, + make_ready_tag, + make_sample_key, +) + + +def _config(args) -> dict: + return { + "enable": True, + "package_version": "0.1.7", + "ray": {"address": args.ray_address, "namespace": args.namespace}, + "partition_id": "speco_drafter_features", + "run_id": args.run_id, + "schema_version": 1, + "controller": {"polling_mode": True}, + "backend": { + "storage_backend": "SimpleStorage", + "SimpleStorage": { + "total_storage_size": 32, + "num_data_storage_units": 1, + }, + }, + } + + +def _record(run_id: str, sequence_no: int): + meta = SampleMetadata( + schema_version=1, + run_id=run_id, + sample_id=f"smoke-{sequence_no:04d}", + sequence_no=sequence_no, + algorithm="DSPARK", + target_model_id="smoke-target", + target_model_revision="smoke-revision", + tokenizer_fingerprint="smoke-tokenizer", + target_layer_ids=[0], + hidden_states_layout="dflash_aux", + hidden_dtype="float32", + hidden_shape=[3, 4], + feature_length=3, + full_sequence_length=3, + feature_start=0, + feature_end=3, + use_logits=False, + ) + sample = DraftFeatureSample( + algorithm="DSPARK", + input_ids=torch.tensor([1, 2, 3]) + sequence_no, + loss_mask=torch.tensor([0.0, 1.0, 1.0]), + position_ids=torch.tensor([0, 1, 2]), + hidden_states=( + torch.arange(12, dtype=torch.float32).reshape(3, 4) + sequence_no + ), + ) + return meta, sample + + +def _wait_for(predicate, timeout: float, description: str): + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + value = predicate() + if value: + return value + time.sleep(0.2) + raise TimeoutError(f"Timed out waiting for {description}") + + +def run_owner(args) -> None: + config = _config(args) + configure_transfer_queue(config) + connect_ray_cluster(args.ray_address, args.namespace) + start_transfer_queue_owner(config) + records = [_record(args.run_id, sequence_no) for sequence_no in range(2)] + keys = [make_sample_key(meta) for meta, _ in records] + done_key = f"control:v1:{args.run_id}:smoke-client-done" + try: + for (meta, sample), key in zip(records, keys, strict=True): + put_sample(key, encode_sample(sample, meta), tag=make_ready_tag(meta)) + print(f"OWNER_READY keys={keys}", flush=True) + _wait_for( + lambda: list_samples().get(done_key), + args.timeout, + "client done marker", + ) + clear_samples([done_key]) + print("OWNER_OBSERVED_CLIENT_DONE", flush=True) + finally: + close_transfer_queue_owner() + print("OWNER_CLOSED", flush=True) + + +def run_client(args) -> None: + config = _config(args) + configure_transfer_queue(config) + connect_ray_cluster(args.ray_address, args.namespace) + connect_transfer_queue_client() + metas = [_record(args.run_id, sequence_no)[0] for sequence_no in range(2)] + keys = [make_sample_key(meta) for meta in metas] + done_key = f"control:v1:{args.run_id}:smoke-client-done" + try: + _wait_for( + lambda: all(key in list_samples() for key in keys), + args.timeout, + "owner sample", + ) + tags = list_samples() + fetched = get_samples(keys) + restored = [ + decode_sample( + key, + tags[key], + fields, + ExpectedFeatureConfig( + run_id=args.run_id, + target_model_id="smoke-target", + hidden_states_layout="dflash_aux", + hidden_dtype="float32", + ), + ) + for key, fields in fetched + ] + assert restored[0].hidden_states.tolist() == torch.arange(12).reshape(3, 4).tolist() + assert restored[1].hidden_states.tolist() == ( + torch.arange(12).reshape(3, 4) + 1 + ).tolist() + clear_samples(keys) + put_sample( + done_key, + {"marker": torch.tensor([1], dtype=torch.uint8)}, + tag={ + "record_type": "control", + "status": "client_done", + "schema_version": 1, + "run_id": args.run_id, + }, + ) + print( + f"CLIENT_OK samples={len(restored)} shape={tuple(restored[0].hidden_states.shape)}", + flush=True, + ) + finally: + close_transfer_queue_client() + print("CLIENT_CLOSED_LOCAL_ONLY", flush=True) + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("role", choices=("owner", "client")) + parser.add_argument("--ray-address", required=True) + parser.add_argument("--namespace", default="speco-drafter-smoke") + parser.add_argument("--run-id", default="tq-smoke") + parser.add_argument("--timeout", type=float, default=60.0) + args = parser.parse_args() + if args.role == "owner": + run_owner(args) + else: + run_client(args) + + +if __name__ == "__main__": + main() diff --git a/pyproject.toml b/pyproject.toml index ca8a9c7b..445e7810 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -16,6 +16,9 @@ dependencies = [ "packaging>=24", ] +[project.optional-dependencies] +transfer-queue = ["TransferQueue==0.1.7"] + [project.urls] Repository = "https://github.com/verl-project/verl-SpeCo" @@ -23,6 +26,7 @@ Repository = "https://github.com/verl-project/verl-SpeCo" verl-speco = "verl_speco.main:main" verl-speco-draft-train = "verl_speco.draft_train_launcher:main" verl-speco-inspect-features = "verl_speco.inspect_feature_store:main" +verl-speco-tq-owner = "verl_speco.tq_owner:main" [tool.setuptools.dynamic] version = { attr = "verl_speco.__version__" } diff --git a/tests/unit/test_drafter_sample_protocol.py b/tests/unit/test_drafter_sample_protocol.py new file mode 100644 index 00000000..74defffe --- /dev/null +++ b/tests/unit/test_drafter_sample_protocol.py @@ -0,0 +1,126 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from __future__ import annotations + +from dataclasses import replace + +import pytest +import torch + +from verl_speco.trainer.feature_store import DraftFeatureSample +from verl_speco.transport.drafter_sample_protocol import ( + ExpectedFeatureConfig, + SampleMetadata, + decode_sample, + encode_sample, + make_eos_record, + make_ready_tag, + make_sample_key, +) + + +def _metadata() -> SampleMetadata: + return SampleMetadata( + schema_version=1, + run_id="run-a", + sample_id="train-000017", + sequence_no=17, + algorithm="DSPARK", + target_model_id="/models/Qwen3-8B", + target_model_revision="rev-a", + tokenizer_fingerprint="sha256:tokenizer", + target_layer_ids=[2, 8, 14, -1], + hidden_states_layout="dflash_aux_plus_last", + hidden_dtype="bfloat16", + hidden_shape=[4, 16], + feature_length=4, + full_sequence_length=10, + feature_start=6, + feature_end=10, + use_logits=False, + ) + + +def _sample() -> DraftFeatureSample: + return DraftFeatureSample( + algorithm="DSPARK", + input_ids=torch.tensor([10, 11, 12, 13]), + loss_mask=torch.tensor([0.0, 1.0, 1.0, 1.0]), + position_ids=torch.tensor([6, 7, 8, 9]), + hidden_states=torch.arange(64, dtype=torch.bfloat16).reshape(4, 16), + metadata={"ignored_at_wire_boundary": True}, + ) + + +def _expected() -> ExpectedFeatureConfig: + meta = _metadata() + return ExpectedFeatureConfig( + run_id=meta.run_id, + target_model_id=meta.target_model_id, + target_model_revision=meta.target_model_revision, + tokenizer_fingerprint=meta.tokenizer_fingerprint, + target_layer_ids=meta.target_layer_ids, + hidden_states_layout=meta.hidden_states_layout, + hidden_dtype=meta.hidden_dtype, + ) + + +def test_sample_round_trip() -> None: + meta = _metadata() + key = make_sample_key(meta) + fields = encode_sample(_sample(), meta) + restored = decode_sample(key, make_ready_tag(meta), fields, _expected()) + + assert key == "drafter:v1:run-a:000000000017:train-000017" + assert tuple(restored.hidden_states.shape) == (4, 16) + assert restored.hidden_states.dtype == torch.bfloat16 + assert restored.input_ids.tolist() == [10, 11, 12, 13] + assert restored.metadata["sequence_no"] == 17 + assert restored.metadata["hidden_states_layout"] == "dflash_aux_plus_last" + assert fields["metadata_json"].dtype == torch.uint8 + + +def test_decode_rejects_identity_mismatch() -> None: + meta = _metadata() + fields = encode_sample(_sample(), meta) + bad_tag = {**make_ready_tag(meta), "sample_id": "wrong"} + with pytest.raises(ValueError, match="tag mismatch for sample_id"): + decode_sample(make_sample_key(meta), bad_tag, fields, _expected()) + + +def test_decode_rejects_consumer_contract_mismatch() -> None: + meta = _metadata() + fields = encode_sample(_sample(), meta) + expected = replace(_expected(), target_model_revision="different") + with pytest.raises(ValueError, match="target_model_revision"): + decode_sample(make_sample_key(meta), make_ready_tag(meta), fields, expected) + + +def test_encode_rejects_shape_mismatch() -> None: + meta = replace(_metadata(), hidden_shape=[4, 32]) + with pytest.raises(ValueError, match="hidden_states shape mismatch"): + encode_sample(_sample(), meta) + + +def test_eos_record_is_control_only() -> None: + key, fields, tag = make_eos_record("run-a", 18) + assert key == "control:v1:run-a:eos" + assert fields["marker"].tolist() == [1] + assert tag == { + "record_type": "control", + "status": "eos", + "schema_version": 1, + "run_id": "run-a", + "total_samples": 18, + } diff --git a/tests/unit/test_transferqueue_bridge.py b/tests/unit/test_transferqueue_bridge.py new file mode 100644 index 00000000..66049d68 --- /dev/null +++ b/tests/unit/test_transferqueue_bridge.py @@ -0,0 +1,183 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from __future__ import annotations + +import sys +from types import SimpleNamespace + +import pytest +import torch + +from verl_speco.integration import transferqueue_bridge as bridge + + +class _FakeClient: + def __init__(self) -> None: + self.closed = False + + def close(self) -> None: + self.closed = True + + +class _FakeTQ: + def __init__(self) -> None: + self.init_calls = [] + self.close_calls = 0 + self.clear_calls = [] + self.records: dict[str, tuple[dict, dict]] = {} + self.client = _FakeClient() + + def init(self, config=None): + self.init_calls.append(config) + + def kv_put(self, *, key, partition_id, fields, tag): + self.records[key] = (dict(fields), dict(tag)) + + def kv_list(self, *, partition_id): + return {key: tag for key, (_, tag) in self.records.items()} + + def kv_batch_get(self, *, keys, partition_id): + return {key: self.records[key][0] for key in keys} + + def kv_clear(self, *, keys, partition_id): + self.clear_calls.append((list(keys), partition_id)) + for key in keys: + self.records.pop(key, None) + + def get_client(self): + return self.client + + def close(self): + self.close_calls += 1 + + +class _FakeRay: + def __init__(self) -> None: + self.initialized = False + self.init_calls = [] + self.shutdown_calls = 0 + + def is_initialized(self): + return self.initialized + + def init(self, **kwargs): + self.initialized = True + self.init_calls.append(kwargs) + + def shutdown(self): + self.initialized = False + self.shutdown_calls += 1 + + +@pytest.fixture +def fake_runtime(monkeypatch): + fake_tq = _FakeTQ() + fake_ray = _FakeRay() + monkeypatch.setattr(bridge, "tq", fake_tq) + monkeypatch.setattr(bridge, "_TQ_IMPORTABLE", True) + monkeypatch.setitem(sys.modules, "ray", fake_ray) + monkeypatch.setattr( + bridge, + "_state", + { + "enabled": False, + "configured": False, + "initialized": False, + "config": None, + "owner": False, + "ray_initialized_here": False, + "ray_address": None, + "ray_namespace": None, + }, + ) + return fake_tq, fake_ray + + +def _config(): + return { + "enable": True, + "package_version": "0.1.7", + "partition_id": "speco_drafter_features", + "run_id": "run-a", + "schema_version": 1, + "ray": {"address": "ray-head:6379", "namespace": "speco-drafter"}, + "controller": {"polling_mode": True}, + "backend": { + "storage_backend": "SimpleStorage", + "SimpleStorage": {"total_storage_size": 16, "num_data_storage_units": 1}, + }, + } + + +def test_owner_connects_ray_and_receives_only_native_tq_config(fake_runtime) -> None: + fake_tq, fake_ray = fake_runtime + config = _config() + assert bridge.configure_transfer_queue(config) + bridge.connect_ray_cluster("ray-head:6379", "speco-drafter") + bridge.start_transfer_queue_owner(config) + + assert fake_ray.init_calls == [ + {"address": "ray-head:6379", "namespace": "speco-drafter"} + ] + native = bridge._to_plain_dict(fake_tq.init_calls[0]) + assert set(native) == {"controller", "backend"} + assert bridge._state["owner"] is True + + bridge.close_transfer_queue_owner() + assert fake_tq.close_calls == 1 + assert fake_ray.shutdown_calls == 1 + + +def test_client_put_list_get_many_clear_and_local_close(fake_runtime) -> None: + fake_tq, fake_ray = fake_runtime + assert bridge.configure_transfer_queue(_config()) + bridge.connect_ray_cluster("ray-head:6379", "speco-drafter") + bridge.connect_transfer_queue_client() + assert fake_tq.init_calls == [None] + + bridge.put_sample( + "k0", + {"hidden_states": torch.ones(2, 4), "ignored": None}, + tag={"status": "ready"}, + ) + bridge.put_sample( + "k1", + {"hidden_states": torch.zeros(3, 4)}, + tag={"status": "ready"}, + ) + + assert bridge.list_samples() == { + "k0": {"status": "ready"}, + "k1": {"status": "ready"}, + } + records = bridge.get_samples(["k1", "k0"]) + assert [key for key, _ in records] == ["k1", "k0"] + assert records[0][1]["hidden_states"].shape == (3, 4) + + bridge.clear_samples(["k0", "k1"]) + assert bridge.list_samples() == {} + bridge.close_transfer_queue_client() + assert fake_tq.client.closed is True + assert fake_tq.close_calls == 0 + assert fake_ray.shutdown_calls == 1 + + +def test_client_close_cannot_be_used_by_owner(fake_runtime) -> None: + _, _ = fake_runtime + config = _config() + bridge.configure_transfer_queue(config) + bridge.connect_ray_cluster("ray-head:6379", "speco-drafter") + bridge.start_transfer_queue_owner(config) + with pytest.raises(RuntimeError, match="owner"): + bridge.close_transfer_queue_client() diff --git a/verl_speco/config/speco_base.yaml b/verl_speco/config/speco_base.yaml index 664fa3be..13152ca9 100644 --- a/verl_speco/config/speco_base.yaml +++ b/verl_speco/config/speco_base.yaml @@ -194,8 +194,32 @@ actor_rollout_ref: # Ray path. Requires `pip install TransferQueue==0.1.7`. transfer_queue: enable: false + package_version: "0.1.7" + # TQ 0.1.7 discovers its named TransferQueueController through Ray. + # Standalone owner, Producer and every torchrun rank must use the same + # address and namespace. PR #48 Ray actors already have this context. + ray: + address: null + namespace: speco-drafter + partition_id: speco_drafter_features + run_id: null + schema_version: 1 + connect_timeout_seconds: 120 + poll_interval_seconds: 0.5 + drop_last: true + controller: + polling_mode: true backend: storage_backend: SimpleStorage SimpleStorage: total_storage_size: 100000 num_data_storage_units: 8 + MooncakeStore: + auto_init: false + metadata_server: localhost:50050 + master_server_address: localhost:50051 + local_hostname: "" + protocol: tcp + global_segment_size: 4294967296 + local_buffer_size: 1073741824 + device_name: "" diff --git a/verl_speco/integration/transferqueue_bridge.py b/verl_speco/integration/transferqueue_bridge.py index 7062b7f4..79407aa5 100644 --- a/verl_speco/integration/transferqueue_bridge.py +++ b/verl_speco/integration/transferqueue_bridge.py @@ -34,6 +34,7 @@ import logging import os import threading +from collections.abc import Mapping, Sequence from typing import Any, Optional import torch @@ -83,6 +84,9 @@ def _raise(*args: Any, **kwargs: Any) -> Any: "initialized": False, # tq.init() has run in this process "config": None, # the transfer_queue sub-config (plain dict) "owner": False, # this process created the task-level TQ system + "ray_initialized_here": False, + "ray_address": None, + "ray_namespace": None, } @@ -118,7 +122,13 @@ def _extract_tq_config(training_cfg: Any) -> Optional[dict]: elif isinstance(training_cfg, dict): transfer_queue_cfg = training_cfg.get("transfer_queue", None) if transfer_queue_cfg is None: - return None + plain = _to_plain_dict(training_cfg) + if isinstance(plain, dict) and any( + key in plain for key in ("enable", "backend", "controller", "ray") + ): + transfer_queue_cfg = plain + else: + return None return _to_plain_dict(transfer_queue_cfg) @@ -143,6 +153,71 @@ def is_transfer_queue_enabled() -> bool: return bool(_state["enabled"]) and _TQ_IMPORTABLE +def connect_ray_cluster( + ray_address: str | None, + namespace: str | None = None, +) -> None: + """Connect this ordinary process to the Ray cluster hosting TQ. + + PR #48 workers are already Ray actors and therefore do not call this + function. Standalone owner, Producer and torchrun ranks must call it + before ``tq.init`` so TQ 0.1.7 can discover its named Controller actor. + """ + + try: + import ray + except ImportError as exc: # pragma: no cover - depends on optional package + raise RuntimeError( + "Ray is required by TransferQueue 0.1.7. Install TransferQueue==0.1.7 " + "and connect all standalone processes to the same Ray cluster." + ) from exc + + if ray.is_initialized(): + return + address = str(ray_address or "auto").strip() or "auto" + kwargs: dict[str, Any] = {"address": address} + if namespace: + kwargs["namespace"] = str(namespace) + ray.init(**kwargs) + with _state_lock: + _state["ray_initialized_here"] = True + _state["ray_address"] = address + _state["ray_namespace"] = namespace + + +def start_transfer_queue_owner(tq_config: Any) -> None: + """Create the task-level named TQ Controller in the current Ray cluster.""" + + plain = _extract_tq_config(tq_config) + if plain is None: + raise ValueError("TransferQueue owner configuration is missing") + if not bool(plain.get("enable", True)): + raise ValueError("TransferQueue owner requires transfer_queue.enable=true") + if not _TQ_IMPORTABLE: + raise RuntimeError("TransferQueue==0.1.7 is required to start the TQ owner") + with _state_lock: + if _state["initialized"]: + raise RuntimeError("TransferQueue is already initialized in this process") + tq.init(_as_tq_config(_native_tq_config(plain))) + with _state_lock: + _state["config"] = plain + _state["enabled"] = True + _state["configured"] = True + _state["initialized"] = True + _state["owner"] = True + logger.info("[SpeCo TQ] standalone owner started (partition=%s)", _partition_id()) + + +def connect_transfer_queue_client() -> None: + """Attach this process to the named TQ Controller on its Ray cluster.""" + + if not _TQ_IMPORTABLE: + raise RuntimeError("TransferQueue==0.1.7 is required to connect a TQ client") + if not bool(_state["enabled"]): + raise RuntimeError("configure_transfer_queue() must enable TQ before client connect") + _ensure_initialized() + + def init_transfer_queue(config: Any) -> bool: """Cluster-wide TQ bootstrap, called once from the SpecoTaskRunner. @@ -155,7 +230,7 @@ def init_transfer_queue(config: Any) -> bool: tq_cfg = _extract_tq_config(_drafter_training_cfg(config)) if tq_cfg is None or not bool(tq_cfg.get("enable")) or not _TQ_IMPORTABLE: return False - tq.init(_to_plain_dict(tq_cfg)) + tq.init(_as_tq_config(_native_tq_config(tq_cfg))) with _state_lock: _state["config"] = _to_plain_dict(tq_cfg) _state["enabled"] = True @@ -189,6 +264,40 @@ def _ensure_initialized() -> None: _state["initialized"] = True +def _native_tq_config(tq_cfg: Mapping[str, Any]) -> dict[str, Any]: + """Remove SPECO-only connection/protocol fields before ``tq.init``.""" + + project_keys = { + "enable", + "package_version", + "ray", + "partition_id", + "run_id", + "schema_version", + "connect_timeout_seconds", + "poll_interval_seconds", + "drop_last", + } + return {key: value for key, value in tq_cfg.items() if key not in project_keys} + + +def _as_tq_config(value: Mapping[str, Any]) -> Any: + """TQ 0.1.7 annotates its config as DictConfig; keep tests dependency-light.""" + + try: + from omegaconf import OmegaConf + + return OmegaConf.create(dict(value)) + except ImportError: # pragma: no cover - project normally depends on OmegaConf + return dict(value) + + +def _partition_id() -> str: + config = _state.get("config") or {} + value = config.get("partition_id") if isinstance(config, Mapping) else None + return str(value or _SPECO_TQ_PARTITION) + + # --------------------------------------------------------------------------- # Key / put / get / close # --------------------------------------------------------------------------- @@ -228,7 +337,7 @@ def put_sample( # incorrect. (Exact kwarg names verified against TQ 0.1.7 on first run.) tq.kv_put( key=key, - partition_id=_SPECO_TQ_PARTITION, + partition_id=_partition_id(), fields=payload, tag=tag or {}, ) @@ -247,13 +356,111 @@ def get_sample(key: str) -> dict: _ensure_initialized() # TQ returns the stored sample (TensorDict-like). Return shape is version # dependent, so handle both a direct value and a {key: value} mapping. - result = tq.kv_batch_get(keys=[key], partition_id=_SPECO_TQ_PARTITION) + result = tq.kv_batch_get(keys=[key], partition_id=_partition_id()) value = _extract_value(result, key) if value is None: return {} return _tensordict_to_dict(value) +def list_samples() -> dict[str, dict[str, Any]]: + """Return key -> tag for the configured partition without fetching fields.""" + + if not is_transfer_queue_enabled(): + raise RuntimeError("list_samples called while TransferQueue is not enabled.") + _ensure_initialized() + result = tq.kv_list(partition_id=_partition_id()) + if result is None: + return {} + if not isinstance(result, Mapping): + raise TypeError(f"tq.kv_list returned unsupported type {type(result)!r}") + # 0.1.7 returns key -> tag when partition_id is supplied. Accept the + # partition -> (key -> tag) wrapper as well to keep the bridge version-safe. + nested = result.get(_partition_id()) + if isinstance(nested, Mapping) and all(isinstance(v, Mapping) for v in nested.values()): + result = nested + records: dict[str, dict[str, Any]] = {} + for key, tag in result.items(): + if tag is None: + records[str(key)] = {} + elif isinstance(tag, Mapping): + records[str(key)] = dict(tag) + else: + raise TypeError(f"TQ tag for key {key!r} must be a mapping, got {type(tag)!r}") + return records + + +def get_samples(keys: Sequence[str]) -> list[tuple[str, dict[str, Any]]]: + """Batch-fetch records and return one plain field dict per input key.""" + + normalized_keys = [str(key) for key in keys] + if not normalized_keys: + return [] + if len(set(normalized_keys)) != len(normalized_keys): + raise ValueError("get_samples keys must be unique") + if not is_transfer_queue_enabled(): + raise RuntimeError("get_samples called while TransferQueue is not enabled.") + _ensure_initialized() + result = tq.kv_batch_get(keys=normalized_keys, partition_id=_partition_id()) + values = _split_batch_result(result, normalized_keys) + return [ + (key, _tensordict_to_dict(value)) + for key, value in zip(normalized_keys, values, strict=True) + ] + + +def clear_samples(keys: Sequence[str]) -> None: + """Delete consumed records from the configured partition.""" + + normalized_keys = [str(key) for key in keys] + if not normalized_keys: + return + if not is_transfer_queue_enabled(): + raise RuntimeError("clear_samples called while TransferQueue is not enabled.") + _ensure_initialized() + tq.kv_clear(keys=normalized_keys, partition_id=_partition_id()) + + +def close_transfer_queue_client() -> None: + """Close only this process's TQ client; never kill the shared Controller.""" + + with _state_lock: + if _state["owner"]: + raise RuntimeError("TQ owner must use close_transfer_queue_owner()") + initialized = bool(_state["initialized"]) + _state["initialized"] = False + if initialized and _TQ_IMPORTABLE: + try: + client = tq.get_client() + if client is not None: + client.close() + except Exception: # noqa: BLE001 + logger.debug("[SpeCo TQ] local client close raised", exc_info=True) + _shutdown_local_ray_connection() + + +def close_transfer_queue_owner() -> None: + """Close global TQ resources. Only the process that started them may call.""" + + close_transfer_queue() + _shutdown_local_ray_connection() + + +def _shutdown_local_ray_connection() -> None: + with _state_lock: + initialized_here = bool(_state["ray_initialized_here"]) + _state["ray_initialized_here"] = False + if not initialized_here: + return + try: + import ray + + if ray.is_initialized(): + ray.shutdown() + except ImportError: # pragma: no cover + return + + def close_transfer_queue() -> None: """Close task-level TQ resources if this process initialized them. @@ -277,6 +484,42 @@ def close_transfer_queue() -> None: logger.debug("[SpeCo TQ] tq.close() raised; ignoring shutdown error") +def _split_batch_result(result: Any, keys: Sequence[str]) -> list[Any]: + if result is None: + raise KeyError(f"TQ returned no payload for keys={list(keys)!r}") + # TensorDict exposes a batch_size and indexes rows with ``result[index]``. + # Check this before the generic Mapping branch because some TensorDict + # versions also satisfy mapping-like protocols for their field columns. + if getattr(result, "batch_size", None) is not None: + return _index_batch_rows(result, len(keys)) + if isinstance(result, Mapping): + if all(key in result for key in keys): + return [result[key] for key in keys] + if isinstance(result, (list, tuple)): + if len(result) != len(keys): + raise RuntimeError( + f"TQ returned {len(result)} rows for {len(keys)} requested keys" + ) + return list(result) + if len(keys) == 1: + return [result] + return _index_batch_rows(result, len(keys)) + + +def _index_batch_rows(result: Any, expected_rows: int) -> list[Any]: + try: + rows = [result[index] for index in range(expected_rows)] + except Exception as exc: # noqa: BLE001 + raise TypeError( + f"Unable to split TQ batch result of type {type(result)!r} into {expected_rows} rows" + ) from exc + if len(rows) != expected_rows: + raise RuntimeError( + f"TQ returned {len(rows)} rows for {expected_rows} requested keys" + ) + return rows + + def _extract_value(result: Any, key: str) -> Any: if result is None: return None @@ -296,10 +539,18 @@ def _tensordict_to_dict(value: Any) -> dict: __all__ = [ "KVBatchMeta", "configure_transfer_queue", + "connect_ray_cluster", + "connect_transfer_queue_client", + "start_transfer_queue_owner", "close_transfer_queue", + "close_transfer_queue_client", + "close_transfer_queue_owner", + "clear_samples", "init_transfer_queue", "is_transfer_queue_enabled", + "list_samples", "make_sample_key", "put_sample", "get_sample", + "get_samples", ] diff --git a/verl_speco/tq_owner.py b/verl_speco/tq_owner.py new file mode 100644 index 00000000..0d8d1400 --- /dev/null +++ b/verl_speco/tq_owner.py @@ -0,0 +1,115 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Long-lived TransferQueue owner for standalone Producer/Consumer jobs.""" + +from __future__ import annotations + +import logging +import signal +import threading +from typing import Any + +import hydra +import torch +from omegaconf import OmegaConf + +from verl_speco.integration.transferqueue_bridge import ( + close_transfer_queue_owner, + configure_transfer_queue, + connect_ray_cluster, + put_sample, + start_transfer_queue_owner, +) + + +logger = logging.getLogger(__name__) + + +def install_signal_handlers(stop_event: threading.Event) -> None: + def _request_stop(signum: int, _frame: Any) -> None: + logger.info("TQ owner received signal %s", signum) + stop_event.set() + + signal.signal(signal.SIGINT, _request_stop) + signal.signal(signal.SIGTERM, _request_stop) + + +def publish_owner_ready(run_id: str, schema_version: int) -> str: + if not run_id: + raise ValueError("transfer_queue.run_id must be set for standalone owner") + key = f"control:v{int(schema_version)}:{run_id}:owner-ready" + put_sample( + key, + {"marker": torch.tensor([1], dtype=torch.uint8)}, + tag={ + "record_type": "control", + "status": "owner_ready", + "schema_version": int(schema_version), + "run_id": run_id, + }, + ) + return key + + +def wait_until_stopped(stop_event: threading.Event) -> None: + stop_event.wait() + + +def run_owner(config: Any, *, stop_event: threading.Event | None = None) -> int: + training_cfg = config.actor_rollout_ref.rollout.drafter.training + tq_cfg = OmegaConf.to_container(training_cfg.transfer_queue, resolve=True) + if not isinstance(tq_cfg, dict): + raise TypeError("transfer_queue configuration must resolve to a mapping") + # Invoking the dedicated owner entrypoint is itself the request to enable + # TQ. Keep speco_base.yaml disabled by default for ordinary training jobs, + # and enable only this process's copied configuration. + tq_cfg["enable"] = True + if not configure_transfer_queue(tq_cfg): + raise RuntimeError( + "Standalone TQ owner requires TransferQueue==0.1.7" + ) + ray_cfg = tq_cfg.get("ray", {}) + ray_address = ray_cfg.get("address") + if not ray_address: + raise ValueError("transfer_queue.ray.address must point to a running Ray cluster") + namespace = ray_cfg.get("namespace") + event = stop_event or threading.Event() + if stop_event is None: + install_signal_handlers(event) + + started = False + try: + connect_ray_cluster(str(ray_address), str(namespace) if namespace else None) + start_transfer_queue_owner(tq_cfg) + started = True + ready_key = publish_owner_ready( + str(tq_cfg.get("run_id") or ""), + int(tq_cfg.get("schema_version", 1)), + ) + logger.info("TQ owner ready key=%s", ready_key) + wait_until_stopped(event) + return 0 + finally: + if started: + close_transfer_queue_owner() + + +@hydra.main(config_path="config", config_name="speco_base", version_base=None) +def main(config: Any) -> None: + logging.basicConfig(level=logging.INFO) + raise SystemExit(run_owner(config)) + + +if __name__ == "__main__": + main() diff --git a/verl_speco/transport/__init__.py b/verl_speco/transport/__init__.py new file mode 100644 index 00000000..4acd2a06 --- /dev/null +++ b/verl_speco/transport/__init__.py @@ -0,0 +1,38 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Transport protocols shared by standalone SPECO producers and consumers.""" + +from verl_speco.transport.drafter_sample_protocol import ( + DRAFTER_TQ_PARTITION, + PROTOCOL_SCHEMA_VERSION, + ExpectedFeatureConfig, + SampleMetadata, + decode_sample, + encode_sample, + make_eos_record, + make_ready_tag, + make_sample_key, +) + +__all__ = [ + "DRAFTER_TQ_PARTITION", + "PROTOCOL_SCHEMA_VERSION", + "ExpectedFeatureConfig", + "SampleMetadata", + "decode_sample", + "encode_sample", + "make_eos_record", + "make_ready_tag", + "make_sample_key", +] diff --git a/verl_speco/transport/drafter_sample_protocol.py b/verl_speco/transport/drafter_sample_protocol.py new file mode 100644 index 00000000..de54b00f --- /dev/null +++ b/verl_speco/transport/drafter_sample_protocol.py @@ -0,0 +1,393 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Wire protocol for standalone drafter samples stored in TransferQueue. + +One TQ key represents one training sample. Tensor payloads live in TQ fields; +small discovery attributes live in the TQ tag; richer metadata is JSON encoded +as a uint8 tensor so Producer and Consumer use one versioned contract. +""" + +from __future__ import annotations + +import json +from dataclasses import asdict, dataclass +from typing import Any, Mapping + +import torch + +from verl_speco.trainer.feature_store import DraftFeatureSample + + +PROTOCOL_SCHEMA_VERSION = 1 +DRAFTER_TQ_PARTITION = "speco_drafter_features" +_REQUIRED_FIELDS = ( + "input_ids", + "loss_mask", + "position_ids", + "hidden_states", + "metadata_json", +) +_OPTIONAL_TENSOR_FIELDS = ( + "last_hidden_states", + "target", + "target_logprobs", +) + + +@dataclass(frozen=True) +class SampleMetadata: + schema_version: int + run_id: str + sample_id: str + sequence_no: int + algorithm: str + target_model_id: str + target_model_revision: str + tokenizer_fingerprint: str + target_layer_ids: list[int] + hidden_states_layout: str + hidden_dtype: str + hidden_shape: list[int] + feature_length: int + full_sequence_length: int + feature_start: int + feature_end: int + use_logits: bool + + def validate(self) -> None: + if self.schema_version != PROTOCOL_SCHEMA_VERSION: + raise ValueError( + f"Unsupported drafter sample schema_version={self.schema_version}; " + f"expected {PROTOCOL_SCHEMA_VERSION}" + ) + if not self.run_id: + raise ValueError("SampleMetadata.run_id must not be empty") + if not self.sample_id: + raise ValueError("SampleMetadata.sample_id must not be empty") + if self.sequence_no < 0: + raise ValueError("SampleMetadata.sequence_no must be non-negative") + if self.algorithm.strip().upper() != "DSPARK": + raise ValueError( + f"Standalone TQ protocol currently requires algorithm=DSPARK, got {self.algorithm!r}" + ) + if len(self.hidden_shape) != 2: + raise ValueError( + f"SampleMetadata.hidden_shape must be [rows, hidden_dim], got {self.hidden_shape!r}" + ) + if self.feature_length <= 0: + raise ValueError("SampleMetadata.feature_length must be positive") + if self.hidden_shape[0] != self.feature_length: + raise ValueError( + "SampleMetadata hidden_shape/feature_length mismatch: " + f"{self.hidden_shape[0]} vs {self.feature_length}" + ) + if not (0 <= self.feature_start < self.feature_end <= self.full_sequence_length): + raise ValueError( + "SampleMetadata feature window must satisfy " + "0 <= feature_start < feature_end <= full_sequence_length" + ) + if self.feature_end - self.feature_start != self.feature_length: + raise ValueError( + "SampleMetadata feature window length does not match feature_length: " + f"{self.feature_end - self.feature_start} vs {self.feature_length}" + ) + + def to_dict(self) -> dict[str, Any]: + self.validate() + payload = asdict(self) + payload["algorithm"] = self.algorithm.strip().upper() + return payload + + @classmethod + def from_dict(cls, payload: Mapping[str, Any]) -> "SampleMetadata": + try: + meta = cls( + schema_version=int(payload["schema_version"]), + run_id=str(payload["run_id"]), + sample_id=str(payload["sample_id"]), + sequence_no=int(payload["sequence_no"]), + algorithm=str(payload["algorithm"]), + target_model_id=str(payload["target_model_id"]), + target_model_revision=str(payload["target_model_revision"]), + tokenizer_fingerprint=str(payload["tokenizer_fingerprint"]), + target_layer_ids=[int(v) for v in payload["target_layer_ids"]], + hidden_states_layout=str(payload["hidden_states_layout"]), + hidden_dtype=str(payload["hidden_dtype"]), + hidden_shape=[int(v) for v in payload["hidden_shape"]], + feature_length=int(payload["feature_length"]), + full_sequence_length=int(payload["full_sequence_length"]), + feature_start=int(payload["feature_start"]), + feature_end=int(payload["feature_end"]), + use_logits=bool(payload["use_logits"]), + ) + except KeyError as exc: + raise ValueError(f"metadata_json missing required field {exc.args[0]!r}") from exc + meta.validate() + return meta + + +@dataclass(frozen=True) +class ExpectedFeatureConfig: + """Consumer-side contract. ``None`` fields are intentionally unchecked.""" + + run_id: str + schema_version: int = PROTOCOL_SCHEMA_VERSION + algorithm: str = "DSPARK" + target_model_id: str | None = None + target_model_revision: str | None = None + tokenizer_fingerprint: str | None = None + target_layer_ids: list[int] | None = None + hidden_states_layout: str | None = None + hidden_dtype: str | None = None + + +def make_sample_key(meta: SampleMetadata) -> str: + meta.validate() + return ( + f"drafter:v{meta.schema_version}:{meta.run_id}:" + f"{meta.sequence_no:012d}:{meta.sample_id}" + ) + + +def make_ready_tag(meta: SampleMetadata) -> dict[str, Any]: + meta.validate() + return { + "record_type": "sample", + "status": "ready", + "schema_version": meta.schema_version, + "run_id": meta.run_id, + "sequence_no": meta.sequence_no, + "sample_id": meta.sample_id, + "algorithm": meta.algorithm.strip().upper(), + } + + +def encode_sample( + sample: DraftFeatureSample | Mapping[str, Any], meta: SampleMetadata +) -> dict[str, torch.Tensor]: + """Encode one normalized feature sample into TQ tensor fields.""" + + meta.validate() + normalized = ( + sample + if isinstance(sample, DraftFeatureSample) + else DraftFeatureSample.from_dict(dict(sample), strict=True) + ) + normalized.validate(strict=True) + if isinstance(normalized.hidden_states, (list, tuple)): + raise TypeError("TQ drafter protocol requires hidden_states to be one dense tensor") + hidden = _cpu_contiguous(normalized.hidden_states) + input_ids = _cpu_contiguous(normalized.input_ids, dtype=torch.int64).reshape(-1) + loss_mask = _cpu_contiguous(normalized.loss_mask, dtype=torch.float32).reshape(-1) + if normalized.position_ids is None: + position_ids = torch.arange(input_ids.numel(), dtype=torch.int64) + else: + position_ids = _cpu_contiguous(normalized.position_ids, dtype=torch.int64).reshape(-1) + + _validate_primary_tensors(input_ids, loss_mask, position_ids, hidden, meta) + metadata_json = _json_to_tensor(meta.to_dict()) + fields: dict[str, torch.Tensor] = { + "input_ids": input_ids, + "loss_mask": loss_mask, + "position_ids": position_ids, + "hidden_states": hidden, + "metadata_json": metadata_json, + } + for field_name in _OPTIONAL_TENSOR_FIELDS: + value = getattr(normalized, field_name) + if value is not None: + fields[field_name] = _cpu_contiguous(value) + return fields + + +def decode_sample( + key: str, + tag: Mapping[str, Any], + fields: Mapping[str, Any], + expected_config: ExpectedFeatureConfig | Mapping[str, Any], +) -> DraftFeatureSample: + """Validate a TQ record and restore the existing training sample type.""" + + expected = ( + expected_config + if isinstance(expected_config, ExpectedFeatureConfig) + else ExpectedFeatureConfig(**dict(expected_config)) + ) + missing = [name for name in _REQUIRED_FIELDS if name not in fields] + if missing: + raise ValueError(f"TQ sample {key!r} missing required fields: {missing}") + metadata = SampleMetadata.from_dict(_tensor_to_json(fields["metadata_json"])) + expected_key = make_sample_key(metadata) + if key != expected_key: + raise ValueError(f"TQ sample key mismatch: got {key!r}, expected {expected_key!r}") + _validate_tag(tag, metadata) + _validate_expected(metadata, expected) + + input_ids = _require_tensor(fields, "input_ids").detach().cpu().to(torch.int64).reshape(-1) + loss_mask = _require_tensor(fields, "loss_mask").detach().cpu().to(torch.float32).reshape(-1) + position_ids = _require_tensor(fields, "position_ids").detach().cpu().to(torch.int64).reshape(-1) + hidden = _require_tensor(fields, "hidden_states").detach().cpu().contiguous() + _validate_primary_tensors(input_ids, loss_mask, position_ids, hidden, metadata) + + payload: dict[str, Any] = { + "schema_version": metadata.schema_version, + "algorithm": metadata.algorithm, + "input_ids": input_ids, + "loss_mask": loss_mask, + "position_ids": position_ids, + "hidden_states": hidden, + "metadata": metadata.to_dict(), + } + for field_name in _OPTIONAL_TENSOR_FIELDS: + if field_name in fields and fields[field_name] is not None: + payload[field_name] = _require_tensor(fields, field_name).detach().cpu().contiguous() + return DraftFeatureSample.from_dict(payload, strict=True) + + +def make_eos_record( + run_id: str, total_samples: int +) -> tuple[str, dict[str, torch.Tensor], dict[str, Any]]: + if not run_id: + raise ValueError("run_id must not be empty") + if total_samples < 0: + raise ValueError("total_samples must be non-negative") + key = f"control:v{PROTOCOL_SCHEMA_VERSION}:{run_id}:eos" + fields = {"marker": torch.tensor([1], dtype=torch.uint8)} + tag = { + "record_type": "control", + "status": "eos", + "schema_version": PROTOCOL_SCHEMA_VERSION, + "run_id": run_id, + "total_samples": int(total_samples), + } + return key, fields, tag + + +def _validate_tag(tag: Mapping[str, Any], meta: SampleMetadata) -> None: + expected = make_ready_tag(meta) + for name, expected_value in expected.items(): + if tag.get(name) != expected_value: + raise ValueError( + f"TQ sample tag mismatch for {name}: got {tag.get(name)!r}, " + f"expected {expected_value!r}" + ) + + +def _validate_expected(meta: SampleMetadata, expected: ExpectedFeatureConfig) -> None: + checks = { + "run_id": expected.run_id, + "schema_version": expected.schema_version, + "algorithm": expected.algorithm.strip().upper(), + "target_model_id": expected.target_model_id, + "target_model_revision": expected.target_model_revision, + "tokenizer_fingerprint": expected.tokenizer_fingerprint, + "target_layer_ids": expected.target_layer_ids, + "hidden_states_layout": expected.hidden_states_layout, + "hidden_dtype": expected.hidden_dtype, + } + for name, expected_value in checks.items(): + if expected_value is None: + continue + actual = getattr(meta, name) + if name == "algorithm": + actual = str(actual).strip().upper() + if actual != expected_value: + raise ValueError( + f"TQ sample metadata mismatch for {name}: got {actual!r}, " + f"expected {expected_value!r}" + ) + + +def _validate_primary_tensors( + input_ids: torch.Tensor, + loss_mask: torch.Tensor, + position_ids: torch.Tensor, + hidden: torch.Tensor, + meta: SampleMetadata, +) -> None: + if hidden.dim() == 3 and hidden.size(0) == 1: + hidden = hidden.squeeze(0) + if hidden.dim() != 2: + raise ValueError(f"hidden_states must have shape [L,D], got {tuple(hidden.shape)}") + lengths = { + "input_ids": int(input_ids.numel()), + "loss_mask": int(loss_mask.numel()), + "position_ids": int(position_ids.numel()), + "hidden_states": int(hidden.size(0)), + } + if any(value != meta.feature_length for value in lengths.values()): + raise ValueError( + f"TQ sample tensor lengths must equal feature_length={meta.feature_length}: {lengths}" + ) + if list(hidden.shape) != meta.hidden_shape: + raise ValueError( + f"hidden_states shape mismatch: got {list(hidden.shape)}, expected {meta.hidden_shape}" + ) + actual_dtype = _dtype_name(hidden.dtype) + if actual_dtype != meta.hidden_dtype: + raise ValueError( + f"hidden_states dtype mismatch: got {actual_dtype!r}, expected {meta.hidden_dtype!r}" + ) + + +def _json_to_tensor(payload: Mapping[str, Any]) -> torch.Tensor: + raw = json.dumps(dict(payload), sort_keys=True, separators=(",", ":")).encode("utf-8") + return torch.tensor(list(raw), dtype=torch.uint8) + + +def _tensor_to_json(value: Any) -> dict[str, Any]: + tensor = value + if not torch.is_tensor(tensor): + raise TypeError("metadata_json must be a torch.Tensor") + tensor = tensor.detach().cpu().to(torch.uint8).reshape(-1) + try: + decoded = json.loads(bytes(tensor.tolist()).decode("utf-8")) + except (UnicodeDecodeError, json.JSONDecodeError) as exc: + raise ValueError("metadata_json is not valid UTF-8 JSON") from exc + if not isinstance(decoded, dict): + raise ValueError("metadata_json must decode to a JSON object") + return decoded + + +def _require_tensor(fields: Mapping[str, Any], name: str) -> torch.Tensor: + value = fields.get(name) + if not torch.is_tensor(value): + raise TypeError(f"TQ field {name!r} must be a torch.Tensor") + return value + + +def _cpu_contiguous(value: torch.Tensor, *, dtype: torch.dtype | None = None) -> torch.Tensor: + if not torch.is_tensor(value): + raise TypeError(f"Expected torch.Tensor, got {type(value)!r}") + result = value.detach().cpu() + if dtype is not None: + result = result.to(dtype) + return result.contiguous() + + +def _dtype_name(dtype: torch.dtype) -> str: + return str(dtype).removeprefix("torch.") + + +__all__ = [ + "DRAFTER_TQ_PARTITION", + "PROTOCOL_SCHEMA_VERSION", + "ExpectedFeatureConfig", + "SampleMetadata", + "decode_sample", + "encode_sample", + "make_eos_record", + "make_ready_tag", + "make_sample_key", +] From c61fa393a7fca4ea1c50906d9580cea9f13dc7e9 Mon Sep 17 00:00:00 2001 From: xyr <1615681145@qq.com> Date: Wed, 19 Aug 2026 15:26:50 +0800 Subject: [PATCH 32/50] Remove internal TQ design documents Keep implementation planning and explanatory notes out of the tracked project tree. Assisted-by: OpenAI Codex --- ...sync_vllm_mooncake_dspark_training_plan.md | 2896 ----------------- ...standalone_tq_foundation_implementation.md | 1037 ------ ...standalone_vllm_tq_dspark_training_plan.md | 1131 ------- 3 files changed, 5064 deletions(-) delete mode 100644 docs/async_vllm_mooncake_dspark_training_plan.md delete mode 100644 docs/standalone_tq_foundation_implementation.md delete mode 100644 docs/standalone_vllm_tq_dspark_training_plan.md diff --git a/docs/async_vllm_mooncake_dspark_training_plan.md b/docs/async_vllm_mooncake_dspark_training_plan.md deleted file mode 100644 index 6cda0e40..00000000 --- a/docs/async_vllm_mooncake_dspark_training_plan.md +++ /dev/null @@ -1,2896 +0,0 @@ -# verl-SpeCo 异步 vLLM → TransferQueue(Mooncake 后端)→ DSpark 流式训练方案 - -## 1. 文档范围 - -本文只讨论当前 `verl-SpeCo-ls` 项目的独立草稿模型训练入口: - -```text -examples/run_qwen3-8b_drafter_separate_training.sh - → python -m verl_speco.draft_train_launcher - → torch.distributed.run - → python -m verl_speco.draft_train - → run_standalone_draft_training() -``` - -目标是训练 `mode=train/offline` 的草稿模型,不是 RL 与草稿模型一起训练,也不是先生成全量 hidden states 再长期保存到磁盘。 - -输入文件已经包含提前生成好的 response。新流水线需要: - -1. Producer 读取 prompt 和预生成 response,构造完整 token 序列; -2. Producer 并行请求一个或多个 vLLM endpoint 做 prefill; -3. hidden states 写入 TransferQueue(简称 TQ),TQ 的数据后端使用 Mooncake; -4. DSpark 的各训练 rank 从 TQ/Mooncake 并行读取自己的 local batch; -5. 所有 rank 完成同一个 optimizer step 后,清理这一批 TQ 数据; -6. 不使用自研 Stream Coordinator;第一版复用 PR #48 已验证的 TQ KV 传输模式,由 TQ 保存 payload 和 tags,rank 0 通过 `kv_list` 发现 ready key 并组成 global batch。 - -本文不照搬 AngelSpec 的进程组织。实现依据是旁边只读参考仓库 `verl-SpeCo` 的 PR #48 分支,尤其是: - -```text -verl_speco/integration/transferqueue_bridge.py -verl_speco/integration/sglang_runtime.py -verl_speco/integration/oldlogprob_runtime.py -verl_speco/workers/speco_worker.py -verl_speco/integration/task_runner.py -``` - -参考仓库只用于理解和复用设计,不修改其代码。真正实现仍放在本项目 `verl-SpeCo-ls`,服务于 standalone drafter training。 - -建议按下面顺序阅读: - -1. 第一部分先介绍 verl/SpeCo 中的进程、`TokenOutput`、`DataProto`、WorkerGroup、replica 和 owner rank; -2. 再完整解释 PR #48 的 SGLang TQ 路径; -3. 再解释 PR #48 的 old-logprob TQ 路径; -4. 再解释 hidden 恢复后如何进入 online buffer 和 optimizer step; -5. 第二部分才讨论如何把相同 TQ 传输能力改造成独立 Producer/DSpark Consumer。 - -## 第一部分:PR #48 原始 TQ 流程 - -这一部分只解释只读参考仓库 `../verl-SpeCo` 的 PR #48。这里的 Producer、Consumer、Ray driver、数据结构和生命周期全部是 PR #48 当前代码的行为,不是本文为 standalone 设计的行为。 - -### 1A. 阅读 PR #48 前必须知道的项目对象 - -#### SGLang server - -SGLang server 是 rollout 推理进程。它接收 prompt,执行 target model 推理并生成 response。SpeCo 的 SGLang patch 还会在推理过程中收集指定层 hidden states。 - -它不是 drafter trainer,也不执行 optimizer step。在 PR #48 中它是 hidden-state Producer。 - -#### TokenOutput - -`TokenOutput` 是一次 rollout request 的返回对象。核心字段可以概括为: - -```python -TokenOutput( - token_ids=list[int], - log_probs=..., - routed_experts=..., - extra_fields={ - "global_steps": int, - "drafter_sample": dict | None, - }, -) -``` - -`token_ids` 是生成结果;`extra_fields` 是 SpeCo 添加的旁路字段。`drafter_sample` 不影响正常 response 返回,它用于把草稿训练所需信息从 rollout 侧带回训练控制层。 - -#### DataProto 和 non_tensor_batch - -verl 将一批 rollout 结果整理成 `DataProto`。它通常分为: - -```python -DataProto( - batch=TensorDict(...), - non_tensor_batch={...}, - meta_info={...}, -) -``` - -- `batch`:规则的批量 tensor,例如 prompts、responses、attention mask; -- `non_tensor_batch`:不能直接组成规则 dense batch 的 Python 对象或 object array; -- `meta_info`:批次级配置和指标。 - -每个 request 的 `TokenOutput.extra_fields["drafter_sample"]` 最终会被 rollout/agent-loop 聚合到: - -```python -gen_batch_output.non_tensor_batch["drafter_sample"] -``` - -所以 SGLang 侧写的是单 request `TokenOutput`,driver 侧拿到的是批量 `DataProto`。 - -#### RayPPOTrainer driver - -`SpecoRayPPOTrainer` 所在进程是控制进程,简称 driver。它负责按训练 step 依次调用 rollout、old-logprob、actor update,以及向各 WorkerGroup 发 RPC。 - -driver 不执行 drafter 模型 forward。PR #48 之前,它会接触包含 hidden tensor 的 `drafter_sample`;PR #48 之后,它只处理中小 tensor、标量 metadata 和 TQ key。 - -#### WorkerGroup - -WorkerGroup 是 verl 对一组 Ray worker 的调用封装。代码: - -```python -self.drafter_wg.collect_rollout_features(buckets) -``` - -不是普通本地函数调用,而是根据注册的 dispatch rule,将参数分发到多张卡上的 `SpecoWorker.collect_rollout_features()`。 - -#### Rollout replica - -rollout replica 是一组共同承载一个 rollout 模型副本的进程/GPU。一个任务可能存在多个 data-parallel rollout replica。`replica_rank` 标识当前 sample 是哪个 rollout replica 生成的: - -```text -replica_rank = 0, 1, 2, ... -``` - -#### Drafter training replica、DP rank 和 SP rank - -drafter 训练也可能按 data parallel 和 sequence parallel 组织: - -```text -drafter replica / DP rank 0 - ├─ SP rank 0 - └─ SP rank 1 - -drafter replica / DP rank 1 - ├─ SP rank 0 - └─ SP rank 1 -``` - -同一个 drafter DP replica 内的 SP ranks 共同执行一个模型副本的训练。`replica_rank` 用来把 rollout replica 产生的数据路由到对应 drafter DP replica。 - -#### Owner rank - -`collect_rollout_features` 注册了: - -```python -@register( - dispatch_mode=make_nd_compute_dispatch_fn( - mesh_name="drafter_owner_route" - ) -) -``` - -每个 drafter DP replica 的 `SP rank 0` 被标记为 collect leader/owner。driver 传入的是按 replica 分好的 bucket,dispatch 层负责把 bucket 发到对应训练组。一个 replica 内可能有多个 rank 需要共同训练,但只有指定 leader 负责汇总 RPC 返回。 - -这里的 owner route 是 PR #48 为什么不能简单让任意 worker 随机拿 sample 的原因:sample 必须进入与当前 drafter device mesh 一致的训练组。 - -## 2. PR #48 改造前的 online 特征流程 - -PR #48 改造的是 SpeCo 的 online drafter feature transport。hidden states 有两条主要来源。 - -### 2.1 SGLang rollout hidden 路径 - -改造前: - -```text -SGLang server - → 生成 drafter_sample,其中直接包含 hidden_states CPU tensor - → TokenOutput.extra_fields - → RayPPOTrainer driver 收集 drafter_sample - → driver 按 drafter replica/owner 分桶 - → Ray dispatch / object store - → SpecoWorker.collect_rollout_features(samples) - → _store_rollout_sample() - → online drafter buffer/train -``` - -此时 `drafter_sample` 类似: - -```python -{ - "input_ids": int64[1, total_len], - "prompts": int64[1, prompt_len], - "responses": int64[1, response_len], - "hidden_states": bf16[1, hidden_rows, hidden_dim], - "hidden_positions": int64[1, hidden_rows], - "target_logprobs": tensor | None, - "global_step": 42, - "replica_rank": 1, -} -``` - -问题是整个字典经过 driver,而 `hidden_states` 是其中最大的字段。driver 的 host memory 和 Ray object store 都会承载这些 tensor。 - -### 2.2 old-logprob hook hidden 路径 - -另一条路径在 actor old-logprob forward 中捕获 hidden states。改造前,大 chunk 通过: - -```python -chunk_ref = ray.put(hidden_chunk) -``` - -sample 不直接带 tensor,而是带: - -```python -{ - "hidden_states_ref_chunks": [ - { - "ref": ray_object_ref, - "start": 0, - "length": 512, - }, - ], -} -``` - -drafter worker 再 `ray.get(ref)`,根据 `start/length` 切出每条 sample 所需行。 - -### 2.3 PR #48 要改变的边界 - -PR #48 没有改变: - -- rollout 什么时候产生 sample; -- driver 如何触发 drafter worker; -- drafter worker 如何调用 `_store_rollout_sample()`; -- drafter model 的训练逻辑; -- drafter 权重发布。 - -它只改变大 tensor 的跨进程介质: - -```text -改造前:Producer → Ray driver/object store → Consumer -改造后:Producer → TQ storage → Consumer - key 仍走原 Ray 控制路径 -``` - -## 3. PR #48 改造后的完整 TQ 流程 - -### 3.0 总览 - -PR #48 的目标不是让 TQ 自己产生训练 batch,而是把原来经过 Ray driver/Ray object store 的大 hidden tensor 攁到 TQ。原来的控制路径继续存在,只是控制路径上从“大 tensor”变成“小 key”。 - -整体结构是: - -```text - 原 Ray 控制路径 - drafter_sample / chunk ref -Producer ───────────────────── key ───────────────────▶ Consumer - │ │ - │ kv_put(large tensor) │ kv_batch_get(key) - ▼ ▼ -TransferQueue storage ─────────────────────────────────────┘ -``` - -因此 PR #48 同时保留两条通道: - -```text -控制通道:Producer → Ray driver → drafter worker -数据通道:Producer → TQ storage → drafter worker -``` - -控制通道负责告诉 Consumer “本次应该处理哪个样本、对应哪个 TQ key”;数据通道负责传输 hidden states 等大 tensor。 - -#### 3.0.1 配置放在哪里 - -PR #48 在 drafter training 配置下增加: - -```yaml -actor_rollout_ref: - rollout: - drafter: - training: - transfer_queue: - enable: false - backend: - storage_backend: SimpleStorage - SimpleStorage: - total_storage_size: 100000 - num_data_storage_units: 8 -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/config/speco_base.yaml -``` - -这里的 `transfer_queue` 是 SpeCo drafter 自己的配置,不是 standalone `feature_store.type`,也不是简单把上游 verl 的 `transfer_queue.enable` 打开。 - -#### 3.0.2 TaskRunner 创建整套 TQ - -RL 任务启动时,`SpecoTaskRunner.run()` 在创建 workers 之前调用: - -```python -from verl_speco.integration.transferqueue_bridge import ( - close_transfer_queue, - init_transfer_queue, -) - -transfer_queue_started = init_transfer_queue(config) -try: - trainer.init_workers() - trainer.fit() -finally: - if transfer_queue_started: - close_transfer_queue() -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/integration/task_runner.py:319 -``` - -`init_transfer_queue(config)` 内部读取: - -```python -config.actor_rollout_ref.rollout.drafter.training.transfer_queue -``` - -然后执行: - -```python -tq.init(_to_plain_dict(tq_cfg)) -``` - -并记录: - -```python -_state["initialized"] = True -_state["owner"] = True -``` - -这里的 owner 是“创建/拥有 TQ 生命周期的进程”。只有 owner 在任务结束时执行 `tq.close()`。 - -关键顺序是: - -```text -SpecoTaskRunner -→ tq.init(完整配置) -→ trainer.init_workers() -→ Ray workers 启动 -``` - -也就是说,PR #48 假设 TQ Controller/storage 已经由 TaskRunner 在 Ray 集群环境中建立,后启动的 worker 只需要连接。 - -#### 3.0.3 每个 Producer/Consumer 进程怎么连接 TQ - -bridge 中的 `_ensure_initialized()` 是进程级懒初始化: - -```python -def _ensure_initialized(): - if _state["initialized"]: - return - - with _state_lock: - if _state["initialized"]: - return - - tq.init() - _state["initialized"] = True -``` - -注意这里是: - -```python -tq.init() -``` - -不是: - -```python -tq.init(config) -``` - -无参初始化的含义是连接 TaskRunner 已创建的同一套 TQ。它依赖 PR #48 所处的 Ray 运行环境完成服务发现。 - -因此 PR #48 不是“每个 worker 各创建一套 TQ”,而是: - -```text -TaskRunner:tq.init(config),创建一次 -SGLang producer:tq.init(),连接 -actor producer:tq.init(),连接 -drafter consumer:tq.init(),连接 -``` - -#### 3.0.4 SGLang Producer 怎么写 hidden states - -SGLang 已经完成 rollout,并组装出 `drafter_sample` 后,PR #48 执行: - -```python -configure_transfer_queue(training_cfg) - -if is_transfer_queue_enabled(): - tq_key = make_sample_key( - collection_global_steps, - self.replica_rank, - request_id, - ) - - tq_payload = { - "hidden_states": hidden_states.unsqueeze(0).cpu(), - } - - if target_logprobs is not None: - tq_payload["target_logprobs"] = ( - target_logprobs.unsqueeze(0).cpu() - ) - - put_sample( - tq_key, - tq_payload, - tag={ - "global_step": collection_global_steps, - "replica_rank": self.replica_rank, - }, - ) - - drafter_sample["hidden_states_tq_key"] = tq_key - drafter_sample["hidden_states"] = None -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/integration/sglang_runtime.py:2020 -``` - -这里发生了两条不同的数据流: - -```text -大 tensor:SGLang → TQ -小字典/key:SGLang → 原 Ray/driver 路径 → drafter worker -``` - -写进 TQ 后将: - -```python -drafter_sample["hidden_states"] = None -``` - -是为了避免 hidden states 继续经过 driver/Ray object store。driver 仍收到 sample,但大 tensor 已替换成: - -```python -drafter_sample["hidden_states_tq_key"] -``` - -#### 3.0.5 `put_sample()` 实际怎么写 - -bridge 中: - -```python -def put_sample(key, tensor_dict, *, tag=None): - payload = { - k: v - for k, v in tensor_dict.items() - if torch.is_tensor(v) - } - - _ensure_initialized() - - tq.kv_put( - key=key, - partition_id="speco_drafter_features", - fields=payload, - tag=tag or {}, - ) -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/integration/transferqueue_bridge.py:206 -``` - -这里可以明确看到: - -- PR #48 使用 TQ 高层 KV API; -- 一个 key 对应一个 sample; -- `fields` 是 tensor 字典; -- `tag` 是小 metadata; -- partition 当前写死为 `speco_drafter_features`; -- 写入前 tensor 已 `.cpu()`; -- 写入失败直接抛异常,不静默回退。 - -key 的生成代码是: - -```python -def make_sample_key(global_step, replica_rank, request_id): - return f"speco:{global_step}:{replica_rank}:{request_id}" -``` - -这个 key 对 RL rollout 是合理的,因为它用 step、rollout replica 和 request ID 标识一次在线采样。 - -#### 3.0.5.1 PR #48 中“数据”和“元数据”分别长什么样 - -PR #48 一条 SGLang 样本在写 TQ 之前,`drafter_sample` 同时包含大 tensor 和控制信息。简化后类似: - -```python -drafter_sample = { - # 普通训练输入,仍走原 sample/Ray 控制路径 - "input_ids": int64[1, prompt_len + response_len], - "prompts": int64[1, prompt_len], - "responses": int64[1, response_len], - - # 大 tensor,开启 TQ 后从这个字典移除 - "hidden_states": bf16[1, hidden_rows, hidden_dim], - "target_logprobs": fp32[1, rows, topk_or_vocab] | None, - - # hidden 与 token 对齐所需的小字段 - "hidden_positions": int64[1, hidden_rows] | None, - "hidden_position_start": int, - "hidden_position_end": int, - "hidden_window_start": int, - "hidden_window_end": int, - - # 控制信息 - "global_step": int, - "replica_rank": int, -} -``` - -执行 `put_sample()` 时,并不是把整个 `drafter_sample` 放进 TQ。PR #48 只抽取占用大的 tensor: - -```python -tq_payload = { - "hidden_states": bf16[1, hidden_rows, hidden_dim], - "target_logprobs": fp32[1, rows, ...], # 可选 - "hidden_raw_target_logprobs": ..., # 可选 - "hidden_raw_target_logprobs_positions": ..., # 可选 -} -``` - -这就是 TQ 的 data payload。它被传给: - -```python -tq.kv_put(fields=tq_payload) -``` - -另外还有 TQ tag: - -```python -tag = { - "global_step": 42, - "replica_rank": 1, -} -``` - -tag 是 TQ 侧轻量 metadata,用于描述/检索对象,不承载 hidden tensor。 - -写入完成后,仍经 Ray 传递的轻量 `drafter_sample` 变成: - -```python -drafter_sample = { - "input_ids": int64[1, total_len], - "prompts": int64[1, prompt_len], - "responses": int64[1, response_len], - - "hidden_states": None, - "target_logprobs": None, - "hidden_states_tq_key": "speco:42:1:req-007", - - "hidden_positions": int64[1, hidden_rows] | None, - "hidden_position_start": 128, - "hidden_position_end": 640, - "global_step": 42, - "replica_rank": 1, -} -``` - -因此 PR #48 实际存在三类对象: - -| 对象 | 内容 | 传输路径 | 作用 | -|---|---|---|---| -| TQ fields/payload | `hidden_states` 等大 tensor | Producer → TQ storage → Consumer | 避免 driver 搬运大 tensor | -| TQ tag | `global_step`、`replica_rank` | TQ control/index metadata | 描述该 key | -| 轻量 `drafter_sample` | tokens、位置、标量 metadata、`hidden_states_tq_key` | 原 Ray driver 路径 | 告诉 Consumer 取哪个 key,以及如何解释取回的 tensor | - -代码实现解耦的关键不是“所有内容都进 TQ”,而是: - -```python -drafter_sample["hidden_states_tq_key"] = tq_key -drafter_sample["hidden_states"] = None -``` - -第一行给 Consumer 留下寻址信息;第二行阻止大 tensor 继续沿旧路径传输。 - -#### 3.0.5.2 Consumer 如何把两部分重新合成一个训练样本 - -Consumer 最初拿到的是轻量 sample: - -```python -sample["hidden_states"] is None -sample["hidden_states_tq_key"] == "speco:42:1:req-007" -``` - -它执行: - -```python -payload = get_sample(sample["hidden_states_tq_key"]) -sample["hidden_states"] = payload["hidden_states"] -``` - -合并后: - -```python -sample = { - "input_ids": ..., - "prompts": ..., - "responses": ..., - "hidden_positions": ..., - "hidden_states": bf16[1, hidden_rows, hidden_dim], - "hidden_states_tq_key": "speco:42:1:req-007", - ... -} -``` - -后面的 `_store_rollout_sample()` 看到的结构与关闭 TQ 时基本一致,所以训练主体不需要增加 TQ 分支。TQ bridge 只改变“大 tensor 从哪里恢复”,不改变 drafter trainer 的输入语义。 - -#### 3.0.6 old-logprob Producer 怎么写 chunk - -PR #48 的另一条 Producer 路径来自 actor old-logprob hidden hook。原来是: - -```python -chunk_ref = ray.put(hidden_chunk) -``` - -开启 TQ 后改成: - -```python -tq_key = make_sample_key( - global_step, - owner, - f"chunk{len(chunk_refs)}", -) - -put_sample( - tq_key, - {"hidden": hidden_chunk}, - tag={ - "global_step": global_step, - "owner": owner, - }, -) - -chunk_ref = tq_key -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/integration/oldlogprob_runtime.py:533 -``` - -后面的 driver 不需要知道 ref 是 Ray ObjectRef 还是 TQ string key,它只把 ref 当不透明 token 继续传递。 - -#### 3.0.7 Consumer 怎么根据 key 读取 - -drafter worker 收到原来的 sample 小字典后: - -```python -tq_key = sample.get("hidden_states_tq_key") - -if tq_key is not None and self._speco_tq_enabled: - payload = get_sample(tq_key) - - for field in ( - "hidden_states", - "target_logprobs", - "hidden_raw_target_logprobs", - "hidden_raw_target_logprobs_positions", - ): - if payload.get(field) is not None: - sample[field] = payload[field] - - if sample.get("hidden_states") is None: - raise RuntimeError( - "TQ key exists but hidden_states payload is missing" - ) -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/workers/speco_worker.py:848 -``` - -恢复 tensor 后,后面的逻辑仍使用原 `sample`/`batch`,drafter trainer 不需要知道 tensor 来自 Ray 还是 TQ。 - -#### 3.0.8 `get_sample()` 实际怎么读 - -```python -def get_sample(key): - _ensure_initialized() - - result = tq.kv_batch_get( - keys=[key], - partition_id="speco_drafter_features", - ) - - value = _extract_value(result, key) - return _tensordict_to_dict(value) -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/integration/transferqueue_bridge.py:237 -``` - -`_extract_value()` 兼容三种返回形态: - -```python -if isinstance(result, dict): - return result.get(key) -if isinstance(result, (list, tuple)): - return result[0] -return result -``` - -这是因为不同 TQ 版本/后端返回包装可能不同。 - -#### 3.0.9 为什么需要 `_densify_tq_tensor()` - -PR #48 后续修复发现,TQ 把单样本 tensor 放进 TensorDict 后,`kv_batch_get` 可能返回 NestedTensor,并额外带 batch 维。旧代码要执行: - -```python -tensor[start:start + length] -``` - -但 NestedTensor 不支持在 jagged dim 直接 slice。因此加入: - -```python -def _densify_tq_tensor(tensor): - if tensor.is_nested: - parts = [ - part - for part in tensor.unbind() - if part.numel() > 0 - ] - tensor = torch.cat(parts, dim=0) - - if tensor.dim() == 3: - tensor = tensor.squeeze(0) - elif tensor.dim() == 1: - tensor = tensor.unsqueeze(0) - - return tensor.contiguous() -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/workers/speco_worker.py:72 -``` - -standalone Consumer 同样必须做这个转换,不能假定 `kv_batch_get` 返回普通 dense `[seq, hidden]`。 - -#### 3.0.10 为什么需要 per-step cache - -old-logprob 路径里,一个 owner hidden chunk 可能被约 16 个 sample 共同引用。如果每个 sample 都: - -```python -get_sample(same_tq_key) -``` - -就会重复传输同一个数百 MB chunk。PR #48 在每次 `collect_rollout_features()` 开始时创建: - -```python -self._tq_chunk_cache = {} -``` - -解析 ref 时: - -```python -cache_key = ref if isinstance(ref, str) else id(ref) - -if cache_key not in cache: - cache[cache_key] = _resolve_tq_or_ray_ref(ref) - -tensor = cache[cache_key] -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/workers/speco_worker.py:98 -../verl-SpeCo/verl_speco/workers/speco_worker.py:854 -``` - -独立训练如果每个 sample 都是独立 TQ key,主要使用 `kv_batch_get(keys=[...])` 批量读取,不一定需要跨 sample chunk cache;但 prefetch 重试或共享 packed object 时仍应保留 key cache。 - -#### 3.0.11 PR #48 什么时候删除数据 - -PR #48 没有在 `get_sample()` 后删除。其注释明确说明:同一 drafter replica 的多个 TP/SP rank 可能读取同一 key,第一次读取后立刻删除会让剩余 rank 失败。 - -当前策略是任务结束时由 owner: - -```python -tq.close() -``` - -统一结束 TQ 生命周期。也就是说,PR #48 当前没有实现精细的逐 step `kv_clear`。 - -#### 3.0.12 PR #48 的完整时序 - -```text -SpecoTaskRunner - → tq.init(config) - → 启动 Ray workers - -SGLang/actor Producer process - → configure_transfer_queue() - → 第一次 put 时 tq.init() - → kv_put(key, tensor fields, tag) - → 把 key 塞回原 sample/ref - -Ray driver - → 只中转小 sample/key - -drafter worker Consumer process - → 第一次 get 时 tq.init() - → kv_batch_get([key]) - → 解包 TensorDict/NestedTensor - → 恢复 sample["hidden_states"] - → 原 drafter collect/train 逻辑 - -任务结束 - → TaskRunner owner tq.close() -``` - -### 3.1 已经实现的可复用能力 - -PR #48 新增 `verl_speco/integration/transferqueue_bridge.py`,锁定思路是把 TQ 当成独立传输库,不修改上游 verl。它提供: - -```python -configure_transfer_queue(training_cfg) -init_transfer_queue(config) -make_sample_key(global_step, replica_rank, request_id) -put_sample(key, tensor_dict, tag=...) -get_sample(key) -close_transfer_queue() -``` - -实际写入调用是: - -```python -tq.kv_put( - key=key, - partition_id="speco_drafter_features", - fields=payload, - tag=tag, -) -``` - -实际读取调用是: - -```python -result = tq.kv_batch_get( - keys=[key], - partition_id="speco_drafter_features", -) -``` - -另外,PR #48 已经处理了多项 standalone 方案也需要的问题: - -1. TQ 返回值可能是 direct value、`{key: value}` 或 list,需要统一解包; -2. TQ 返回的 tensor 可能是 NestedTensor,必须通过 `unbind + cat` 恢复为生产端写入的 dense tensor; -3. 多个样本引用同一 hidden chunk 时,需要 per-step cache,避免重复 `kv_batch_get` 同一个大对象; -4. TQ key 存在但 payload 缺失时 fail loud,不能静默丢样本; -5. `enable=false` 时保留原传输路径。 - -这些逻辑应直接作为本项目 TQ adapter 的参考。 - -### 3.2 PR #48 的数据流 - -PR #48 优化的是 RL online 路径: - -```text -SGLang/actor worker - → kv_put(hidden states) - → 把 hidden_states_tq_key 塞进原 drafter_sample - → 原 Ray driver 继续传递小 sample/key - → drafter worker collect_rollout_features() - → kv_batch_get(key) -``` - -它没有让 consumer 自己从 TQ 发现下一批 key;key 仍沿原来的 Ray 控制路径到达 drafter worker。 - -### 3.3 PR #48 没有提供的 standalone 能力 - -PR #48 当前没有实现: - -- 从预生成 response 文件读取数据的独立 Producer; -- Producer 并行请求外部 vLLM endpoint; -- standalone DSpark trainer 主动发现 ready key; -- global batch 到各 torchrun rank 的分片; -- 每个 optimizer step 后精确 `kv_clear`; -- EOS; -- standalone 无 Ray 的 TQ bootstrap; -- MooncakeStore 的实际运行验证。 - -PR #48 当前配置是: - -```yaml -transfer_queue: - enable: false - backend: - storage_backend: SimpleStorage - SimpleStorage: - total_storage_size: 100000 - num_data_storage_units: 8 -``` - -并且 bridge 注释针对 `TransferQueue==0.1.7`。旁边当前 verl 主线已使用 `0.1.8` 文案并包含 `MooncakeStore` 配置。当前机器没有安装 `transfer_queue` 包,因此正式实现前必须锁定版本并实机验证 API 签名,不能把 0.1.7 和 0.1.8 混用。 - -### 3.4 standalone 方案对 PR #48 的扩展 - -不再自行实现 `publish_ready/claim/lease/ack` HTTP 服务。第一版在 PR #48 KV 模式上补四个操作: - -```python -tq.kv_batch_put(...) # Producer 批量写 -tq.kv_list(...) # rank 0 列出 key + tag -tq.kv_batch_get(...) # 各 rank 并行读 -tq.kv_clear(...) # optimizer step 成功后删 -``` - -第一版不依赖 TQ `Sampler/StreamingDataLoader/get_meta`,因为 PR #48 并未使用或验证这些接口。等 KV 流程稳定后再升级为 TQ StreamingDataLoader。 - -### 3.5 PR #48 与 standalone 独立训练逐项映射 - -| PR #48 online RL | standalone drafter training | -|---|---| -| `SpecoTaskRunner` 调用 `tq.init(config)` | 新增独立 TQ owner/bootstrap 进程,或在确认安全后由 Producer owner 调用 `tq.init(config)` | -| SGLang/actor worker 是 Producer | `feature_producer.py` 是独立 Producer | -| rollout 过程中已经得到 hidden states | Producer 读取预生成 response,再请求外部 vLLM prefill | -| `make_sample_key(global_step, replica_rank, request_id)` | `make_replay_sample_key(dataset,row,tokens,target fingerprint)` | -| `put_sample()` / `kv_put()` | 继续复用同一写入模式,可扩展为 `kv_batch_put()` | -| key 通过 Ray driver/sample 传给 consumer | 没有 driver;rank 0 用 `kv_list()` 主动发现 ready keys | -| drafter worker `get_sample(key)` | 每个 torchrun rank `kv_batch_get(local_keys)` | -| `collect_rollout_features()` 恢复 sample | 转换为 `DraftFeatureSample` 后调用现有 `prepare_training_batch_from_samples()` | -| 多 TP/SP rank 可能读同一 key,所以不立即删除 | data-parallel rank 读取互不重叠的 key;全 rank step 成功后统一 `kv_clear(global_keys)` | -| TaskRunner 结束时 `tq.close()` | 每 step clear;输入 drain 完成后 owner 最后 `tq.close()` | - -standalone 需要新增的控制流是: - -```text -Producer DSpark rank 0 其他 ranks - │ │ │ - │ kv_put(sample key, fields, tag) │ │ - ├────────────────────────────────────▶│ │ - │ │ kv_list READY keys │ - │ │ │ - │ │ broadcast selected_keys ──▶│ - │ │ │ - │ │ kv_batch_get(local keys) │ kv_batch_get(local keys) - │ │ │ - │ ├──── DSpark synchronized step ────┤ - │ │ │ - │ │ kv_clear(global keys) │ -``` - -这个映射中,TQ 同时承担: - -- 大 tensor 存储/传输; -- key、tag 和 partition 的轻量索引。 - -但第一版 global batch 的选择仍由单个 DSpark job 的 rank 0 完成。这样最接近 PR #48 的 KV API,避免在同一次改造中再引入未经该 PR 验证的 Sampler/StreamingDataLoader。 - -### 3.6 PR #48 的 SGLang 路径:逐函数、逐对象完整流程 - -下面从一次生成请求开始,不省略中间层。 - -#### 阶段 1:SGLang完成生成并收集 hidden states - -执行进程:SGLang rollout server。 - -输入是一次 request 对应的 prompt 和生成配置。生成结束时,代码已经持有: - -```python -prompt_tensor: int64[prompt_len] -response_tensor: int64[response_len] -hidden_states: bf16[hidden_rows, hidden_dim] -hidden_positions: int64[hidden_rows] | None -target_logprobs: tensor | None -request_id: str -collection_global_steps: int -self.replica_rank: int -``` - -这些变量的语义: - -- `prompt_tensor`:输入 prompt token IDs; -- `response_tensor`:SGLang生成的 response token IDs; -- `hidden_states`:target model 指定层在部分 token positions 上的输出; -- `hidden_positions`:每个 hidden row 对应完整 `prompt+response` 序列中的哪个 token position; -- `target_logprobs`:可选的目标概率监督; -- `request_id`:当前 rollout request 标识; -- `replica_rank`:执行该 request 的 rollout replica。 - -SGLang 先构造完整 sample: - -```python -drafter_sample = { - "input_ids": torch.cat( - [prompt_tensor, response_tensor], dim=0 - ).unsqueeze(0), - "prompts": prompt_tensor.unsqueeze(0), - "responses": response_tensor.unsqueeze(0), - "hidden_states": hidden_states.unsqueeze(0).cpu(), - "hidden_positions": hidden_positions.unsqueeze(0).cpu(), - "target_logprobs": ( - target_logprobs.unsqueeze(0).cpu() - if target_logprobs is not None - else None - ), - "global_step": collection_global_steps, - "replica_rank": self.replica_rank, - # 还有 hidden window/alignment metadata -} -``` - -前导维 `1` 表示这是一个单样本 batch。`.cpu()` 表示跨进程传输前把大 tensor 放到 CPU 内存。 - -#### 阶段 2:PR #48 将大 fields 写入 TQ - -同一个 SGLang进程执行: - -```python -tq_key = make_sample_key( - collection_global_steps, - self.replica_rank, - request_id, -) - -tq_payload = { - "hidden_states": hidden_states.unsqueeze(0).cpu(), -} - -put_sample( - tq_key, - tq_payload, - tag={ - "global_step": collection_global_steps, - "replica_rank": self.replica_rank, - }, -) -``` - -调用展开后是: - -```python -tq.init() # 当前进程第一次使用时 -tq.kv_put( - key=tq_key, - partition_id="speco_drafter_features", - fields=tq_payload, - tag=tag, -) -``` - -效果是 TQ 中增加一行: - -```text -partition = speco_drafter_features -key = speco:42:1:req-007 -fields = {hidden_states: bf16[1, H, D], ...} -tag = {global_step: 42, replica_rank: 1} -``` - -`kv_put` 返回后,SGLang侧将旧 sample 改成: - -```python -drafter_sample["hidden_states_tq_key"] = tq_key -drafter_sample["hidden_states"] = None -``` - -此后该 Python 字典不再携带 hidden tensor,只携带定位它的 key。 - -#### 阶段 3:把 drafter_sample 放进 TokenOutput.extra_fields - -SGLang返回: - -```python -TokenOutput( - token_ids=token_ids, - log_probs=log_probs, - routed_experts=routed_experts, - extra_fields={ - "global_steps": collection_global_steps, - "drafter_sample": drafter_sample, - }, -) -``` - -此时 `TokenOutput` 中有两类输出: - -- 正常 rollout 输出:`token_ids/log_probs`; -- SpeCo 训练旁路输出:`extra_fields.drafter_sample`。 - -TQ payload 不在 `TokenOutput` 中,只有 `hidden_states_tq_key` 在其中。 - -#### 阶段 4:多个 TokenOutput 聚合成 gen_batch_output - -rollout/agent-loop 层将多个 request 的输出合并为批量 `DataProto`: - -```python -gen_batch_output.non_tensor_batch["drafter_sample"] -``` - -可能是 object array: - -```python -array([ - {"hidden_states_tq_key": "speco:42:0:req-A", ...}, - {"hidden_states_tq_key": "speco:42:1:req-B", ...}, -], dtype=object) -``` - -之所以进入 `non_tensor_batch`,是因为每条 sample 的 hidden window、Python metadata 和可选字段不一定具有统一 dense shape。 - -#### 阶段 5:driver 从 DataProto 取出 drafter samples - -`generate_sequences_with_speco()` 包装原 rollout 调用: - -```python -gen_batch_output = original_generate_sequences(...) -collected = self._speco_collect_generation_samples(gen_batch_output) -``` - -`_speco_collect_generation_samples()` 调用: - -```python -samples = pop_drafter_samples(gen_batch_output) -``` - -`pop_drafter_samples()` 实际执行: - -```python -non_tensor_batch = gen_batch_output.non_tensor_batch -samples_array = non_tensor_batch.pop("drafter_sample", None) -samples = normalize_drafter_samples(samples_array) -``` - -这里 `pop` 有两个作用: - -1. 取得 SpeCo drafter side-channel samples; -2. 从正常 PPO 的 `gen_batch_output` 中移除该旁路字段,避免后续 PPO batch 继续携带它。 - -`normalize_drafter_samples()` 将 dict、object array 或 list 统一成: - -```python -samples: list[dict] -``` - -#### 阶段 6:driver 按 replica_rank 分桶 - -假设有两个 rollout/drafter replicas,收到: - -```python -samples = [ - {"replica_rank": 1, "hidden_states_tq_key": "k1", ...}, - {"replica_rank": 0, "hidden_states_tq_key": "k2", ...}, - {"replica_rank": 1, "hidden_states_tq_key": "k3", ...}, -] -``` - -执行: - -```python -buckets = bucket_drafter_samples_by_replica( - samples, - num_replicas=2, -) -``` - -结果: - -```python -buckets = [ - [sample_k2], # bucket 0 - [sample_k1, sample_k3], # bucket 1 -] -``` - -分桶依据只有: - -```python -owner_rank = int(sample["replica_rank"]) -buckets[owner_rank].append(sample) -``` - -这一步没有读取 TQ,也没有处理 hidden tensor;只对小字典做路由。 - -#### 阶段 7:driver 通过 WorkerGroup RPC 分发 buckets - -driver 调用: - -```python -self._speco_set_drafter_global_step() -self._speco_collect_rollout_features_rpc( - "rollout", - buckets, -) -``` - -RPC 内部调用: - -```python -self.drafter_wg.collect_rollout_features(buckets) -``` - -因为 worker 方法注册了 `drafter_owner_route` dispatch,WorkerGroup 将 `buckets[0]` 发给 drafter DP replica 0,将 `buckets[1]` 发给 drafter DP replica 1。一个训练 replica 内的 SP ranks 根据 mesh dispatch 规则参与对应调用。 - -这里传输的对象仍是: - -```python -list[dict] -``` - -其中包含 tokens、position metadata 和 TQ key,不包含被置空的 hidden states。 - -#### 阶段 8:SpecoWorker 根据 key 从 TQ 恢复 tensor - -目标 worker 执行: - -```python -def collect_rollout_features(self, samples): - for sample in samples: - tq_key = sample.get("hidden_states_tq_key") - payload = get_sample(tq_key) - sample["hidden_states"] = payload["hidden_states"] -``` - -`get_sample()` 展开为: - -```python -tq.init() # 此 Consumer 进程第一次使用时 -result = tq.kv_batch_get( - keys=[tq_key], - partition_id="speco_drafter_features", -) -payload = _extract_value(result, tq_key) -payload = _tensordict_to_dict(payload) -``` - -现在 `sample` 再次包含: - -```python -{ - "input_ids": ..., - "hidden_positions": ..., - "hidden_states": bf16[1, hidden_rows, hidden_dim], - "hidden_states_tq_key": "...", -} -``` - -这与关闭 TQ 时 worker 收到的逻辑内容一致。 - -#### 阶段 9:worker 构造 DrafterBaseTrainer 所需 batch dict - -worker 先保留 token fields: - -```python -batch = { - "input_ids": sample["input_ids"], - "prompts": sample["prompts"], - "responses": sample["responses"], -} -``` - -再复制 hidden alignment metadata,例如: - -```python -batch["hidden_positions"] -batch["hidden_position_start"] -batch["hidden_position_end"] -batch["hidden_states_layout"] -batch["global_step"] -``` - -hidden tensor 单独作为参数: - -```python -self._store_rollout_sample( - batch=batch, - hidden_states=hidden, - target_logprobs=target_logprobs, -) -``` - -#### 阶段 10:样本进入在线 buffer 或落盘 - -`_store_rollout_sample()` 根据 training mode 分支: - -```python -if mode == "collect_only": - self._write_rollout_feature_sample( - batch, - hidden_states, - target_logprobs, - ) -else: - self.trainer.collect_online_data( - batch, - hidden_states, - target_logprobs, - ) -``` - -`collect_only` 会转换成 `DraftFeatureSample` 并写 `TorchShardFeatureStore`。online 模式则进入 `DrafterBaseTrainer.collect_online_data()`。 - -`collect_online_data()` 做: - -1. 将 `input_ids/hidden_states/positions/logprobs` 规范化到 CPU; -2. 按 batch 维拆成逐样本; -3. 根据 `hidden_positions` 校验 hidden row 与 token position; -4. 截取可训练窗口; -5. 构造内部 training item; -6. 保存到当前 step 的 `collected_data`,或在启用 data buffer 时保存到跨 step buffer。 - -因此 TQ get 完成不代表立即 optimizer step。它先恢复 online training sample,再进入现有数据准备逻辑。 - -#### 阶段 11:driver 在 actor update 周期触发 drafter 训练 - -driver 包装了 `update_actor()`: - -```python -should_train_drafter = ( - self._speco_should_attempt_drafter_train_this_step() -) - -actor_output = original_update_actor(...) - -if should_train_drafter: - drafter_trained, metrics = self._speco_train_drafter() -``` - -`_speco_train_drafter()` 再向 WorkerGroup 发: - -```python -self.drafter_wg.train_drafter() -``` - -每个 `SpecoWorker.train_drafter()`: - -1. 检查是否属于 drafter training group; -2. 检查 `training_interval_steps`; -3. 激活 drafter training model; -4. 循环 `train_steps_per_trigger` 次; -5. 每次调用 `self.trainer.training_step(global_step)`; -6. 成功时准备需要发布的 drafter state dict; -7. 清理训练期间临时状态。 - -`training_step()` 从刚才的 online `collected_data/DataBuffer` 组成 batch,执行 drafter forward、loss、backward 和 optimizer step。 - -所以 SGLang TQ 路径的最终效果是: - -```text -TQ 只替换 hidden tensor 跨进程传输 -→ sample 收集逻辑不变 -→ online buffer 不变 -→ drafter training trigger 不变 -→ loss/optimizer 不变 -``` - -### 3.7 PR #48 old-logprob 路径的完整差异 - -old-logprob 路径没有 `TokenOutput.extra_fields`。它从 PPO 的 actor old-logprob forward 开始。 - -#### 阶段 1:driver 构造 collect plan - -driver 根据 batch、collect interval 和 drafter owner 数量决定: - -```python -collect_mask: bool[batch] -hidden_positions: list/tensor per sample -owner_rank: int64[batch] -prompt_lens: int64[batch] -response_lens: int64[batch] -``` - -并把 `global_step` 等控制字段放入 old-logprob micro-batch。 - -#### 阶段 2:actor forward hook 选择 hidden rows - -actor worker 在 old-logprob forward 中捕获指定层 hidden states,根据 `collect_mask/position_mask` 只保留需要训练的样本和 token rows。 - -输出可以是 dense selected tensor,也可以是 sparse rows。随后 `_put_oldlogprob_hidden_refs()` 将同一个 owner 的多条 sample rows 拼成一个较大的 `hidden_chunk`。 - -#### 阶段 3:hidden chunk 写入 TQ - -改造前: - -```python -chunk_ref = ray.put(hidden_chunk) -``` - -PR #48: - -```python -tq_key = make_sample_key( - global_step, - owner, - f"chunk{chunk_index}", -) - -put_sample( - tq_key, - {"hidden": hidden_chunk}, - tag={"global_step": global_step, "owner": owner}, -) - -chunk_ref = tq_key -``` - -TQ fields: - -```python -{"hidden": bf16[total_owner_rows, hidden_dim]} -``` - -控制路径中的 chunk metadata: - -```python -chunk_info = { - "sample_indices": [0, 3, 5], - "starts": [0, 128, 384], - "lengths": [128, 256, 96], - "row_indices": [...], - "dtype": "bfloat16", - "shape": [480, hidden_dim], -} -``` - -`starts/lengths` 描述每条 sample 在共享 chunk 中对应的行区间。 - -#### 阶段 4:driver 将 chunk ref 还原为逐样本引用 - -driver 的 `_speco_collect_oldlogprob_features()` 读取: - -```python -chunk_refs = ["speco:42:0:chunk0", ...] -chunk_meta = [chunk_info, ...] -``` - -然后为每个 batch sample 构造: - -```python -sample["hidden_states_ref_chunks"] = [ - { - "ref": "speco:42:0:chunk0", - "chunk_start": 128, - "chunk_length": 256, - "chunk_row_indices": ..., - "dtype": "bfloat16", - "shape": [480, hidden_dim], - } -] -``` - -同时构造该 sample 的: - -```python -input_ids -prompts -responses -hidden_positions -hidden_states_layout -replica_rank=owner -``` - -再按 owner 放入 `buckets[owner]`,通过同一个 `collect_rollout_features()` RPC 发给 drafter worker。 - -#### 阶段 5:Consumer 获取共享 chunk 并切片 - -drafter worker 发现: - -```python -sample.get("hidden_states") is None -sample.get("hidden_states_ref_chunks") is not None -``` - -于是调用 `_resolve_hidden_state_chunks()`。对字符串 ref: - -```python -if ref.startswith("speco:"): - full_chunk = get_sample(ref)["hidden"] - full_chunk = _densify_tq_tensor(full_chunk) -``` - -然后按 sample metadata 取行: - -```python -sample_hidden = full_chunk[ - chunk_start : chunk_start + chunk_length -] -``` - -同一个 chunk 被多个 sample 复用,所以使用: - -```python -self._tq_chunk_cache[ref] = full_chunk -``` - -保证一次 `collect_rollout_features()` 中同一个 TQ key 只 get 一次。 - -得到逐样本 hidden 后,后续 `_store_rollout_sample → collect_online_data → train_drafter` 与 SGLang 路径相同。 - -### 3.8 PR #48 数据生命周期和清理 - -PR #48 的 TQ row 生命周期是: - -```text -TaskRunner tq.init(config) -→ Producer kv_put -→ key 经 Ray 控制路径传递 -→ 一个或多个 drafter TP/SP rank kv_batch_get -→ online drafter 收集/训练继续执行 -→ 整个 trainer.fit() 结束 -→ TaskRunner finally 调用 tq.close() -``` - -当前没有: - -```python -tq.kv_clear(key) -``` - -原因是同一 sample/key 可能被一个 drafter replica 的多个 rank 读取。若第一个 rank get 后删除,其余 rank 可能 get 失败。 - -因此 PR #48 采用任务级生命周期,而不是样本级确认和回收。这简化了并发正确性,但意味着长任务中的 TQ storage 占用可能持续增长;代码注释也把“leader 在 barrier 后精细 clear”留作后续工作。 - -### 3.9 PR #48 开启与关闭时的行为差异 - -`configure_transfer_queue()` 返回: - -```python -enabled_in_config and transfer_queue_importable -``` - -关闭时: - -```text -SGLang drafter_sample 继续内联 hidden_states -old-logprob 继续 ray.put(hidden_chunk) -Consumer 继续 ray.get/ref resolve -``` - -开启时: - -```text -SGLang hidden fields → TQ,sample 只带 key -old-logprob hidden chunk → TQ,ref 变成字符串 key -Consumer 根据 key 类型走 TQ get -``` - -如果配置开启但 `transfer_queue` 包未安装,bridge 会记录 warning,并让 `is_transfer_queue_enabled()` 返回 false,保留旧 Ray 路径。若已经进入 `put_sample/get_sample` 却发生 TQ 错误,则抛异常,不静默丢 hidden data。 - -## 第二部分:基于 PR #48 的 standalone drafter training 适配 - -从这一部分开始才讨论 `verl-SpeCo-ls` 的独立训练。下面的 `kv_list`、rank 0 选 global keys、逐 step `kv_clear`、独立 vLLM Producer 都是需要在本项目新增的逻辑,不是 PR #48 已有逻辑。 - -### 当前 standalone 基线 - -当前独立训练是: - -```text -draft_train_launcher -→ torch.distributed.run -→ 每个 rank 创建 DraftFeatureDataLoader -→ 每个 rank 在训练循环内同步 TargetFeatureReplayer.materialize() -→ vLLM/file hidden payload -→ prepare_training_batch_from_samples() -→ training_step_from_batch() -``` - -新方案要把 `TargetFeatureReplayer.materialize()` 从训练 rank 的同步取数路径移到独立 Producer,同时保持后两步训练接口不变。 - -## 4. TQ metadata 到底记录什么 - -### 4.1 Partition - -一次训练运行使用一个独立 partition: - -```python -partition_id = f"speco:{run_id}:dspark_train" -``` - -partition 用来隔离: - -- 不同训练 run; -- train 和 validation; -- 不同 target checkpoint 生成的 hidden states。 - -不能让两个 target 模型共用同一 partition,否则训练侧可能消费错误的 hidden states。 - -### 4.2 Sample key - -每条输入样本使用稳定 key: - -```python -sample_key = sha256( - dataset_id - + row_id - + prompt_token_ids - + response_token_ids - + tokenizer_fingerprint - + target_model_fingerprint - + target_layer_ids - + hidden_states_layout -).hexdigest() -``` - -稳定 key 用于: - -- vLLM HTTP 请求重试时不生成不同对象; -- Producer 重启后识别相同样本; -- 检查 hidden states 是否属于正确模型和正确层; -- TQ/Mooncake 清理时准确定位对象。 - -### 4.3 Fields 与 READY 约定 - -每个样本包含固定字段: - -```python -{ - "input_ids": int64[seq], - "loss_mask": float32[seq], - "position_ids": int64[seq], - "hidden_states": bf16[seq, aux_hidden_dim or aux_hidden_dim+hidden_size], -} -``` - -这里必须保持当前 `DraftFeatureSample` 的契约:`TargetFeatureReplayer._feature_from_vllm_payload()` 会把多个 aux layers flatten;当 layout 是 `dflash_aux_plus_last` 时,还会把 final hidden 拼到同一个 `hidden_states` tensor 尾部。`DSparkTrainerBackend.preprocess_individual_items()` 再根据 metadata 中的 `hidden_states_layout` 将 final hidden 切出来,生成训练 batch 的 `target_last_hidden_states`。 - -因此 TQ 不需要新增一个独立 `target_last_hidden_states` field。开启 DSpark L1 时,要求: - -```text -metadata.hidden_states_layout = dflash_aux_plus_last -hidden_states.shape[-1] = num_context_layers * hidden_size + hidden_size -``` - -完整 TQ native metadata API 可以做字段级 ready 判定,但 PR #48 当前走的是高层 KV API:一次 `kv_put` 把一个样本的多个 tensor fields 一起写入。因此第一版采用更直接的约定: - -```python -required_fields = [ - "input_ids", - "loss_mask", - "position_ids", - "hidden_states", -] -``` - -Producer 先在内存中验证所有必需字段,再进行一次 `kv_put`。只有 `kv_put` 成功返回,key 才会出现在 `kv_list` 结果中,并带有: - -```python -tag={ - "status": "ready", - "run_id": run_id, - "sample_id": sample_key, -} -``` - -Consumer 只选择 `status=ready` 且 `run_id` 匹配的 key。不要先 put `input_ids`、再单独 put `hidden_states`,否则 Consumer 可能观察到半成品。 - -### 4.4 Tags - -tags 是轻量 metadata,不放大 tensor: - -```python -tags = { - "sample_id": sample_key, - "source_row": row_id, - "seq_len": seq_len, - "payload_bytes": payload_bytes, - "target_model_fp": target_model_fingerprint, - "target_layers": "8,16,24", - "hidden_layout": "dflash_aux_plus_last", - "producer_status": "success", -} -``` - -tags 可用于过滤、监控、背压统计和错误排查。它不能替代 tensor shape/dtype 校验。 - -### 4.5 Run ID,而不是先依赖 task_name - -PR #48 的 `kv_put/kv_batch_get` 路径没有使用 `task_name` 或 Sampler consumption history。第一版按它的已验证接口,在 partition 和 tags 中放 `run_id`: - -```python -partition_id = "speco_drafter_features" -tag = { - "run_id": run_id, - "status": "ready", -} -``` - -不同 run 最好直接使用不同 partition: - -```python -partition_id = f"speco_drafter_features_{run_id}" -``` - -这样 job 重启和清理更简单。`task_name=dspark_train` 留到后续迁移 TQ native metadata/Sampler 时再使用。 - -### 4.6 standalone 中一条样本的完整对象形态 - -standalone 没有 PR #48 的轻量 `drafter_sample → Ray driver` 路径,因此需要让 TQ tag 承担“如何找到和解释 payload”的 metadata 作用。 - -#### Producer 读到的原始记录 - -```python -source_record = { - "dataset_id": "math-train", - "row_id": 12345, - "prompt": "...", - "response": "已经提前生成的 response", -} -``` - -#### Token replay 样本 - -分词和对齐后: - -```python -replay_sample = DraftReplaySample( - input_ids=int64[full_seq], - loss_mask=float32[full_seq], - position_ids=int64[full_seq], - feature_positions=int64[feature_rows], - draft_position_ids=int64[feature_rows], - metadata={ - "dataset_id": "math-train", - "row_id": 12345, - }, -) -``` - -这里的 `full_seq` 是 prompt 与预生成 response 拼接后的长度;`feature_positions` 指明哪些 token 位置最终进入 drafter 监督样本。 - -#### vLLM 返回的原始 hidden payload - -当前文件协议要求 safetensors 至少包含: - -```python -vllm_payload = { - "token_ids": int64[prefill_rows], - "hidden_states": bf16[prefill_rows, returned_layers, hidden_size], -} -``` - -这还不能直接给 DSpark。Producer 应复用当前: - -```python -TargetFeatureReplayer._feature_from_vllm_payload(...) -``` - -完成 token 校验、position 对齐、选层和 flatten。 - -#### Producer 最终得到的 DraftFeatureSample - -```python -feature = DraftFeatureSample( - algorithm="DSpark", - input_ids=int64[feature_rows], - loss_mask=float32[feature_rows], - position_ids=int64[feature_rows], - hidden_states=bf16[feature_rows, feature_hidden_dim], - metadata={ - "hidden_states_layout": "dflash_aux_plus_last", - "target_layer_ids": [8, 16, 24], - "target_model_path": "...", - "target_config_fingerprint": "...", - "feature_start": 128, - "feature_end": 640, - "sequence_length": 512, - }, -) -``` - -若有 3 个 aux layer、target hidden size 为 4096,并包含 final hidden: - -```text -feature_hidden_dim = 3 * 4096 + 4096 = 16384 -hidden_states.shape = [feature_rows, 16384] -``` - -前 `12288` 维是 aux/context hidden,最后 `4096` 维是 DSpark L1 所需 final hidden。现有 DSpark backend 会根据 `dflash_aux_plus_last` 自动拆分。 - -#### 写入 TQ 的 data fields - -第一版建议一个 key 对应一个已经规范化完成的 `DraftFeatureSample`: - -```python -tq_fields = { - "input_ids": feature.input_ids.cpu(), - "loss_mask": feature.loss_mask.cpu(), - "position_ids": feature.position_ids.cpu(), - "hidden_states": feature.hidden_states.cpu(), -} -``` - -这里只放 tensor,因为 PR #48 的 `put_sample()` 会过滤非 tensor: - -```python -payload = { - key: value - for key, value in tensor_dict.items() - if torch.is_tensor(value) -} -``` - -#### 写入 TQ 的 tag metadata - -```python -tq_tag = { - "run_id": "run-20260818-001", - "status": "ready", - "sample_id": sample_key, - "sequence_no": 12345, - "algorithm": "DSpark", - "hidden_states_layout": "dflash_aux_plus_last", - "target_model_fingerprint": "sha256:...", - "target_layer_ids": "8,16,24", - "feature_rows": 512, - "hidden_dim": 16384, - "payload_bytes": 16777216, -} -``` - -tag 中只放 TQ 版本支持序列化的小标量/字符串。列表等复杂对象可以编码为稳定字符串或 JSON。`payload_bytes` 用于背压统计。 - -#### TQ 中逻辑上保存的 row - -```text -partition: speco_drafter_features_run-20260818-001 -key: 86a4...ef2 - -fields: - input_ids → int64[512] - loss_mask → float32[512] - position_ids → int64[512] - hidden_states → bf16[512, 16384] - -tag: - status → ready - sequence_no → 12345 - hidden_layout → dflash_aux_plus_last - target_model_fp → sha256:... -``` - -#### Consumer 恢复出的对象 - -rank 0 通过 `kv_list` 同时取得 key 和 tag;每个 rank 用 local keys 调用 `kv_batch_get` 取得 fields,然后组合: - -```python -feature = DraftFeatureSample( - algorithm=tag["algorithm"], - input_ids=densify(fields["input_ids"]).reshape(-1), - loss_mask=densify(fields["loss_mask"]).reshape(-1), - position_ids=densify(fields["position_ids"]).reshape(-1), - hidden_states=densify(fields["hidden_states"]), - metadata={ - "hidden_states_layout": tag["hidden_states_layout"], - "target_model_fingerprint": tag["target_model_fingerprint"], - }, -) - -feature.validate(strict=True) -``` - -这样传给: - -```python -trainer.prepare_training_batch_from_samples([feature, ...]) -``` - -的数据结构,与现有文件 feature store 读出的 `DraftFeatureSample` 一致。也就是说,TQ 替换的是存取介质和调度方式,不改变 DSpark backend 的样本契约。 - -## 5. 新的整体架构 - -```text - 小 metadata - ┌──────────────────────────┐ - │ TransferQueueController │ - │ KV metadata / key / tags │ - │ partition / storage map │ - └────────────┬─────────────┘ - │ -JSONL/token replay │ - │ │ - ▼ │ -Feature Producer │ - ├─ tokenizer/window │ - ├─ asyncio bounded concurrency │ - ├─ vLLM endpoint pool │ - ├─ validate/pack │ - └─ TQ put ─────────────────────┤ - ▼ - TQ Mooncake backend - hidden-state tensors - │ - ┌───────────────────┼───────────────────┐ - ▼ ▼ ▼ - DSpark rank 0 DSpark rank 1 DSpark rank N - TQ get TQ get TQ get - └───────────────────┼───────────────────┘ - ▼ - synchronized optimizer step - │ - ▼ - TQ clear after success -``` - -大 tensor 的路径是: - -```text -vLLM/Producer memory → TQ Mooncake backend → each training rank -``` - -不会走: - -```text -Mooncake → rank 0 → rank 1/2/3 -``` - -rank 0 最多只广播 `kv_list` 得到的 key 字符串列表;tags 只在选 batch 时由 rank 0 使用。 - -### 5.1 standalone 每一步为什么能实现推理和训练异步 - -#### 步骤 A:Producer 独立推进输入 cursor - -Producer 自己维护: - -```python -reader_cursor = 12346 -``` - -它不等待 Trainer 请求样本。只要 TQ ready bytes 没超过背压上限,就继续读取文件并创建 vLLM task。 - -效果是 Producer 的执行进度与 `optimizer_step` 解耦: - -```text -Producer sequence_no: 1200,1201,1202,... -Trainer optimizer_step: 87 -``` - -两者通过 TQ 中的 ready rows 衔接,不互相直接调用。 - -#### 步骤 B:并发 vLLM task 完成顺序可以乱序 - -例如 Producer 同时提交: - -```text -sequence_no 100 → endpoint 0 -sequence_no 101 → endpoint 1 -sequence_no 102 → endpoint 0 -``` - -完成顺序可能是: - -```text -101 → 100 → 102 -``` - -每个 task 完成后独立执行 `kv_put`,所以慢请求不会阻塞已经完成的请求写入。tag 中的 `sequence_no` 保留原数据顺序。 - -#### 步骤 C:`kv_put` 成功是 READY 可见性的边界 - -Producer 在调用前已经得到完整 `DraftFeatureSample`。一次 `kv_put` 写入该 sample 的全部 tensor fields,并在 tag 中标记 `status=ready`。 - -因此 Consumer 的判断规则是: - -```text -kv_list 能列出该 key -且 tag.run_id 匹配 -且 tag.status == ready -→ 可以尝试 kv_batch_get -``` - -Consumer 仍需对取回 fields 做完整性校验;tag 是调度 metadata,不是正确性证明。 - -#### 步骤 D:rank 0 只负责选 key - -rank 0 执行: - -```python -entries = list_ready_keys() -selected = sorted(entries, key=sequence_no)[:global_batch_size] -``` - -这一步处理的数据只是: - -```python -[ - {"key": "k100", "sequence_no": 100, ...}, - {"key": "k101", "sequence_no": 101, ...}, -] -``` - -不包含 `[seq, hidden_dim]` hidden tensor,所以 rank 0 不成为大数据中转瓶颈。 - -#### 步骤 E:广播保证所有 rank 对同一个 global step 达成一致 - -所有 rank 调用同一次: - -```python -dist.broadcast_object_list(holder, src=0) -``` - -广播结束后,每个 rank 看到完全相同的 global key list。然后按确定性区间切分: - -```text -rank 0: keys[0:per_rank] -rank 1: keys[per_rank:2*per_rank] -... -``` - -这样不会出现 rank 0 训练 batch A、rank 1 训练 batch B 数量不同,或者某个 rank 没进入 backward 的情况。 - -#### 步骤 F:各 rank 直接读取 Mooncake 后端 - -每个 rank 执行: - -```python -tq.kv_batch_get(keys=local_keys, partition_id=partition_id) -``` - -TQ 根据 key 找到 storage backend 中的数据。使用 MooncakeStore 时,大 tensor 数据路径是 Mooncake → 本 rank;rank 0 不读取其他 rank 的 local payload。 - -因此: - -```text -控制面:rank 0 → broadcast small keys -数据面:Mooncake → each rank directly -``` - -#### 步骤 G:恢复现有 DraftFeatureSample 契约 - -每个 rank 将 TQ fields 与 tag 合并、densify、校验,得到 `list[DraftFeatureSample]`。从这里开始继续执行项目现有代码: - -```python -batch = trainer.prepare_training_batch_from_samples( - materialized_samples, - step=optimizer_step, -) - -ok = await trainer.training_step_from_batch( - batch, - optimizer_step, -) -``` - -所以 TQ 不进入 DSpark model/backend 内部,训练数学逻辑不变。 - -#### 步骤 H:全 rank 成功以后才能清理 - -每个 rank 的 `ok` 通过现有 `_all_ranks_true()` 聚合: - -```text -rank 0 ok = true -rank 1 ok = true -rank 2 ok = true -rank 3 ok = true -→ global_ok = true -``` - -只有此时 rank 0 执行: - -```python -tq.kv_clear(keys=global_keys, partition_id=partition_id) -``` - -这样保证被删除的数据已经参与完成的 optimizer step。若任何 rank get/OOM/backward 失败,不执行 clear,便于作业失败后的诊断或恢复。 - -#### 步骤 I:异步重叠如何形成 - -时间线上: - -```text -时间 ─────────────────────────────────────────▶ - -Producer: vLLM(batch N+1) ─ put ─ vLLM(batch N+2) ─ put -Trainer: get(batch N) ─ train(batch N) ─ get/train(batch N+1) -``` - -Producer 和 Trainer 是不同进程,互相不调用;TQ ready rows 是缓冲区。因此 vLLM prefill、网络传输和 DSpark GPU 训练可以重叠。背压只在缓冲区达到容量上限时暂停 Producer。 - -## 6. Producer:读取预生成 response 并并行请求 vLLM - -### 6.1 输入处理 - -Producer 从现有 JSONL/token replay 数据源读取: - -```python -sample = { - "row_id": "12345", - "prompt": "...", - "response": "提前生成好的文本", -} -``` - -构造: - -```python -prompt_ids = tokenizer.encode(sample["prompt"]) -response_ids = tokenizer.encode(sample["response"]) -input_ids = prompt_ids + response_ids -``` - -同时产生: - -```python -loss_mask -position_ids -feature_positions -sample_key -``` - -### 6.2 有界并发 - -不能按样本串行请求: - -```python -for sample in samples: - result = request_vllm(sample) -``` - -改成: - -```python -async def run_producer(samples): - semaphore = asyncio.Semaphore(max_inflight_requests) - - async def run_one(sample): - async with semaphore: - result = await vllm_pool.prefill(sample) - feature = validate_and_pack(sample, result) - await tq_transport.put(feature) - - async with asyncio.TaskGroup() as group: - for sample in samples: - group.create_task(run_one(sample)) -``` - -`max_inflight_requests` 是 Producer 同时未完成的请求数量,不是 vLLM batch size。vLLM 服务端仍会对同时到达的请求做自己的 continuous batching。 - -### 6.3 多 endpoint - -多个 endpoint 例如: - -```yaml -vllm_endpoints: - - http://node0:8000/v1 - - http://node1:8000/v1 - - http://node2:8000/v1 -``` - -调度器维护每个 endpoint 的 inflight 数: - -```python -endpoint = min( - endpoints, - key=lambda item: item.inflight, -) -``` - -请求前 `inflight += 1`,在 `finally` 中 `inflight -= 1`。失败只重试对应样本,不阻塞全部 Producer。 - -### 6.4 当前 vLLM 文件桥接 - -当前客户端协议期望: - -```python -response.kv_transfer_params["hidden_states_path"] -``` - -所以第一阶段仍然是: - -```text -vLLM 写临时 safetensors -→ Producer load_file -→ 校验 token_ids/hidden_states -→ TQ put 到 Mooncake backend -→ TQ put 成功后删除临时文件 -``` - -删除必须发生在 TQ put 成功之后: - -```python -path = request_vllm_hidden_file(sample) -try: - feature = load_and_validate(path) - await tq_transport.put(feature) -finally: - if put_succeeded: - Path(path).unlink(missing_ok=True) -``` - -### 6.5 目标版本:vLLM 直接写 TQ/Mooncake - -目标响应可改成: - -```json -{ - "kv_transfer_params": { - "backend": "transfer_queue", - "partition_id": "speco:run-1:dspark_train", - "sample_key": "abc123" - } -} -``` - -服务端顺序必须是: - -```text -prefill -→ 捕获指定层 hidden states -→ TQ/Mooncake put 完成 -→ 返回 HTTP success 和 sample key -``` - -这需要定制 vLLM exporter;当前 `verl-SpeCo-ls` 中没有服务端 writer 实现。 - -## 7. 按 PR #48 扩展 TQ bridge - -不要重新发明一套 transport。将 PR #48 的 bridge 设计移植到本项目并增加 standalone 所需方法: - -```python -class StandaloneTQTransport: - def put_sample(self, key, tensor_dict, tag): ... - def list_ready_keys(self, run_id): ... - def get_samples(self, keys, fields=None): ... - def clear_samples(self, keys): ... - def put_control(self, key, tag): ... - def close(self): ... -``` - -写入延续 PR #48 的真实形式: - -```python -tq.kv_put( - key=key, - partition_id=partition_id, - fields={ - "input_ids": feature.input_ids.cpu(), - "loss_mask": feature.loss_mask.cpu(), - "position_ids": feature.position_ids.cpu(), - "hidden_states": feature.hidden_states.cpu(), - }, - tag={ - "run_id": run_id, - "status": "ready", - "sequence_no": sequence_no, - "payload_bytes": payload_bytes, - }, -) -``` - -批量读取延续 PR #48 的 `kv_batch_get`: - -```python -result = tq.kv_batch_get( - keys=keys, - partition_id=partition_id, - fields=required_fields, # 0.1.7 是否支持该参数需实机确认 -) -``` - -新增发现和清理: - -```python -items = tq.kv_list(partition_id=partition_id) -tq.kv_clear(keys=keys, partition_id=partition_id) -``` - -这里的 `kv_list/kv_clear` 参数名必须根据锁定的 TQ 版本验证。当前参考仓库只实际调用了 `kv_put/kv_batch_get`,没有为这两个接口提供运行证据。 - -读取结果继续复用 PR #48 的两个适配函数: - -```python -value = _extract_value(result, key) -row = _tensordict_to_dict(value) -row["hidden_states"] = _densify_tq_tensor(row["hidden_states"]) -``` - -## 8. DSpark 多 rank 如何消费 - -### 8.1 第一版:rank 0 用 kv_list 发现 READY keys - -PR #48 中 key 由 Ray driver 传给 drafter worker;standalone 没有这条控制路径,所以 rank 0 需要主动列出 key: - -```python -rank = dist.get_rank() -world_size = dist.get_world_size() -global_batch_size = batch_size_per_gpu * world_size - -if rank == 0: - entries = tq_transport.list_ready_keys(run_id=run_id) - entries.sort(key=lambda x: (x.tag["sequence_no"], x.key)) - selected_keys = [x.key for x in entries[:global_batch_size]] -else: - selected_keys = None - -holder = [selected_keys] -dist.broadcast_object_list(holder, src=0) -selected_keys = holder[0] -``` - -`sequence_no` 是输入文件顺序。它让多个 Producer 并发完成顺序不同的情况下,Trainer 仍能确定性地组成 batch。 - -rank 0 广播的是字符串 key 列表,不是 hidden-state tensor。 - -### 8.2 各 rank 切自己的 keys - -例如 global batch keys: - -```text -[s0, s1, s2, s3, s4, s5, s6, s7] -``` - -world size 为 4、每卡 batch size 为 2: - -```text -rank 0 → [s0, s1] -rank 1 → [s2, s3] -rank 2 → [s4, s5] -rank 3 → [s6, s7] -``` - -代码: - -```python -def shard_keys(keys, rank, world_size): - assert len(keys) % world_size == 0 - per_rank = len(keys) // world_size - start = rank * per_rank - end = start + per_rank - return keys[start:end] -``` - -### 8.3 每个 rank 并行 get - -所有进程执行: - -```python -local_keys = shard_keys( - selected_keys, - rank=rank, - world_size=world_size, -) - -local_payloads = tq_transport.get_samples(local_keys) -``` - -数据路径: - -```text -rank 0 ← Mooncake(s0,s1) -rank 1 ← Mooncake(s2,s3) -rank 2 ← Mooncake(s4,s5) -rank 3 ← Mooncake(s6,s7) -``` - -不是 rank 0 get 全部后再 scatter。 - -### 8.4 转成当前训练格式 - -TQ 返回的数据需要先按 PR #48 的规则解包、densify,再转换成现有 `DraftFeatureSample`: - -```python -def tq_row_to_feature(row, tag): - return DraftFeatureSample( - algorithm="DSpark", - input_ids=row["input_ids"], - loss_mask=row["loss_mask"], - position_ids=row["position_ids"], - hidden_states=row["hidden_states"], - metadata={ - "hidden_states_layout": tag["hidden_states_layout"], - **row.get("metadata", {}), - }, - ) -``` - -`hidden_states_layout` 不能只放在无法恢复的临时 Python 对象里;应随 TQ tag 保存。rank 0 从 `kv_list` 取得 key/tag 后,需要把每个 local key 对应的 tag 一起交给本 rank 的转换逻辑。Consumer 必须把 layout 放回 `DraftFeatureSample.metadata`。DSpark L1 所需的 `target_last_hidden_states` 随后由现有 backend 从拼接的 `hidden_states` 中切出,transport 层不计算 loss,也不拆 layout。 - -## 9. 修改当前训练循环 - -在 `run_standalone_draft_training()` 中增加数据源分支: - -```python -feature_store_type = str(feature_store_cfg.get("type", "torch_shard")) - -if feature_store_type == "transfer_queue": - tq_stream = build_transfer_queue_stream( - config=config, - rank=rank, - world_size=world_size, - ) - store = None - loader = None - feature_replayer = None -else: - store = build_feature_store_from_config( - feature_store_cfg, - read_only=True, - ) - loader = DraftFeatureDataLoader(...) -``` - -流式训练循环: - -```python -while successful_steps < max_steps: - global_keys, materialized_samples = tq_stream.next_local_batch() - - batch = trainer.prepare_training_batch_from_samples( - materialized_samples, - step=optimizer_step, - ) - - has_batch = batch is not None - if not _all_ranks_true(has_batch, trainer.runtime_device): - raise RuntimeError("at least one rank failed to fetch its TQ batch") - - ok = await trainer.training_step_from_batch( - batch, - optimizer_step, - ) - - if not _all_ranks_true(ok, trainer.runtime_device): - raise RuntimeError("DSpark step failed on at least one rank") - - dist.barrier() - if rank == 0: - tq_stream.clear_global_batch(global_keys) - dist.barrier() -``` - -TQ 接在数据输入层,不放进 `DSparkTrainerBackend`。Backend 只负责模型、forward、loss、backward 和 optimizer。 - -## 10. READY key、inflight key 和训练提交 - -### 10.1 Ready - -在本方案中 ready 表示: - -> `dspark_train` 所需的所有 fields 已经成功写入 TQ storage,训练侧可以读取。 - -第一版不是通过 PR #48 尚未验证的 Sampler 自动判定,而是 Producer 完整校验 payload 后一次 `kv_put`,并设置 `tag.status=ready`。rank 0 的 `kv_list` 只选择这些 key。 - -### 10.2 Inflight key - -rank 0 选出一个 global batch 后,在本地保存: - -```python -inflight_global_keys = selected_keys -``` - -其他 step 不应再次选中这批 key。因为 standalone 只有一个同步 DSpark job,第一版由 rank 0 进程内 `set` 排除 inflight key 即可: - -```python -ready = [x for x in listed if x.key not in inflight_keys] -``` - -若以后允许多个独立 Trainer job 同时消费同一 partition,进程内 set 就不够,届时必须使用 TQ Controller/Sampler 的原子消费分配能力。 - -### 10.3 Optimizer committed - -optimizer committed 表示所有 DSpark rank 已经完成: - -```text -forward → backward → gradient synchronization → optimizer.step -``` - -它比 `kv_list` 发现 key、甚至 `kv_batch_get` 取出 tensor 都更晚。 - -第一版推荐简单语义: - -```text -TQ 负责 key/tag 和 tensor 传输 -rank 0 负责单 Trainer job 的 batch 选择和 inflight set -训练失败 → 整个作业 fail-fast -训练成功 → kv_clear payload,并从 inflight set 移除 -恢复 → 从最近 checkpoint + 输入 cursor 重新启动 -``` - -这样不需要重新实现 lease/ack 状态机,也没有夸大 PR #48 当前尚未使用的 Sampler 能力。 - -## 11. 为什么训练完一个 step 才清理 - -不能在 `kv_batch_get()` 后立即 clear: - -```text -get 成功 -→ clear -→ forward OOM -→ 数据已不存在,无法重试 -``` - -正确顺序: - -```text -rank 0..N get -→ 所有 rank 确认 batch 有效 -→ training_step_from_batch -→ _all_ranks_true(ok) -→ rank 0 kv_clear global keys -``` - -当前 `_all_ranks_true()` 已经是项目现有的跨 rank 同步工具,可以继续复用。 - -## 12. 背压 - -背压表示 Producer 生成速度高于 Trainer 消费速度时,Producer 必须暂停继续提交,以免 Mooncake 内存无限增长。 - -建议限制: - -```yaml -max_vllm_inflight_requests: 32 -max_pending_put_bytes: 8589934592 -max_tq_ready_samples: 256 -max_tq_ready_bytes: 68719476736 -``` - -Producer 在 tags 中写: - -```python -{"payload_bytes": payload_bytes} -``` - -周期性通过 TQ list/metadata 统计当前 partition 尚未消费的数据量: - -```python -while ready_bytes >= max_tq_ready_bytes: - await asyncio.sleep(backpressure_poll_interval) -``` - -如果锁定的 TQ 版本提供容量/ready 统计接口,应直接使用,避免全量扫描 keys。 - -## 13. Stable ID、幂等和孤儿数据 - -### 13.1 幂等 - -幂等表示同一操作重复执行,最终逻辑结果仍只有一份。 - -Producer 对同一样本重试时必须使用相同 `sample_key`: - -```python -await tq.put(key="abc123", ...) -await tq.put(key="abc123", ...) -``` - -不能每次生成随机 key: - -```text -abc123-retry-1 -abc123-retry-2 -``` - -否则一个输入可能训练多次并持续占用 Mooncake。 - -### 13.2 孤儿数据 - -孤儿数据表示 tensor 已经写进 storage,但由于 Producer 崩溃或 metadata 更新失败,没有进入正常消费路径。 - -使用 TQ 后,metadata 和 storage 由同一套系统管理,可以减少“Mooncake有对象、自研 Coordinator 没记录”的双系统窗口,但仍要配置 TTL/partition cleanup: - -```text -训练正常结束 → clear partition -训练异常退出 → 下次启动检查旧 partition -超过 TTL → 清理未消费数据 -``` - -## 14. EOS 和 drop-last - -EOS 表示 Producer 已经读完输入,并且所有 vLLM/TQ put 都已完成。 - -TQ 中需要一种结束条件,具体使用 Controller API、特殊 metadata 或单独的运行状态记录取决于锁定版本。不能把“当前暂时没有 ready sample”当成 EOS,因为 Producer 可能仍在请求 vLLM。 - -最后不足一个 global batch 时: - -```python -global_batch_size = batch_size_per_gpu * world_size -``` - -第一版使用 `drop_last=true`,避免某些 rank 有数据、某些 rank 没数据导致分布式训练不同步。 - -结束条件: - -```text -producer_done == true -and ready_samples < global_batch_size -and inflight_requests == 0 -and pending_puts == 0 -``` - -## 15. 双缓冲预取 - -训练 batch N 时,CPU 后台线程预取 batch N+1: - -```python -next_future = executor.submit(tq_stream.next_local_batch) - -current_batch = first_batch -while current_batch is not None: - next_batch = next_future.result() - next_future = executor.submit(tq_stream.next_local_batch) - - train(current_batch) - current_batch = next_batch -``` - -实际顺序应调整为避免等待 future 后才训练。推荐: - -```python -current = tq_stream.next_local_batch() - -while current is not None: - future = executor.submit(tq_stream.next_local_batch) - train_and_clear(current) - current = future.result() -``` - -第一版只预取一个 global batch,避免 Trainer 崩溃时大量数据已被采样但未训练。 - -如果 Mooncake/TQ Python get 是阻塞函数,使用专用 `ThreadPoolExecutor`,不要阻塞 Producer 的 asyncio event loop。 - -## 16. 建议代码结构 - -```text -verl_speco/ - trainer/ - tq_transport.py # TQ client、put/get/meta/clear 封装 - tq_feature_stream.py # kv_list、global keys、rank shard、decode/prefetch - feature_producer.py # JSONL → 并发 vLLM → TQ - draft_training_loop.py # 增加 transfer_queue 数据源分支 - target_feature_replay.py # 复用 vLLM payload 校验/转换逻辑 -``` - -不要新增: - -```text -coordinator.py -coordinator_client.py -``` - -建议抽象: - -```python -class StreamingFeatureSource(Protocol): - def next_local_batch(self) -> tuple[list[str], list[DraftFeatureSample]]: ... - def clear_global_batch(self, keys: list[str]) -> None: ... - def close(self) -> None: ... -``` - -这样训练循环不依赖 TQ 的具体类型。 - -## 17. 配置草案 - -```yaml -actor_rollout_ref: - rollout: - drafter: - training: - mode: offline - backend: dspark - batch_size_per_gpu: 2 - max_steps: 1000 - - feature_store: - type: transfer_queue - partition_id: speco_drafter_features_${run_id} - drop_last: true - prefetch_steps: 1 - - transfer_queue: - # 与 PR #48 的配置层级和 init 方式保持一致。 - enable: true - package_version: 0.1.8 # 最终以实测版本为准 - backend: - storage_backend: MooncakeStore - MooncakeStore: - auto_init: false - metadata_server: localhost:50123 - master_server_address: localhost:50124 - local_hostname: localhost - protocol: tcp - global_segment_size: 4294967296 - local_buffer_size: 1073741824 - device_name: "" - - required_fields: - - input_ids - - loss_mask - - position_ids - - hidden_states - - producer: - input_path: /path/to/generated_responses.jsonl - vllm_endpoints: - - http://node0:8000/v1 - - http://node1:8000/v1 - max_inflight_requests: 32 - max_pending_put_bytes: 8589934592 - max_ready_samples: 256 - max_ready_bytes: 68719476736 -``` - -当前 examples 中的: - -```bash -transfer_queue.enable=False -``` - -属于 verl RL 主入口配置,当前 standalone `draft_train_launcher` 不读取它。不能只改成 `True`;必须实现上述 `feature_store.type=transfer_queue` 分支。 - -## 18. 启动顺序 - -逻辑顺序: - -```text -1. 启动 Mooncake metadata/master 服务; -2. 启动 standalone TQ owner 进程,调用一次 `tq.init(tq_config)` 创建/连接 TQ Controller 和 MooncakeStore backend; -3. 启动一个或多个定制 vLLM server -4. 启动 Feature Producer -5. Producer 和所有 Trainer rank 调用无参 `tq.init()`,连接 owner 创建的同一套 TQ; -6. 启动 verl_speco.draft_train_launcher -7. torchrun 启动所有 DSpark rank -8. 各 rank 连接 TQ -9. rank 0 用 `kv_list` 选择 global keys,各 rank 并行 `kv_batch_get`/train; -10. 输入耗尽后 Producer 发布 done 状态 -11. Trainer drain 完整 global batches 后退出 -12. 清理 partition,停止 Producer、vLLM、TQ、Mooncake -``` - -PR #48 的 `init_transfer_queue()` 在 Ray `SpecoTaskRunner` 中调用 `tq.init(config)`,worker 的无参 `tq.init()`依靠同一个 Ray 集群发现 named Controller;这部分不能原样复制到 standalone torchrun。 - -本项目要求一开始不使用 Ray,因此 Phase 0 必须先证明锁定的 TQ 版本支持独立 owner/controller 进程,以及 Producer/torchrun rank 如何获得连接信息。若 TQ 0.1.7/0.1.8 实际只能通过 Ray named actor 完成发现,那么有两个选择: - -1. 接受仅用 Ray 承载 TQ 控制面的最小方案; -2. 给 TQ 增加或使用其已有的独立 ZMQ/server-info bootstrap。 - -在这项验证完成前,文档不能声称 PR #48 已经提供“无 Ray TQ standalone 启动”。 - -## 19. 故障处理 - -### vLLM 请求失败 - -- 对单个 sample 按稳定 key 重试; -- 指数退避; -- 超过次数记录失败,并根据配置 fail-fast 或跳过; -- 不写不完整 TQ fields。 - -### vLLM 文件读取成功,但 TQ put 失败 - -- 暂时保留临时文件; -- 重试 TQ put; -- put 成功后再删除; -- 不把样本视为 ready。 - -### 某个训练 rank get 失败 - -- 该 rank 报告 `local_ok=false`; -- `_all_ranks_true()` 使全部 rank 得到一致失败结果; -- 第一版整个训练 fail-fast; -- 不 clear global batch。 - -### OOM/optimizer step 失败 - -- 不 clear; -- 所有 rank 一致退出; -- 从最近训练 checkpoint 恢复; -- 根据 TQ 消费提交语义决定是否重放当前 batch。 - -### clear 失败 - -- optimizer 已成功,不能再次训练这批; -- 将 batch keys 写入本地小型 `gc_pending` 日志; -- 后台重试 clear; -- checkpoint 保存最近 committed sample IDs,避免恢复时重复消费。 - -## 20. 观测指标 - -Producer: - -```text -producer/vllm_inflight -producer/vllm_requests_per_sec -producer/vllm_prefill_tokens_per_sec -producer/vllm_p50_latency -producer/vllm_p95_latency -producer/tq_put_bytes_per_sec -producer/tq_put_failures -producer/pending_put_bytes -``` - -TQ/Mooncake: - -```text -tq/ready_samples -tq/ready_bytes -tq/consumed_samples -tq/storage_bytes -tq/clear_failures -mooncake/put_bandwidth -mooncake/get_bandwidth -``` - -Trainer: - -```text -trainer/tq_wait_seconds -trainer/tq_get_seconds -trainer/tq_get_bytes_per_sec -trainer/decode_seconds -trainer/h2d_seconds -trainer/step_seconds -trainer/data_stall_ratio -trainer/successful_steps -``` - -## 21. 实施阶段 - -### Phase 0:锁定依赖和契约 - -- 从 PR #48 的 `TransferQueue==0.1.7` 起验证,同时对比当前 verl 文档使用的 0.1.8; -- 实测 `tq.init(config)`、worker `tq.init()`、`kv_put`、`kv_batch_get`、`kv_list`、`kv_clear`; -- 验证该版本是否支持无 Ray owner/controller 部署以及连接信息传递; -- 先用 PR #48 已配置的 `SimpleStorage` 做最小闭环; -- 再把 backend 换成 `MooncakeStore`,验证 tcp,最后再验证 rdma; -- 固定 fields、tags、partition 和 `run_id/sequence_no/status`; -- 写 fake TQ 单元测试。 - -### Phase 1:文件桥接 + TQ KV 模式 - -- 新增独立 Producer; -- 32 个有界并发 vLLM 请求; -- 读取 vLLM 临时 safetensors; -- TQ put 成功后删除文件; -- standalone trainer 由 rank 0 `kv_list` 并广播 global keys; -- 各 rank 并行 `kv_batch_get`; -- 复用 PR #48 的返回值解包、NestedTensor densify 和 per-step cache; -- optimizer 成功后 `kv_clear`。 - -验收:连续训练 1000 step,临时文件数量、TQ ready bytes 和 Mooncake占用均保持有界。 - -### Phase 2:双缓冲与多 endpoint - -- 增加多 endpoint 最少 inflight 调度; -- 增加一个 global batch 预取; -- 动态背压; -- 注入单 rank get 失败,确认所有 rank 一致退出而非死锁。 - -### Phase 3:vLLM 直接写 TQ/Mooncake - -- 修改外部定制 vLLM exporter; -- 去掉 `hidden_states_path` 临时文件; -- HTTP 响应返回 partition/sample key; -- 验证 HTTP 重试的幂等性。 - -### Phase 4:可选升级到 TQ StreamingDataLoader - -- 在当前保守方案稳定后再引入 RankAwareSampler; -- 让每个 rank 自动取得 local micro-batch; -- 去掉 rank 0 手工 key-list 广播; -- 验证与 torchrun/DSpark 的 global step 对齐。 - -## 22. 最终推荐 - -针对当前 `verl-SpeCo-ls`,推荐的第一版不是 AngelSpec 的完整架构,也不是单独写 Coordinator,而是: - -```text -当前预生成 response 文件 -→ 独立 asyncio Producer -→ 并行访问多个 vLLM endpoint -→ 读取并校验临时 hidden-state 文件 -→ TransferQueue put -→ Mooncake storage backend -→ rank 0 kv_list 获取 READY global keys -→ broadcast key list -→ 各 DSpark rank 并行 kv_batch_get -→ 现有 prepare_training_batch_from_samples() -→ 现有 training_step_from_batch() -→ 全 rank 成功 -→ TQ clear -``` - -这套方案保留当前 standalone DSpark 训练主体,只替换 `DraftFeatureDataLoader + TargetFeatureReplayer.materialize()` 所在的数据输入路径。它真正建立在 PR #48 已实现的 KV transport 之上,而不是假设 PR #48 已经实现了 standalone Sampler/StreamingDataLoader。 - -## 23. 参考 - -- verl TransferQueue: -- TransferQueue: -- Mooncake Store: -- 本地只读参考:`../verl-SpeCo/verl_speco/integration/transferqueue_bridge.py`(PR #48) -- 本地只读参考:`../verl-SpeCo/docs/transferqueue_integration_plan.md`(PR #48) diff --git a/docs/standalone_tq_foundation_implementation.md b/docs/standalone_tq_foundation_implementation.md deleted file mode 100644 index 5bb7b48e..00000000 --- a/docs/standalone_tq_foundation_implementation.md +++ /dev/null @@ -1,1037 +0,0 @@ -# Standalone TQ 公共基础层实现说明 - -## 1. 文档范围和已验证结论 - -本文解释当前仓库已经实现并测试通过的 TQ 公共基础层: - -```text -verl_speco/transport/drafter_sample_protocol.py -verl_speco/integration/transferqueue_bridge.py -verl_speco/config/speco_base.yaml -verl_speco/tq_owner.py -examples/run_dspark_tq_owner.sh -examples/tq_connection_smoke.py -tests/unit/test_drafter_sample_protocol.py -tests/unit/test_transferqueue_bridge.py -pyproject.toml -``` - -当前已经实现: - -1. Producer 和 Consumer 共用的样本 key、tag、fields 和 metadata 协议; -2. 普通进程连接 Ray 集群; -3. TQ Owner 创建 named `TransferQueueController`; -4. 独立 Client 发现并连接同一个 Controller; -5. 单样本 put、元数据 list、批量 get 和批量 clear; -6. Owner 与 Client 不同的关闭边界; -7. 独立 Owner 入口和共享 Hydra 配置; -8. mock 单元测试和真实双进程 TQ 0.1.7 smoke test。 - -当前还没有实现: - -1. Producer 读取 JSONL、并发访问多个 vLLM endpoint 的完整 pipeline; -2. `feature_store.type=tq` 工厂分支; -3. `TQFeatureStore` 和 `TQFeatureDataLoader`; -4. rank 0 选择 global keys、各 rank 读取 local keys; -5. TQ batch 接入 DSpark optimizer step; -6. optimizer step 成功后的 rank 0 clear。 - -因此当前代码已经证明“两个独立进程能通过 Ray 连接同一个 TQ,并按共享协议批量传输多个样本”,但尚未接到正式 Producer 和 DSpark Consumer 主循环。 - -## 2. 运行时角色和术语 - -### 2.1 Ray head - -Ray head 是 Ray 集群的控制节点。TQ 0.1.7 使用 Ray 管理 named Controller 和 storage actors。 - -Ray 不传输 standalone hidden-state payload。当前代码没有对这些样本调用: - -```python -ray.put(hidden_states) -``` - -### 2.2 TQ Owner - -TQ Owner 是普通 Python OS 进程,入口为: - -```text -python -m verl_speco.tq_owner -``` - -它不是 Ray actor。它先调用 `ray.init(address=...)` 加入 Ray,再调用带完整配置的 `tq.init(config)` 创建 TQ Controller 和 storage。 - -Owner 是唯一允许调用全局 `tq.close()` 的进程。 - -### 2.3 Named TransferQueueController - -TQ 0.1.7 内部创建: - -```python -TransferQueueController.options( - name="TransferQueueController" -).remote(...) -``` - -`name` 将 Controller actor 注册到 Ray actor registry。其他进程加入同一 Ray address 和 namespace 后,通过: - -```python -ray.get_actor("TransferQueueController") -``` - -取得 actor handle,再读取 TQ backend 配置。 - -Controller 保存控制信息和数据位置;使用 MooncakeStore 时,大 tensor 本身保存在 MooncakeStore。 - -### 2.4 TQ Client - -Owner、Producer、每个 Consumer rank 都在各自 OS 进程内拥有独立 TQ Client。 - -普通 Client 通过无参: - -```python -tq.init() -``` - -发现 named Controller并初始化本进程的 storage manager。Client 不是 Ray actor,Producer 和 torchrun rank 也不需要改成 Ray actor。 - -### 2.5 Partition、key、tag 和 fields - -当前固定 partition: - -```text -speco_drafter_features -``` - -TQ 中一条记录逻辑上是: - -```text -partition_id -└── key - ├── tag:轻量 dict,由 kv_list 发现 - └── fields:Tensor payload,由 kv_batch_get 读取 -``` - -## 3. 共享配置如何工作 - -共享配置定义在 `verl_speco/config/speco_base.yaml`: - -```yaml -transfer_queue: - enable: false - package_version: "0.1.7" - ray: - address: null - namespace: speco-drafter - partition_id: speco_drafter_features - run_id: null - schema_version: 1 - connect_timeout_seconds: 120 - poll_interval_seconds: 0.5 - drop_last: true - controller: - polling_mode: true - backend: - storage_backend: SimpleStorage - SimpleStorage: - total_storage_size: 100000 - num_data_storage_units: 8 - MooncakeStore: - auto_init: false - metadata_server: localhost:50050 - master_server_address: localhost:50051 - local_hostname: "" - protocol: tcp - global_segment_size: 4294967296 - local_buffer_size: 1073741824 - device_name: "" -``` - -### 3.1 Ray 连接字段 - -```yaml -ray: - address: 10.0.0.1:6379 - namespace: speco-drafter -``` - -它们决定当前进程连接哪个 Ray 集群,以及在哪个 namespace 查找 `TransferQueueController`。Owner、Producer 和所有 Consumer ranks 必须使用相同值。 - -### 3.2 SPECO 协议字段 - -```yaml -partition_id: speco_drafter_features -run_id: dspark-20260819-a -schema_version: 1 -``` - -这些字段用于路由和校验,不属于 TQ 原生配置。`run_id` 用于隔离不同 pipeline 的记录。 - -### 3.3 TQ 原生字段 - -```yaml -controller: ... -backend: ... -``` - -只有这些字段应传给 `tq.init(full_config)`。bridge 的 `_native_tq_config()` 会删除: - -```text -enable -package_version -ray -partition_id -run_id -schema_version -connect_timeout_seconds -poll_interval_seconds -drop_last -``` - -对象变化为: - -```text -完整 SPECO transfer_queue dict -→ _native_tq_config() -→ controller/backend等TQ字段 -→ OmegaConf DictConfig -→ tq.init() -``` - -## 4. Bridge 的进程内状态 - -`verl_speco/integration/transferqueue_bridge.py` 在每个 OS 进程内分别维护: - -```python -_state = { - "enabled": False, - "configured": False, - "initialized": False, - "config": None, - "owner": False, - "ray_initialized_here": False, - "ray_address": None, - "ray_namespace": None, -} -``` - -该 dict 不跨进程共享。Owner、Producer、每个 rank 分别拥有自己的 `_state`。 - -| 字段 | 含义 | -|---|---| -| `enabled` | 当前进程配置是否开启 TQ | -| `configured` | 是否调用过 `configure_transfer_queue()` | -| `initialized` | 当前进程是否执行过 `tq.init()` | -| `config` | 当前进程保存的普通 dict 配置 | -| `owner` | 当前进程是否创建了全局 Controller/Storage | -| `ray_initialized_here` | bridge 是否负责调用了本进程的 `ray.init()` | -| `ray_address/namespace` | 本进程的 Ray 连接信息 | - -`_state_lock` 只保护同一进程内多个线程同时初始化,不是分布式锁。 - -## 5. Owner 的完整启动数据流 - -Owner 入口是 `verl_speco/tq_owner.py`,直接加载 Hydra 主配置 -`verl_speco/config/speco_base.yaml`。`run_owner()` 会复制其中的 -`transfer_queue` 子配置,并仅在 Owner 进程的副本中强制设置 -`enable=true`;普通训练任务看到的共享默认值仍是 `false`。 - -### 阶段 1:读取配置 - -执行者:Owner OS 进程。 - -入口: - -```python -run_owner(config) -``` - -取得: - -```python -training_cfg = config.actor_rollout_ref.rollout.drafter.training -tq_cfg = training_cfg.transfer_queue -``` - -然后调用: - -```python -configure_transfer_queue(training_cfg) -``` - -该函数只把 OmegaConf 转成普通 dict并更新当前进程 `_state`,不会连接 Ray,也不会创建 TQ。 - -### 阶段 2:连接 Ray - -Owner 调用: - -```python -connect_ray_cluster(ray_address, namespace) -``` - -内部执行: - -```python -if not ray.is_initialized(): - ray.init(address=ray_address, namespace=namespace) -``` - -边界类型是 Ray control-plane connection。此时还没有传输训练 tensor。 - -### 阶段 3:创建 Controller 和 Storage - -Owner 调用: - -```python -start_transfer_queue_owner(tq_cfg) -``` - -执行顺序: - -1. `_extract_tq_config()` 得到普通 dict; -2. 检查 `enable=true`; -3. 检查 `TransferQueue` 包可用; -4. 防止本进程重复初始化; -5. `_native_tq_config()` 删除 SPECO 字段; -6. `_as_tq_config()` 转 OmegaConf; -7. `tq.init(native_config)` 创建 Controller、Storage 和 Owner Client; -8. 设置 `_state.owner=True`、`initialized=True`。 - -Ray 中形成: - -```text -Ray cluster / namespace -├── named actor: TransferQueueController -└── storage backend - ├── SimpleStorage actors - └── 或 MooncakeStore connection/process -``` - -### 阶段 4:发布 owner-ready - -调用: - -```python -publish_owner_ready(run_id, schema_version) -``` - -生成: - -```python -key = "control:v1::owner-ready" -fields = {"marker": torch.tensor([1], dtype=torch.uint8)} -tag = { - "record_type": "control", - "status": "owner_ready", - "schema_version": 1, - "run_id": run_id, -} -``` - -这是一条控制记录,不进入训练 batch。 - -### 阶段 5:常驻和关闭 - -Owner 安装 `SIGINT/SIGTERM` handler,并等待: - -```python -stop_event.wait() -``` - -收到信号后调用: - -```python -close_transfer_queue_owner() -``` - -执行: - -```text -验证owner身份 -→ tq.close() -→ 清理Controller/Storage -→ ray.shutdown() -``` - -Owner 必须在 Producer 和 Consumer 退出后才能关闭。 - -## 6. 普通 Client 如何连接同一个 TQ - -Producer 和每个 Consumer rank 后续使用相同顺序: - -```python -configure_transfer_queue(training_cfg) -connect_ray_cluster(ray_address, namespace) -connect_transfer_queue_client() -``` - -`connect_transfer_queue_client()` 最终调用无参: - -```python -tq.init() -``` - -TQ 0.1.7 内部通过: - -```python -ray.get_actor("TransferQueueController") -``` - -找到 Owner 创建的 Controller,读取 backend 配置,然后创建当前进程的 TQ Client。 - -对象和边界变化: - -```text -actor名称字符串 -→ Ray actor registry -→ Controller actor handle -→ Controller.get_config.remote() -→ TQ DictConfig -→ 当前进程TransferQueueClient -→ 同一个SimpleStorage/MooncakeStore -``` - -## 7. 一条具体样本的初始对象 - -真实 smoke test使用: - -```python -sample = DraftFeatureSample( - algorithm="DSPARK", - input_ids=torch.tensor([1, 2, 3]), # int64[3], CPU - loss_mask=torch.tensor([0.0, 1.0, 1.0]), # float32[3], CPU - position_ids=torch.tensor([0, 1, 2]), # int64[3], CPU - hidden_states=torch.arange( - 12, dtype=torch.float32 - ).reshape(3, 4), # float32[3,4], CPU -) -``` - -同时构造: - -```python -meta = SampleMetadata( - schema_version=1, - run_id="codex-batch-smoke", - sample_id="smoke-0000", - sequence_no=0, - algorithm="DSPARK", - target_model_id="smoke-target", - target_model_revision="smoke-revision", - tokenizer_fingerprint="smoke-tokenizer", - target_layer_ids=[0], - hidden_states_layout="dflash_aux", - hidden_dtype="float32", - hidden_shape=[3, 4], - feature_length=3, - full_sequence_length=3, - feature_start=0, - feature_end=3, - use_logits=False, -) -``` - -`DraftFeatureSample` 和 `SampleMetadata` 都是进程内 Python 对象,不直接经过 TQ。 - -## 8. Key 的生成和两个同名函数 - -共享协议调用: - -```python -make_sample_key(meta) -``` - -输出: - -```text -drafter:v1:codex-batch-smoke:000000000000:smoke-0000 -``` - -字段顺序: - -```text -drafter / schema version / run_id / 12位sequence_no / sample_id -``` - -`sequence_no` 是输入文件 record 序号,不是训练 batch 编号。 - -bridge 为兼容 PR #48 还保留另一个: - -```python -transferqueue_bridge.make_sample_key( - global_step, - replica_rank, - request_id, -) -``` - -它生成: - -```text -speco::: -``` - -standalone Producer 必须从 `verl_speco.transport.drafter_sample_protocol` import `make_sample_key`,不能使用 bridge 中的 PR #48 旧函数。 - -## 9. Tag 如何生成 - -```python -tag = make_ready_tag(meta) -``` - -输出: - -```python -{ - "record_type": "sample", - "status": "ready", - "schema_version": 1, - "run_id": "codex-batch-smoke", - "sequence_no": 0, - "sample_id": "smoke-0000", - "algorithm": "DSPARK", -} -``` - -tag 只包含发现、过滤和排序需要的标量/字符串。`list_samples()` 只返回 key/tag,不读取 hidden states。 - -## 10. `encode_sample()` 如何生成 fields - -调用: - -```python -fields = encode_sample(sample, meta) -``` - -### 10.1 校验 - -执行: - -```text -SampleMetadata.validate() -DraftFeatureSample.validate(strict=True) -``` - -随后检查: - -1. hidden states 是一个 dense tensor; -2. ids/mask/position 长度等于 `feature_length`; -3. hidden 第一维等于 `feature_length`; -4. hidden shape 等于 metadata; -5. hidden dtype 等于 metadata; -6. feature window 长度正确。 - -### 10.2 Tensor 规范化 - -```text -input_ids → CPU contiguous int64[L] -loss_mask → CPU contiguous float32[L] -position_ids → CPU contiguous int64[L] -hidden_states → CPU contiguous,保持模型dtype -``` - -没有 `position_ids` 时生成 `torch.arange(L, dtype=int64)`。 - -### 10.3 Metadata JSON 编码 - -```text -SampleMetadata dataclass -→ dict -→ JSON UTF-8 bytes -→ torch.uint8[M] -``` - -实现等价于: - -```python -raw = json.dumps(metadata).encode("utf-8") -metadata_json = torch.tensor(list(raw), dtype=torch.uint8) -``` - -### 10.4 最终 fields - -```python -fields = { - "input_ids": int64[3], - "loss_mask": float32[3], - "position_ids": int64[3], - "hidden_states": float32[3,4], - "metadata_json": uint8[M], -} -``` - -如果存在,协议也保留 `last_hidden_states`、`target` 和 `target_logprobs`。 - -## 11. Bridge 如何写入 TQ - -调用: - -```python -put_sample(key, fields, tag=tag) -``` - -bridge 执行: - -1. 检查 TQ 已启用; -2. 丢弃 fields 中非 tensor 值; -3. 确保本进程已经 `tq.init()`; -4. 取得配置中的 partition; -5. 调用: - -```python -tq.kv_put( - key=key, - partition_id="speco_drafter_features", - fields=fields, - tag=tag, -) -``` - -使用 MooncakeStore 时,大 tensor 路径是: - -```text -Producer CPU tensor -→ Producer TQ Client -→ MooncakeStore -``` - -不是 Ray `ObjectRef`。 - -## 12. Consumer 如何发现 key - -调用: - -```python -records = list_samples() -``` - -内部调用: - -```python -tq.kv_list(partition_id="speco_drafter_features") -``` - -标准化返回类型: - -```python -dict[str, dict[str, Any]] -``` - -示例: - -```python -{ - "drafter:v1:...:smoke-0000": { - "record_type": "sample", - "status": "ready", - "run_id": "codex-batch-smoke", - "sequence_no": 0, - ... - } -} -``` - -bridge 兼容 `key → tag` 和 `partition → key → tag` 两种 wrapper。这一步不读取 fields。 - -## 13. Consumer 如何批量取样本 - -输入: - -```python -keys = [key0, key1] -``` - -调用: - -```python -records = get_samples(keys) -``` - -bridge 只调用一次: - -```python -result = tq.kv_batch_get( - keys=keys, - partition_id="speco_drafter_features", -) -``` - -TQ 0.1.7 返回带 batch 维的 TensorDict。bridge 检查 `result.batch_size`,再执行: - -```python -rows = [result[index] for index in range(len(keys))] -``` - -每行转成普通 dict,最终返回: - -```python -[ - (key0, fields0), - (key1, fields1), -] -``` - -返回顺序与输入 keys 一致。重复 key 会提前报错。 - -## 14. `decode_sample()` 如何恢复训练对象 - -调用: - -```python -sample = decode_sample( - key, - tag, - fields, - expected_config, -) -``` - -### 14.1 Metadata 解码 - -```text -metadata_json uint8[M] -→ bytes -→ UTF-8 -→ json.loads -→ dict -→ SampleMetadata.from_dict -``` - -### 14.2 身份一致性 - -代码根据 metadata 重新生成 key,要求输入 key 完全相等;然后逐项校验 tag 的: - -```text -record_type/status/schema_version/run_id/sequence_no/sample_id/algorithm -``` - -所以 key、tag 和 payload metadata 不能来自不同样本。 - -### 14.3 Consumer 合同 - -Consumer 提供: - -```python -ExpectedFeatureConfig( - run_id="codex-batch-smoke", - schema_version=1, - algorithm="DSPARK", - target_model_id="smoke-target", - target_model_revision=None, - tokenizer_fingerprint=None, - target_layer_ids=None, - hidden_states_layout="dflash_aux", - hidden_dtype="float32", -) -``` - -值为 `None` 的字段不检查;其他字段必须完全一致。正式 Consumer 应填写 target checkpoint、tokenizer、layers、layout 和 dtype,避免使用错误 target 特征。 - -### 14.4 输出 - -完成 tensor 类型、长度、shape、dtype 校验后,构造: - -```python -DraftFeatureSample.from_dict(payload, strict=True) -``` - -输出可以交给现有 `trainer.prepare_training_batch_from_samples()`。当前 smoke test验证到这里,正式 `TQFeatureDataLoader` 尚未实现。 - -## 15. EOS 控制记录 - -调用: - -```python -key, fields, tag = make_eos_record(run_id, total_samples) -``` - -输出: - -```python -key = "control:v1::eos" -fields = {"marker": torch.tensor([1], dtype=torch.uint8)} -tag = { - "record_type": "control", - "status": "eos", - "schema_version": 1, - "run_id": run_id, - "total_samples": total_samples, -} -``` - -EOS 表示不会再发布新样本。Consumer 应先 drain ready samples 再退出。协议函数已实现,正式 Producer/Consumer 尚未调用。 - -## 16. Clear 和数据生命周期 - -bridge 提供: - -```python -clear_samples(keys) -``` - -内部调用: - -```python -tq.kv_clear(keys=keys, partition_id="speco_drafter_features") -``` - -基础层只执行删除,不决定删除时机。正式 Consumer 必须遵守: - -```text -rank 0选择global keys -→ 各rank读取local keys -→ 所有rank完成同一optimizer step -→ 汇总global success -→ rank 0 clear global keys -``` - -不能在 `get_samples()` 后立即 clear,因为 get 成功不代表训练 step 成功。 - -## 17. Client close 和 Owner close - -### 17.1 Client close - -Producer/rank 调用: - -```python -close_transfer_queue_client() -``` - -执行: - -```text -tq.get_client() -→ 当前进程client.close() -→ 如果bridge负责ray.init,则ray.shutdown() -``` - -它不调用全局 `tq.close()`,不会 kill Controller。调用后本进程不能继续使用 TQ。 - -### 17.2 Owner close - -Owner 调用: - -```python -close_transfer_queue_owner() -``` - -执行: - -```text -tq.close() -→ Controller/Storage全局清理 -→ ray.shutdown() -``` - -Owner 若误用 Client close,bridge 会抛 `RuntimeError`。 - -## 18. PR #48 兼容边界 - -bridge 继续保留: - -```python -init_transfer_queue(config) -get_sample(key) -close_transfer_queue() -``` - -PR #48 的 `SpecoTaskRunner` 已是 Ray actor,所以不调用 standalone 的 `connect_ray_cluster()`。旧 worker 继续使用单 key `get_sample()`。 - -standalone 后续使用新增的: - -```python -list_samples() -get_samples(keys) -clear_samples(keys) -``` - -因此没有修改 PR #48 现有调用点的函数签名。 - -## 19. 依赖和命令入口 - -`pyproject.toml` 新增: - -```toml -[project.optional-dependencies] -transfer-queue = ["TransferQueue==0.1.7"] -``` - -安装: - -```bash -pip install -e ".[transfer-queue]" -``` - -TQ 0.1.7 会安装 Ray,并要求 `numpy<2.0.0`。若现有包要求 NumPy 2,需要隔离环境或重新解决依赖。 - -Owner 命令: - -```text -verl-speco-tq-owner -``` - -也可以使用 `examples/run_dspark_tq_owner.sh`。 - -## 20. 单元测试 - -协议测试 `tests/unit/test_drafter_sample_protocol.py` 覆盖: - -1. encode/decode round trip; -2. key 格式; -3. tag 身份冲突; -4. Consumer contract 冲突; -5. hidden shape 冲突; -6. EOS 格式。 - -bridge 测试 `tests/unit/test_transferqueue_bridge.py` 覆盖: - -1. Ray address/namespace 参数; -2. Owner 只向 TQ 传原生配置; -3. Client 无参 `tq.init()`; -4. put/list/get-many/clear; -5. batch 返回顺序; -6. Client close 不调用全局 close; -7. Owner 不能误用 Client close。 - -运行: - -```bash -python -m pytest \ - tests/unit/test_drafter_sample_protocol.py \ - tests/unit/test_transferqueue_bridge.py \ - -q -``` - -## 21. 真实双进程 smoke test - -程序: - -```text -examples/tq_connection_smoke.py -``` - -它使用真实 `TransferQueue==0.1.7`、Ray、SimpleStorage、两个独立 Python进程和两个 batch samples。 - -Owner 路径: - -```text -连接Ray -→ tq.init(full config) -→ 写sample 0和sample 1 -→ 等待client-done -→ clear done marker -→ 全局关闭 -``` - -Client 路径: - -```text -连接同一个Ray -→ tq.init() -→ kv_list发现两个key -→ 一次kv_batch_get([k0,k1]) -→ 拆成两个fields dict -→ 分别decode_sample -→ clear两个sample keys -→ 写client-done -→ 只关闭本地client -``` - -已验证输出: - -```text -OWNER_READY keys=[k0, k1] -CLIENT_OK samples=2 shape=(3, 4) -CLIENT_CLOSED_LOCAL_ONLY -OWNER_OBSERVED_CLIENT_DONE -OWNER_CLOSED -``` - -这证明: - -1. 两个普通进程能连接同一个 TQ; -2. named Controller 发现有效; -3. 0.1.7 的 `kv_list/kv_batch_get/kv_clear` 参数有效; -4. TensorDict batch 能按 key 顺序拆开; -5. 共享协议能恢复 `DraftFeatureSample`; -6. Client close 不会杀掉 Owner; -7. Owner 能最终统一关闭。 - -## 22. 当前完整路径总结 - -```text -Owner -→ ray.init(address, namespace) -→ tq.init(native config) -→ named TransferQueueController - -普通Client -→ ray.init(same address, same namespace) -→ tq.init() -→ 找到同一个Controller - -DraftFeatureSample + SampleMetadata -→ make_sample_key -→ make_ready_tag -→ encode_sample -→ fields + metadata_json tensor -→ bridge.put_sample -→ tq.kv_put -→ SimpleStorage/MooncakeStore - -Consumer/测试Client -→ bridge.list_samples -→ key + tag -→ bridge.get_samples(keys) -→ tq.kv_batch_get -→ TensorDict batch -→ 每个key对应一个fields dict -→ decode_sample -→ DraftFeatureSample - -正式训练成功后(待实现) -→ bridge.clear_samples(global_keys) - -Client退出 -→ close_transfer_queue_client - -所有业务进程退出 -→ Owner close_transfer_queue_owner -→ tq.close -→ ray.shutdown -``` - -## 23. 下一阶段接入约束 - -后续代码不能重新定义协议或直接访问 TQ 私有对象。 - -Producer 应复用: - -```text -SampleMetadata -make_sample_key -make_ready_tag -encode_sample -bridge.put_sample -make_eos_record -``` - -Consumer 应复用: - -```text -bridge.list_samples -bridge.get_samples -decode_sample -bridge.clear_samples -``` - -下一阶段需要新增: - -```text -verl_speco/trainer/tq_feature_store.py -verl_speco/trainer/tq_sample_source.py -feature_store.py 的 type=tq 分支 -draft_training_loop.py 的流式训练分支 -Producer入口、输入读取和并发vLLM文件 -``` - -这些文件应建立在本文已经实现和真实验证过的连接、协议、KV 和关闭接口之上。 diff --git a/docs/standalone_vllm_tq_dspark_training_plan.md b/docs/standalone_vllm_tq_dspark_training_plan.md deleted file mode 100644 index 6644c6de..00000000 --- a/docs/standalone_vllm_tq_dspark_training_plan.md +++ /dev/null @@ -1,1131 +0,0 @@ -# 独立 vLLM Producer + TQ + DSpark Consumer 第一版方案 - -## 1. 第一版要实现什么 - -只实现下面这条主链路: - -```text -包含 prompt + 预生成 response 的输入文件 -→ Producer 并发请求 vLLM prefill -→ Producer 将每条训练样本写入 TQ -→ Consumer 从同一个 TQ 取样本 -→ 独立 torchrun/FSDP DSpark 训练 -→ 一个 optimizer step 成功后删除该 step 的 TQ 样本 -``` - -第一版允许 **TQ 基础设施内部使用 Ray**,因为实测 `TransferQueue==0.1.7` 通过 Ray named actor 发现 `TransferQueueController`。Producer 和 Consumer 仍是普通 OS 进程,不改成 Ray actor;二者不使用 Ray RPC、Ray `ObjectRef` 或 Ray object store 传输训练样本。hidden-state payload 仍通过 TQ 的 `kv_put/kv_batch_get` 和配置的 MooncakeStore backend 传输。第一版不使用 DataProto/WorkerGroup,不生成长期 hidden-state feature store,也暂不实现复杂重试、自动重启和严格 checkpoint 数据恢复。 - -需要运行的组件: - -| 组件 | 数量 | 作用 | -|---|---:|---| -| vLLM server | 一个或多个 | 加载 target model,执行并行 prefill | -| Ray head | 1 个集群 | 保存 TQ named Controller/Storage actors;不承载 Producer/Consumer 业务 RPC | -| TQ owner | 1 个普通进程 | 连接 Ray,创建并保持 TQ controller/storage,最后统一关闭 TQ | -| Producer | 1 个进程 | 读文件、并发请求 vLLM、写 TQ | -| Consumer | 1 个 torchrun 任务 | 多个 rank 从 TQ 取数并训练 DSpark | - -Producer 和 Consumer 不互相调用,也不通过 HTTP 或 Ray 传样本。二者先连接同一个 Ray 集群,再通过无参 `tq.init()` 找到同一个 TQ,最后使用 TQ KV API 读写样本。 - -## 2. 共同的数据约定 - -这部分由两位开发者共同完成并先合入。建议文件: - -```text -verl_speco/transport/drafter_sample_protocol.py -tests/unit/test_drafter_sample_protocol.py -``` - -### 2.1 一个 key 对应一条样本 - -第一版固定: - -```text -一个输入文件 record -→ 一个 sequence_no -→ 一个 sample_id -→ 一个 TQ sample_key -→ 一个单样本 payload -``` - -`sequence_no` 是输入文件中的样本序号,不是 batch 编号。多个 sample keys 到 Consumer 后才组成训练 batch。 - -例如: - -```python -run_id = "dspark-20260818-a" -sequence_no = 17 -sample_id = "train-000017" - -partition_id = "speco_drafter_features" # 第一版沿用 PR 固定值 -sample_key = ( - "drafter:v1:dspark-20260818-a:" - "000000000017:train-000017" -) -``` - -### 2.2 Partition、key、tag 和 payload 的关系 - -TQ 中逻辑上是: - -```text -TQ 实例 -└── partition_id - └── sample_key - ├── tag - └── fields/payload -``` - -- TQ 实例:Producer 和 Consumer 共同连接的 controller/storage; -- `partition_id`:第一版沿用 PR 固定的 `speco_drafter_features`;每次启动新的 TQ owner,保证该 TQ 实例初始为空; -- `sample_key`:该分区中一条训练样本的地址; -- `tag`:轻量索引,Consumer 用 `kv_list()` 看见; -- `fields/payload`:真正的 Tensor 数据,Consumer 用 `kv_batch_get()` 读取。 - -Producer 写入: - -```python -tq.kv_put( - partition_id=partition_id, - key=sample_key, - fields=fields, - tag=tag, -) -``` - -Consumer 先发现 key: - -```python -all_records = tq.kv_list() -tags_by_key = all_records[partition_id] -``` - -这一步只拿 key 和 tag,不搬运 hidden states。 - -Consumer 再取数据: - -```python -result = tq.kv_batch_get( - partition_id=partition_id, - keys=selected_keys, -) -``` - -`kv_batch_get` 中的 batch 表示“一次读取多个独立 sample keys”,不是这些样本在 Producer 写入时就属于同一个对象。 - -### 2.3 Payload 字段 - -一个 sample key 对应的 `fields`: - -```python -fields = { - "input_ids": input_ids, # CPU int64[L] - "loss_mask": loss_mask, # CPU float32[L] - "position_ids": position_ids, # CPU int64[L] - "hidden_states": hidden_states, # CPU bf16[L,D] - "metadata_json": metadata_bytes, # CPU uint8[M] -} -``` - -| field | 含义 | Consumer 中的用途 | -|---|---|---| -| `input_ids` | 经过 feature window 选择后的 token IDs | DSpark token 输入 | -| `loss_mask` | 每个 token 是否参与 loss | 排除 prompt/padding/无效位置 | -| `position_ids` | 每个 row 对应的序列位置 | 位置编码和对齐校验 | -| `hidden_states` | target 指定层 hidden 沿最后一维拼接 | DSpark target/context 特征 | -| `metadata_json` | 模型、layers、layout、shape、样本身份 | Consumer 校验并恢复 metadata | - -符号: - -- `L`:这条训练 feature 保留的 token row 数; -- `H`:target model hidden size; -- `C`:DSpark context layer 数; -- L1 关闭:`D=C*H`,layout=`dflash_aux`; -- L1 开启:`D=C*H+H`,layout=`dflash_aux_plus_last`。 - -示例:`H=4096,C=5,L=1536`,开启 L1: - -```python -input_ids.shape == [1536] -loss_mask.shape == [1536] -position_ids.shape == [1536] -hidden_states.shape == [1536, 24576] -``` - -`metadata_json` 使用 JSON UTF-8 编码为 Tensor,因为参考 PR 的 bridge 只把 Tensor 放入 TQ fields: - -```python -raw = json.dumps(metadata, sort_keys=True).encode("utf-8") -metadata_bytes = torch.tensor(list(raw), dtype=torch.uint8) -``` - -### 2.4 Tag 字段 - -```python -tag = { - "record_type": "sample", - "status": "ready", - "schema_version": 1, - "run_id": "dspark-20260818-a", - "sequence_no": 17, - "sample_id": "train-000017", - "algorithm": "DSPARK", -} -``` - -tag 只存用于发现和筛选的 flat scalar/string。Consumer 只选择: - -```text -record_type=sample -status=ready -schema_version=1 -run_id=当前 run -algorithm=DSPARK -``` - -### 2.5 Metadata 字段 - -`metadata_json` 解码后至少包含: - -```python -metadata = { - "schema_version": 1, - "run_id": "dspark-20260818-a", - "sample_id": "train-000017", - "sequence_no": 17, - "algorithm": "DSPARK", - "target_model_id": "/models/Qwen3-8B", - "target_model_revision": "revision-or-checksum", - "tokenizer_fingerprint": "sha256:...", - "target_layer_ids": [2, 8, 14, 20, 26, -1], - "hidden_states_layout": "dflash_aux_plus_last", - "hidden_dtype": "bfloat16", - "hidden_shape": [1536, 24576], - "feature_length": 1536, - "full_sequence_length": 1800, - "feature_start": 264, - "feature_end": 1800, - "use_logits": False, -} -``` - -其中: - -- `target_model_id/revision`:hidden states 来自哪个 target checkpoint; -- `tokenizer_fingerprint`:Producer 使用的 tokenizer/template 版本; -- `target_layer_ids`:vLLM 返回和参与拼接的层; -- `hidden_states_layout`:Consumer 应如何拆 hidden 最后一维; -- `feature_length`:payload 中四个主要 Tensor 的第一维; -- `full_sequence_length`:完整 prompt+response 的 token 数; -- `[feature_start,feature_end)`:feature 在完整序列中的范围。 - -### 2.6 共享协议接口 - -Producer 和 Consumer 都 import 同一个模块,不各自手写字段名: - -```python -@dataclass(frozen=True) -class SampleMetadata: - schema_version: int - run_id: str - sample_id: str - sequence_no: int - algorithm: str - target_model_id: str - target_model_revision: str - tokenizer_fingerprint: str - target_layer_ids: list[int] - hidden_states_layout: str - hidden_dtype: str - hidden_shape: list[int] - feature_length: int - full_sequence_length: int - feature_start: int - feature_end: int - use_logits: bool - -def make_sample_key(meta: SampleMetadata) -> str: ... -def make_ready_tag(meta: SampleMetadata) -> dict: ... -def encode_sample(sample, meta: SampleMetadata) -> dict[str, Tensor]: ... -def decode_sample(key, tag, fields, expected_config) -> DraftFeatureSample: ... -def make_eos_record(run_id: str, total_samples: int): ... -``` - -Producer 使用 `SampleMetadata/make_sample_key/make_ready_tag/encode_sample`;Consumer 使用 `decode_sample`。 - -`SampleMetadata` Python 对象本身不经过 TQ: - -```text -Producer SampleMetadata -→ JSON -→ uint8 Tensor -→ TQ metadata_json -→ uint8 Tensor -→ JSON -→ Consumer metadata dict -``` - -`decode_sample()` 负责: - -1. 解码 `metadata_json`; -2. 校验 key、tag、metadata 中的 sample 身份一致; -3. 校验模型、tokenizer、layers 和 layout 与 Consumer 配置一致; -4. 校验 Tensor 必需字段、dtype 和 shape; -5. 返回现有 `DraftFeatureSample`。 - -### 2.7 EOS - -Producer 完成全部输入后写一个控制 record: - -```python -eos_key = f"control:v1:{run_id}:eos" -eos_fields = {"marker": torch.tensor([1], dtype=torch.uint8)} -eos_tag = { - "record_type": "control", - "status": "eos", - "schema_version": 1, - "run_id": run_id, - "total_samples": total_samples, -} -``` - -EOS 不进入训练 batch。Consumer 看到 EOS 后继续处理剩余 ready samples;`EOS 已出现且 ready 为空` 时结束。 - -## 3. TQ 怎么启动,Producer 和 Consumer 怎么连接同一个 TQ - -### 3.1 已验证的 TQ 0.1.7 连接机制 - -`TransferQueue==0.1.7` 没有提供“把 Controller 地址直接传给第二个进程”的高层连接接口。它的无参 `tq.init()` 内部执行: - -```python -_TQ_CONTROLLER = ray.get_actor("TransferQueueController") -conf = ray.get(_TQ_CONTROLLER.get_config.remote()) -_maybe_create_tq_client(conf) -``` - -因此所有进程必须先加入同一个 Ray 集群。Ray 在本方案中只承担 TQ 控制面:保存 named `TransferQueueController`、返回 TQ 配置、管理 TQ 创建的 actor。业务进程不通过 Ray 发送 sample key 或 hidden-state tensor。 - -实际连接链路是: - -```text -TQ owner:ray.init(address) → tq.init(full_tq_config) → 创建 named Controller -Producer:ray.init(address) → tq.init() → ray.get_actor() → 创建本地 TQ client -Consumer rank 0..N:ray.init(address) → tq.init() → ray.get_actor() → 创建各自 TQ client -``` - -### 3.2 直接移植并扩展 PR #48 的 bridge - -参考文件: - -```text -C:/Users/xxyyrr/Desktop/上班/verl-SpeCo/ - verl_speco/integration/transferqueue_bridge.py -``` - -第一版不新增 `standalone_tq.py`,直接把该文件移植到目标项目并扩展。保留 PR 已有的 import 检查、进程内幂等状态、`kv_put`、`kv_batch_get`、TensorDict 解包和 owner-only shutdown;新增 Ray 显式连接、批量 list/get/clear 和本地 client 关闭接口。 - -目标文件必须实现以下函数,而不只是提供一个笼统的 transport class: - -```python -def configure_transfer_queue(config: Mapping[str, Any]) -> bool: ... -def connect_ray_cluster(ray_address: str, namespace: str | None = None) -> None: ... -def start_transfer_queue_owner(tq_config: Mapping[str, Any]) -> None: ... -def connect_transfer_queue_client() -> None: ... -def make_sample_key(run_id: str, sequence_no: int, sample_id: str) -> str: ... -def put_sample(key: str, fields: dict[str, Tensor], tag: dict[str, Any]) -> None: ... -def list_samples() -> dict[str, dict[str, Any]]: ... -def get_samples(keys: list[str]) -> list[tuple[str, dict[str, Tensor]]]: ... -def clear_samples(keys: list[str]) -> None: ... -def close_transfer_queue_client() -> None: ... -def close_transfer_queue_owner() -> None: ... -``` - -逐个函数的责任如下。 - -#### `configure_transfer_queue(config)` - -- 从 Hydra/OmegaConf 中读取 Ray address、可选 namespace、固定 partition、run ID、schema version 和 TQ native backend 配置; -- 转成普通 Python dict,保存在进程内 `_state`; -- 校验 `TransferQueue==0.1.7` 可 import; -- 不连接 Ray,不创建 TQ,不产生跨进程副作用; -- 返回该进程是否启用了 TQ。 - -#### `connect_ray_cluster(ray_address, namespace)` - -- 若 `ray.is_initialized()` 为 false,调用 `ray.init(address=ray_address, namespace=namespace)`; -- 若已经初始化,校验现有 Ray context 指向期望集群,不能悄悄连到另一个本地 Ray; -- Owner、Producer 和所有 torchrun ranks 都调用它; -- 此函数只建立当前 OS 进程到 Ray control plane 的连接。 - -#### `start_transfer_queue_owner(tq_config)` - -- 仅由 `tq_owner.py` 调用; -- 前置条件是 `connect_ray_cluster()` 已成功; -- 调用一次 `tq.init(OmegaConf.create(tq_config))`; -- 将 `_state.owner=True`、`_state.initialized=True`; -- TQ 0.1.7 会创建名为 `TransferQueueController` 的 Ray actor,并创建所选 storage backend; -- 重复调用必须报错,不能启动第二套同名 Controller。 - -#### `connect_transfer_queue_client()` - -- 由 Producer 和每个 Consumer rank 调用; -- 前置条件是当前进程已经连接 Ray; -- 调用无参 `tq.init()`,通过 `ray.get_actor("TransferQueueController")` 发现 owner; -- 只创建当前进程的 TQ client,不创建新的 Controller; -- 成功后设置 `_state.initialized=True`;重复调用直接返回。 - -#### `put_sample/list_samples/get_samples/clear_samples` - -- 全部固定使用 `_SPECO_TQ_PARTITION = "speco_drafter_features"`; -- `put_sample()` 调用单样本 `tq.kv_put()`; -- `list_samples()` 调用 `tq.kv_list(partition_id=...)`,只返回 key/tag 元数据; -- `get_samples()` 一次调用 `tq.kv_batch_get(keys=...)`,再按输入 key 顺序解包为普通 dict; -- `clear_samples()` 调用 `tq.kv_clear(keys=..., partition_id=...)`; -- 这些函数不调用 Ray RPC,不把 tensor 放入 Ray object store。 - -#### `close_transfer_queue_client()` 与 `close_transfer_queue_owner()` - -TQ 0.1.7 的公共 `tq.close()` 会 kill 共享 Controller,所以两者不能写成同一个实现: - -- `close_transfer_queue_client()`:通过 0.1.7 已公开的 `tq.get_client()` 取得当前进程 client并调用它的 `close()`,然后 `ray.shutdown()`;绝不能调用会 kill Controller 的全局 `tq.close()`。这个 client close 只用于进程退出阶段,关闭后本进程不能再次调用 TQ; -- `close_transfer_queue_owner()`:仅当 `_state.owner=True` 时调用 `tq.close()`,清理 Controller/Storage,最后 `ray.shutdown()`; -- Producer 或任一训练 rank 提前退出都不能关闭全局 TQ。 - -### 3.3 共享配置 - -Owner、Producer 和 Consumer 必须使用相同的 Ray address/namespace,并通过同一个 named Controller 获得 backend 配置。第一版 partition 固定,不按任务动态创建: - -```yaml -transfer_queue: - enable: true - package_version: "0.1.7" - ray: - address: "ray-head-node:6379" - namespace: "speco-drafter" - partition_id: "speco_drafter_features" - run_id: "dspark-20260819-a" - schema_version: 1 - backend: - storage_backend: MooncakeStore - MooncakeStore: - auto_init: false - metadata_server: "node0:50050" - master_server_address: "node0:50051" - local_hostname: "" - protocol: tcp - global_segment_size: 4294967296 - local_buffer_size: 1073741824 - device_name: "" -``` - -`partition_id/run_id/schema_version` 用于过滤和校验样本;它们不能帮助进程发现 TQ。真正让三个任务连接到同一 TQ 的是“连接同一个 Ray 集群和 namespace,然后找到同名 Controller”。 - -依赖也必须锁定并单独验证。实机安装 `TransferQueue==0.1.7` 会安装 Ray,并要求 `numpy<2.0.0`;当前测试环境中的 `twinkle-kit` 要求 `numpy>=2.0.0`,两者冲突。开发时应使用专门的 TQ/训练环境或重新确认整套依赖约束,不能直接把 0.1.7 安装进已有生产环境后忽略 resolver warning。 - -### 3.4 `tq_owner.py` 要实现的入口和函数 - -新增: - -```text -verl_speco/tq_owner.py -examples/run_dspark_tq_owner.sh -``` - -`tq_owner.py` 建议明确实现: - -```python -def install_signal_handlers(stop_event: threading.Event) -> None: ... -def publish_owner_ready(run_id: str, schema_version: int) -> None: ... -def wait_until_stopped(stop_event: threading.Event) -> None: ... -def run_owner(config: DictConfig) -> int: ... -def main() -> None: ... -``` - -`run_owner()` 的执行顺序必须是: - -```text -configure_transfer_queue(config) -→ connect_ray_cluster(ray.address, ray.namespace) -→ start_transfer_queue_owner(full TQ native config) -→ put owner_ready 控制 record -→ 安装 SIGINT/SIGTERM handler -→ 保持 owner 进程存活 -→ 收到停止信号 -→ close_transfer_queue_owner() -``` - -Owner 必须常驻。TQ 0.1.7 创建 Controller 时没有设置 `lifetime="detached"`,不能在初始化后立即退出。 - -### 3.5 启动和关闭顺序 - -第一版由外部脚本管理全生命周期: - -```text -1. ray start --head,记录 Ray address -2. 启动 Mooncake metadata/master(若 auto_init=false) -3. 启动 TQ owner;owner 连接 Ray并调用 tq.init(full config) -4. 等待 owner_ready -5. 启动一个或多个 vLLM servers -6. 启动 Consumer;每个 torchrun rank 连接 Ray,然后 tq.init() -7. 启动 Producer;连接 Ray,然后 tq.init() -8. Producer 写 EOS,关闭本地 client并退出 -9. Consumer drain、保存 final checkpoint,所有 ranks 关闭本地 client并退出 -10. 给 TQ owner 发送 SIGTERM;仅 owner 执行 tq.close() -11. 等 owner 退出后执行 ray stop -12. 停止 Mooncake 服务 -``` - -外部脚本要用 `trap` 保证异常退出也按“Producer/Consumer → owner → Ray → Mooncake”的顺序清理。不能在 Consumer rank 的 `finally` 中调用全局 `tq.close()`。 - -## 4. Producer 要实现什么 - -### 4.1 Producer 完整顺序 - -```text -读取共享配置 -→ 连接 TQ并校验 owner_ready -→ 初始化 tokenizer -→ 初始化多个 vLLM endpoint clients -→ 流式读取输入文件 -→ 为每条输入分配 sequence_no/sample_id -→ 拼接 prompt+预生成 response,得到 input_ids/loss_mask -→ 并发请求 vLLM prefill -→ 读取 vLLM hidden-state 临时结果 -→ 转换成 DSpark DraftFeatureSample -→ 构造 SampleMetadata -→ encode_sample 得到 fields/tag/key -→ TQ kv_put 一条 sample -→ 删除该请求临时文件 -→ 所有输入完成后写 EOS -→ close_transfer_queue_client()并退出 -``` - -### 4.2 并发模型 - -Producer 是一个进程,内部并发请求多个 endpoint: - -```text -InputReader -→ bounded asyncio input_queue -→ N 个 RequestWorker -→ bounded publish_queue -→ TQ Publisher -``` - -- `vllm_endpoints` 是列表; -- 每个 endpoint 有独立 semaphore; -- 总并发由 `max_inflight_requests` 限制; -- input/publish queue 必须有上限; -- TQ `kv_put` 如果是同步 API,用 `asyncio.to_thread()` 调用; -- TQ ready 数达到 `max_pending_samples` 时暂停继续请求,形成简单背压。 - -`sequence_no` 在 InputReader 中分配,不按 vLLM 完成顺序分配。因此并发乱序不会改变 sample key。 - -### 4.3 vLLM 结果转换 - -复用现有 `TargetFeatureReplayer` 的: - -- OpenAI-compatible vLLM 请求; -- `prompt_token_ids` 校验; -- `kv_transfer_params.hidden_states_path`; -- safetensors 加载; -- `[seq,layers,hidden]` 校验; -- feature positions 选择; -- aux layers flatten; -- DSpark L1 时拼 final hidden。 - -不要复制 `_feature_from_vllm_payload()`;把纯转换逻辑提取成公共函数,让旧 file backend 和新 Producer 共同调用。 - -临时文件顺序: - -```text -加载 -→ 校验/转换 -→ TQ put 成功 -→ 删除 -``` - -第一版仍允许每个并发请求产生短期临时文件,但不生成全量 hidden-state 数据集。 - -### 4.4 Producer 文件分工 - -| 文件 | 实现内容 | -|---|---| -| `verl_speco/standalone_tq_producer.py` | CLI/Hydra 入口,连接 TQ,启动 asyncio pipeline,写 EOS和汇总指标 | -| `verl_speco/producer/input_reader.py` | 流式读输入、分配 `sequence_no/sample_id`、tokenize、构造 loss mask | -| `verl_speco/producer/vllm_feature_client.py` | endpoint pool、并发控制、HTTP 请求、临时结果读取和删除 | -| `verl_speco/trainer/target_feature_replay.py` | 提取可复用的 vLLM payload → `DraftFeatureSample` 纯转换函数 | -| `examples/run_dspark_tq_producer.sh` | Producer 配置和启动命令 | -| `tests/unit/test_drafter_sample_protocol.py` | 协议编码、字段和 shape 测试(两人共同) | -| `tests/integration/test_tq_producer_smoke.py` | fake/短 vLLM → TQ sample + EOS | - -Producer 开发者同时负责移植/扩展 `integration/transferqueue_bridge.py` 和实现 `tq_owner.py`,因为这部分与 TQ 写入和连接直接相关。 - -### 4.5 Producer 各文件的函数级实现规格 - -#### `verl_speco/standalone_tq_producer.py` - -需要实现: - -```python -@dataclass -class ProducerStats: - input_count: int - published_count: int - failed_count: int - pending_bytes: int - -async def publish_one(result: PreparedFeature, transport) -> str: ... -async def run_producer(config: DictConfig) -> ProducerStats: ... -def validate_producer_config(config: DictConfig) -> None: ... -def main() -> None: ... -``` - -`main()` 只负责 Hydra/日志/退出码。`run_producer()` 是可测试的业务入口,执行:连接 Ray → 连接 TQ client → 创建 InputReader 和 vLLM client pool → 启动有界 asyncio pipeline → 等待所有请求及 `kv_put` 完成 → 发布 EOS → 关闭本地 client。它不能启动 Ray head、不能创建 TQ owner、不能调用全局 `tq.close()`。 - -`publish_one()` 接收已经完成转换的一条 `PreparedFeature`,调用共享协议的 `encode_sample()` 得到 `(key, fields, tag)`,再通过 bridge 的 `put_sample()` 发布。只有 `put_sample()` 成功返回,`published_count` 才增加,vLLM 临时文件才允许删除。 - -#### `verl_speco/producer/input_reader.py` - -需要实现: - -```python -@dataclass(frozen=True) -class InputRecord: - sequence_no: int - sample_id: str - prompt: str - response: str - source_metadata: dict[str, Any] - -def iter_input_records(path: str) -> Iterator[InputRecord]: ... -def tokenize_record(record: InputRecord, tokenizer, config) -> TokenizedRequest: ... -def build_loss_mask(input_ids: Tensor, prompt_length: int) -> Tensor: ... -``` - -`iter_input_records()` 流式读取,不把全文件载入内存;在这里按文件顺序分配稳定的 `sequence_no`。`tokenize_record()` 拼接已经存在的 prompt/response,不调用模型生成 response;输出至少包含 `input_ids:int64[L]`、`position_ids:int64[L]`、`loss_mask:float32[L]` 和请求 vLLM 所需字段。 - -#### `verl_speco/producer/vllm_feature_client.py` - -需要实现: - -```python -@dataclass(frozen=True) -class VllmEndpoint: - base_url: str - max_concurrency: int - -class VllmFeatureClientPool: - async def start(self) -> None: ... - async def prefill(self, request: TokenizedRequest) -> RawVllmFeature: ... - async def close(self) -> None: ... - -async def request_prefill(endpoint, request) -> VllmResponse: ... -def choose_endpoint(endpoints, state) -> VllmEndpoint: ... -def load_hidden_state_result(response) -> RawVllmFeature: ... -def delete_temporary_result(raw: RawVllmFeature) -> None: ... -``` - -`prefill()` 必须允许多个 coroutine 同时运行;全局 semaphore 限制总并发,每个 endpoint 另有独立 semaphore。`request_prefill()` 只负责 HTTP 请求和响应校验;`load_hidden_state_result()` 负责读取 vLLM 返回的临时 safetensors/path。临时结果的删除不放在 `load_hidden_state_result()`,而由 `publish_one()` 成功后触发。 - -#### `verl_speco/trainer/target_feature_replay.py` - -把当前类内部的纯转换部分抽成: - -```python -def feature_from_vllm_payload( - payload: RawVllmFeature, - request: TokenizedRequest, - feature_config: FeatureContract, -) -> DraftFeatureSample: ... -``` - -它不进行 HTTP、不访问 TQ、不删除文件,只完成 shape/layout/layer 校验、feature-position 选择、aux layer flatten 和可选 L1 final hidden 拼接。现有 file replay backend 和新 Producer 都调用这一函数,避免两套转换规则。 - -#### `examples/run_dspark_tq_producer.sh` - -负责提供同一套: - -```text -RAY_ADDRESS / Ray namespace -run_id / schema_version / 固定 partition -Mooncake/TQ backend 配置 -输入文件和 tokenizer/model 配置 -vLLM endpoint 列表 -max_inflight_requests / per_endpoint_concurrency -``` - -脚本只启动 Producer,不启动 TQ owner 或 Consumer,便于两位开发者独立调试。 - -## 5. Consumer 要实现什么 - -### 5.1 不新写另一套训练器 - -继续使用现有入口: - -```text -draft_train_launcher.py -→ draft_train.py -→ trainer/draft_training_loop.py -→ DrafterBaseTrainer -→ DSparkTrainerBackend -``` - -训练模式继续使用现有 standalone `training.mode=offline`,用 `feature_store.type=tq` 选择 TQ 数据源。不复制 FSDP、loss、optimizer 或 checkpoint 代码。 - -当前 `feature_store.type` 的作用是告诉 `build_feature_store_from_config()` 应创建哪一种训练数据来源: - -| 当前 type | 对象 | 数据来源 | -|---|---|---| -| `torch_shard` | `TorchShardFeatureStore` | 本地 `.pt` shard,当前 separate-training 示例使用它 | -| `token_replay` | `TokenReplayFeatureStore` | 保存 token replay 的 shard | -| `vllm_safetensors` / `safetensors` | `VllmSafetensorsFeatureStore` | 已生成的 safetensors feature | -| `jsonl_token_replay` / `jsonl` | `JsonlTokenReplayFeatureStore` | JSONL token/text replay | - -第一版新增: - -```yaml -actor_rollout_ref: - rollout: - drafter: - training: - mode: offline - feature_store: - type: tq - path: null - shuffle: false - repeat: false -``` - -这里 `offline` 的含义是“独立于 RL 的 standalone drafter training”,不等于数据必须来自磁盘。 - -不能只给现有工厂增加一个 `type=tq` 然后继续使用通用 `DraftFeatureDataLoader`。当前 loader 会: - -```python -keys = list(store.iter_keys(...)) -``` - -它假设数据集 keys 是一个静态快照;空 store 会直接结束,并且各 rank 独立枚举时可能在 Producer 持续写入的过程中看到不同快照。TQ 是流式、会新增并删除 keys 的数据源,所以 `type=tq` 必须选择专用的 `TQFeatureDataLoader`,由 rank 0 统一发现和分配 keys。 - -#### 5.1.1 当前磁盘 feature store 为什么每个 rank 都会取 keys - -`draft_train_launcher.py` 通过 `torchrun` 启动多个训练进程。每个进程都是一个 rank,并且每个 rank 都会独立进入 `run_standalone_draft_training()`、创建 `DraftFeatureDataLoader`、执行它的 `__iter__()`。当前 loader 的核心逻辑是: - -```python -keys = list( - self.store.iter_keys( - shuffle=self.shuffle, - seed=self.seed + epoch, - ) -) -rank_keys = keys[rank::world_size] - -for key in rank_keys: - batch.append(self.store.read(key)) -``` - -因此当前并不是 rank 0 枚举 keys 后再发送给其他 rank,而是所有 rank 都访问相同的静态 feature-store 路径: - -```text -rank 0:iter_keys() 得到完整静态列表 → keys[0::world_size] → read 自己的 samples -rank 1:iter_keys() 得到完整静态列表 → keys[1::world_size] → read 自己的 samples -... -``` - -例如 store 中固定存在: - -```python -keys = ["k0", "k1", "k2", "k3", "k4", "k5", "k6", "k7"] -``` - -当 `world_size=2` 时,两个 rank 都先得到上述完整列表,然后分别计算: - -```python -# rank 0 -rank_keys = keys[0::2] # ["k0", "k2", "k4", "k6"] - -# rank 1 -rank_keys = keys[1::2] # ["k1", "k3", "k5", "k7"] -``` - -这种实现能够成立,是因为 `.pt`、safetensors 或 replay 文件在训练期间是静态数据集。只要所有 rank 使用同一路径和相同 shuffle seed,`iter_keys()` 就会产生相同顺序的 key 快照,各 rank 可以无通信地算出互不重叠的子集。 - -#### 5.1.2 TQ 为什么必须改成 rank 0 发现 keys - -TQ 中的 keys 会在训练期间动态变化:Producer 持续 `kv_put`,训练成功后 Consumer 又执行 `kv_clear`。如果所有 rank 仍然各自调用 `kv_list`,不同调用时刻可能看到不同快照: - -```python -# rank 0 较早调用 -rank0_keys = ["k0", "k1", "k2", "k3"] - -# Producer 随后写入 k4、k5,rank 1 较晚调用 -rank1_keys = ["k0", "k1", "k2", "k3", "k4", "k5"] -``` - -各 rank 再独立切片后,可能得到不同数量的 local samples,进而无法保证它们以相同顺序进入 forward、backward 和梯度 collective。为此,TQ 专用 loader 必须把控制面和数据面分开: - -```text -控制面:rank 0 执行 kv_list,固定本 step 的 global_keys,并向各 rank 分发 local_keys -数据面:每个 rank 使用自己的 local_keys 直接执行 kv_batch_get,从 TQ/Mooncake 读取 hidden-state payload -``` - -rank 0 只发送较小的 key/tag 元数据,不读取并转发其他 rank 的 hidden-state tensor。所有 rank 仍然都会连接同一个 TQ,也都会调用批量读取接口;只有动态 key 的发现和本 step 的 global batch 决策集中在 rank 0。 - -因此 `feature_store.type=tq` 不是只替换底层 `read()`:它同时改变了 key 的发现、分配、等待和删除语义,需要专用 `TQFeatureDataLoader`,并要求训练循环在所有 rank 成功完成 optimizer step 后,由 rank 0 对本 step 的 `global_keys` 执行 `kv_clear`。 - -### 5.2 Consumer 完整顺序 - -```text -torchrun 启动多个 ranks -→ 每个 rank 初始化 torch.distributed -→ 每个 rank 连接同一个 TQ -→ rank 0 校验 owner_ready,并 broadcast 结果 -→ 初始化现有 DSpark trainer -→ rank 0 kv_list 查找 ready sample keys -→ rank 0 选一个 global batch并分给各 rank -→ 每个 rank kv_batch_get 自己的 local keys -→ decode_sample 得到 list[DraftFeatureSample] -→ prepare_training_batch_from_samples() -→ training_step_from_batch() -→ 所有 rank 汇总 success -→ 成功后 rank 0 kv_clear 这个 global batch 的 keys -→ 继续下一批 -→ 看到 EOS 且 ready 为空 -→ 保存 final checkpoint -→ 所有 ranks close_transfer_queue_client()并退出 -``` - -### 5.3 多 rank 如何分 key - -例如: - -```text -world_size=2 -batch_size_per_gpu=2 -global batch size=4 -``` - -rank 0 选出: - -```python -global_keys = ["k10", "k11", "k12", "k13"] -assignments = [ - ["k10", "k11"], # rank 0 - ["k12", "k13"], # rank 1 -] -``` - -通过 `dist.scatter_object_list` 或 broadcast 发送短字符串列表。每个 rank 直接从 TQ 读取自己的 payload,不通过 rank 0 转发 hidden states。 - -### 5.4 从 TQ 到训练 batch - -每个 rank: - -```python -records = tq_transport.get_samples(local_keys) - -samples = [ - decode_sample( - key=key, - tag=tags_by_key[key], - fields=fields, - expected_config=expected_contract, - ) - for key, fields in records -] - -batch = trainer.prepare_training_batch_from_samples( - samples, - step=optimizer_step, -) - -ok = await trainer.training_step_from_batch( - batch, - optimizer_step, -) -``` - -`records` 是多个独立单样本 payload;`samples` 是变长 `DraftFeatureSample` 列表。Consumer transport 层不要直接 `torch.stack()`,现有 trainer/backend 负责对齐和组 batch。 - -### 5.5 删除与结束 - -第一版采用简单逻辑: - -```text -所有 rank get/decode/train 都成功 -→ all_reduce(global_success)=True -→ rank 0 kv_clear(global_batch_keys) -``` - -任何 rank 失败都不 clear。程序报错退出,第一版不自动恢复。 - -EOS 后不足一个 global batch 的尾部,第一版使用 `drop_last=True`:记录数量,rank 0 clear 这些尾部 keys,然后正常结束。 - -checkpoint 仍使用现有 `save_interval_steps` 和 final checkpoint;第一版不保证崩溃后已 clear 数据能够严格重放。 - -### 5.6 Consumer 文件分工 - -| 文件 | 实现内容 | -|---|---| -| `verl_speco/trainer/feature_store.py` | 工厂增加 `type=tq`,返回 `TQFeatureStore`;TQ 类型不要求 `path` | -| `verl_speco/trainer/tq_feature_store.py` | 实现固定 partition 上的 list/get-many/clear/EOS,不伪装静态 `iter_keys()` | -| `verl_speco/trainer/tq_sample_source.py` | 实现专用 `TQFeatureDataLoader`:rank 0 动态 list/filter/sort、key 分配、各 rank get/decode、EOS/tail | -| `verl_speco/trainer/draft_training_loop.py` | 保持 `mode=offline`;`type=tq` 时跳过 `feature_store.path` 必填检查并选择专用 loader;调用现有 prepare/train,global success 后 clear,final checkpoint | -| `verl_speco/draft_train_launcher.py` | `feature_store.type=tq` 时不要求 feature-store path,透传 torchrun/TQ 配置 | -| `verl_speco/config/speco_base.yaml` | 新增 TQ connection、run、poll、batch 等默认配置 | -| `examples/run_dspark_tq_consumer.sh` | Consumer GPU、batch、checkpoint 和共享 TQ 配置 | -| `tests/unit/test_tq_sample_source.py` | key 过滤/排序/分 rank、EOS、decode 调用测试 | -| `tests/integration/test_dspark_tq_consumer_smoke.py` | 两 rank 取不同样本,完成训练并 clear | - -### 5.7 Consumer 各文件的函数级实现规格 - -#### `verl_speco/trainer/feature_store.py` - -修改现有工厂: - -```python -def build_feature_store_from_config(feature_store_cfg, read_only=False): - store_type = str(feature_store_cfg.get("type", "torch_shard")).lower() - if store_type == "tq": - return TQFeatureStore.from_config(feature_store_cfg) - ... -``` - -要求: - -- 保留 `torch_shard/token_replay/vllm_safetensors/jsonl_token_replay` 的当前行为; -- `type=tq` 时不读取 `path`; -- 工厂只创建对象,不在 import 阶段连接 Ray/TQ; -- `TQFeatureStore` 是流式数据源适配器,不强行实现有误导性的静态 `iter_keys()` 和逐条 `read()`。 - -#### `verl_speco/trainer/tq_feature_store.py` - -需要实现: - -```python -@dataclass(frozen=True) -class ReadyEntry: - key: str - tag: dict[str, Any] - -class TQFeatureStore: - @classmethod - def from_config(cls, cfg) -> "TQFeatureStore": ... - def connect(self) -> None: ... - def list_ready(self, run_id: str) -> list[ReadyEntry]: ... - def get_many(self, entries: list[ReadyEntry]) -> list[DraftFeatureSample]: ... - def clear_many(self, keys: list[str]) -> None: ... - def read_eos(self, run_id: str) -> EosMetadata | None: ... - def close_local(self) -> None: ... -``` - -`connect()` 调用共享 bridge 的 `connect_ray_cluster()` 和 `connect_transfer_queue_client()`。`list_ready()` 调用 `list_samples()` 后只保留 `tag.status=ready`、匹配 `run_id/schema_version` 的数据 key,并按 `(sequence_no,key)` 排序。`get_many()` 一次批量读取 fields,然后逐条调用共享 `decode_sample()`;返回顺序必须和 entries 相同。`clear_many()` 只允许 rank 0 在全局 step 成功后调用。`close_local()` 只关闭本 rank client,不能关闭 Controller。 - -#### `verl_speco/trainer/tq_sample_source.py` - -需要实现: - -```python -@dataclass -class TQLocalBatch: - local_keys: list[str] - local_samples: list[DraftFeatureSample] - global_keys: list[str] | None - -class TQFeatureDataLoader: - def __iter__(self) -> Iterator[TQLocalBatch]: ... - def _select_global_batch(self) -> tuple[list[str], dict[str, ReadyEntry]]: ... - def _broadcast_assignments(self, assignments) -> list[ReadyEntry]: ... - def _handle_eos_and_tail(self) -> bool: ... - def clear_completed_batch(self, global_keys: list[str] | None) -> None: ... -``` - -执行责任必须明确: - -- 所有 rank 创建 loader 并调用 `store.connect()`; -- 只有 rank 0 执行 `_select_global_batch()` 和 `kv_list`; -- rank 0 生成 `list[list[ReadyEntry]]` assignments,使用 `torch.distributed.scatter_object_list` 或 broadcast 分发小型 key/tag; -- 每个 rank 对自己的 local entries 调用 `store.get_many()`,hidden states 从 TQ/Mooncake 直接进入该 rank; -- loader yield `TQLocalBatch`,不能只 yield samples,因为训练后清理还需要 `global_keys`; -- 暂时为空时轮询等待,不能像静态 loader 一样结束;只有 EOS 已出现并且 ready 数据 drain 完才停止; -- `drop_last=True` 时由 rank 0 记录并清理不足 global batch 的尾部 key。 - -#### `verl_speco/trainer/draft_training_loop.py` - -需要新增或调整: - -```python -def build_training_source(config, rank, world_size): ... -def all_ranks_succeeded(local_ok: bool, device) -> bool: ... -async def run_tq_training_loop(trainer, loader, config) -> dict[str, Any]: ... -``` - -`build_training_source()` 根据 `feature_store.type` 分支:静态类型继续创建 `DraftFeatureDataLoader`;`tq` 创建 `TQFeatureDataLoader`,跳过 `feature_store.path` 必填检查,并禁止 `shuffle/repeat`。`run_tq_training_loop()` 对每个 `TQLocalBatch` 调用现有 `prepare_training_batch_from_samples()` 和 `training_step_from_batch()`;所有 rank 通过 collective 得到 global success 后,才让 rank 0 调用 `clear_completed_batch(global_keys)`。异常路径不 clear,finally 只执行 `store.close_local()`。 - -#### `verl_speco/draft_train_launcher.py` - -保留当前 `torch.distributed.run` 启动方式,只增加配置校验和环境透传: - -```python -def validate_tq_launch_config(overrides, launch_config) -> None: ... -def build_child_env(config) -> dict[str, str]: ... -``` - -它不启动 Ray head、不调用 `tq.init()`。职责是确认 `type=tq` 时提供了 Ray address/namespace,且不要求 feature-store path;随后把同一 Ray connection 配置传给每个 torchrun 子进程。每个子 rank 自己建立 Ray/TQ client,不能在 launcher 父进程建一个 client后期待 fork 继承。 - -#### `verl_speco/config/speco_base.yaml` - -增加默认字段: - -```yaml -feature_store: - type: torch_shard - path: null - shuffle: true - repeat: true - tq: - ray_address: null - ray_namespace: speco-drafter - partition_id: speco_drafter_features - run_id: null - schema_version: 1 - poll_interval_seconds: 0.5 - connect_timeout_seconds: 120 - drop_last: true -``` - -当 `type=tq` 时运行期覆盖为 `shuffle=false/repeat=false`。backend 的完整 owner 配置只需要 TQ owner 使用;Producer/Consumer client 从 named Controller 获取它,不各自重新创建 backend。 - -#### Consumer 测试必须覆盖的函数边界 - -- `test_feature_store_factory_builds_tq_without_path()`; -- `test_rank0_filters_and_sorts_ready_entries()`; -- `test_nonzero_rank_never_calls_kv_list()`; -- `test_assignments_are_disjoint_and_global_batch_complete()`; -- `test_each_rank_gets_only_local_keys()`; -- `test_decode_preserves_hidden_states_layout()`; -- `test_clear_only_after_all_ranks_success()`; -- `test_failure_does_not_clear()`; -- `test_eos_drains_ready_then_stops()`; -- `test_client_close_does_not_kill_owner()`。 - -## 6. 两个人怎么分工 - -### 共同先完成 - -1. `drafter_sample_protocol.py`; -2. Ray/TQ connection 配置字段; -3. 一个小型 golden sample; -4. 启动一个测试 Ray head 和 TQ owner,独立进程 A put、独立进程 B list/get/clear 的 smoke test; -5. 验证 client 退出不会 kill owner,只有 owner shutdown 才销毁 Controller。 - -### 开发者 A:Producer/TQ - -负责: - -```text -integration/transferqueue_bridge.py -tq_owner.py -standalone_tq_producer.py -producer/input_reader.py -producer/vllm_feature_client.py -target_feature_replay.py 的公共转换函数 -owner/producer 启动脚本 -Producer/TQ 测试 -``` - -开发者 A 的可交付接口不是“提供一个 TQ 类”,而是: - -```text -bridge:connect_ray_cluster/start_owner/connect_client/put/list/get_many/clear/client_close/owner_close -owner:run_owner/main/signal handler/owner_ready -producer:run_producer/publish_one/统计与 EOS -input reader:iter_input_records/tokenize_record/build_loss_mask -vLLM client:endpoint pool/request_prefill/load/delete -feature conversion:feature_from_vllm_payload -``` - -开发者 B 可以先针对这些接口写 fake transport,不需要等待真实 vLLM 和 Mooncake 联通。 - -### 开发者 B:Consumer/训练 - -负责: - -```text -feature_store.py 的 type=tq 工厂分支 -tq_feature_store.py -tq_sample_source.py / TQFeatureDataLoader -draft_training_loop.py 的 offline + type=tq 分支 -draft_train_launcher.py 配置适配 -speco_base.yaml Consumer 配置 -Consumer 启动脚本 -Consumer/DSpark 测试 -``` - -开发者 B 的可交付接口是: - -```text -feature-store factory:type=tq 分支 -TQFeatureStore:connect/list_ready/get_many/clear_many/read_eos/close_local -TQFeatureDataLoader:select global batch/distribute local keys/get/yield/EOS-tail -training loop:build source/train/global success/clear/final checkpoint -launcher:TQ 配置校验和 torchrun 子进程环境透传 -``` - -### 联调入口 - -建议再提供: - -```text -examples/run_dspark_tq_pipeline_local.sh -``` - -只用于单机联调,顺序启动: - -```text -ray start --head -→ Mooncake metadata/master -→ TQ owner(ray.init + tq.init(full config)) -→ owner_ready -→ vLLM health check -→ Consumer -→ Producer -→ 等 Producer/Consumer 退出 -→ SIGTERM TQ owner(owner 执行 tq.close) -→ ray stop -→ 停止 Mooncake -``` - -最小联调:Producer 发布 8 条样本,2 个 Consumer ranks、每 rank batch size 2,完成 2 个 optimizer steps,8 个 sample keys 被清理,EOS 后保存 final checkpoint。 - -## 7. 第一版验收标准 - -1. TQ owner、Producer 和所有 Consumer ranks 日志显示相同 Ray address/namespace、TQ Controller、partition 和 run ID。 -2. 在同一 Ray 集群中,独立 owner 创建 TQ 后,独立进程 A put,独立进程 B 能 list/get/clear。 -3. Producer 对多个 vLLM endpoints 并发请求,不串行访问。 -4. 一个输入 record 只生成一个 sample key 和一个 payload。 -5. `kv_list` 只拿 tag;`kv_batch_get` 才拿 hidden states。 -6. Consumer 各 rank 读取不同 local keys,不通过 rank 0 搬运 Tensor。 -7. `decode_sample` 能拒绝模型、layer、layout、dtype 或 shape 不匹配的数据。 -8. 所有 rank 训练成功后才 clear 当前 global batch。 -9. Producer 先完成时,Consumer 能 drain 后再退出。 -10. 不产生长期 hidden-state feature store。 -11. Producer 或任一 Consumer rank 退出不会销毁 TQ Controller;只有 owner 调用全局 `tq.close()`。 -12. Ray object store 中不承载 hidden-state payload,训练 tensor 通过 TQ/Mooncake 路径读取。 - -## 8. 后续建议:第一版跑通后再做 - -以下内容不进入第一版开发: - -- Producer HTTP/TQ 复杂重试和 endpoint 熔断; -- Producer 发布 journal,避免重启后重复生成已 clear 样本; -- 自动重启;第一版失败后人工停止整条 pipeline并使用新 `run_id` 重跑; -- Consumer 从最新 checkpoint 自动恢复; -- checkpoint 成功后再 clear 的严格提交窗口; -- TQ owner/storage 整体丢失后的数据重建; -- 多个独立 Consumer 竞争同一 partition; -- lease、ack、超时回收和 exactly-once; -- 动态扩缩容; -- vLLM server 直接写 TQ。 - -第一版先保证:同一个 TQ 能连通、Producer 能并发生产、Consumer 能正确取数训练、每步成功后能及时清理。 From 7a71151f5370634d945215e349ba84bf841541a7 Mon Sep 17 00:00:00 2001 From: xxyyrr598 <1615681145@qq.com> Date: Thu, 20 Aug 2026 16:14:16 +0800 Subject: [PATCH 33/50] Add standalone TQ consumer implementation - Add TQFeatureStore and TQFeatureDataLoader streaming consumer - Wire standalone draft training loop to the TQ consumer path - Move smoke/owner scripts from examples/ to tools/ - Add TQ consumer unit tests and implementation docs --- ...sync_vllm_mooncake_dspark_training_plan.md | 2896 +++++++++++++++++ docs/standalone_tq_consumer_implementation.md | 725 +++++ ...standalone_tq_foundation_implementation.md | 1037 ++++++ ...standalone_vllm_tq_dspark_training_plan.md | 1131 +++++++ tests/unit/test_draft_train_launcher.py | 44 + tests/unit/test_draft_training_loop.py | 58 + tests/unit/test_tq_consumer.py | 265 ++ tools/run_dspark_tq_consumer.sh | 48 + tools/run_dspark_tq_consumer_test.sh | 84 + {examples => tools}/run_dspark_tq_owner.sh | 5 +- tools/run_tq_connection_smoke.sh | 53 + {examples => tools}/tq_connection_smoke.py | 390 ++- tools/tq_delayed_test_producer.py | 234 ++ verl_speco/config/speco_base.yaml | 1 + verl_speco/draft_train_launcher.py | 28 + verl_speco/trainer/draft_training_loop.py | 133 +- verl_speco/trainer/feature_store.py | 16 +- verl_speco/trainer/tq_feature_store.py | 235 ++ verl_speco/trainer/tq_sample_source.py | 202 ++ 19 files changed, 7365 insertions(+), 220 deletions(-) create mode 100644 docs/async_vllm_mooncake_dspark_training_plan.md create mode 100644 docs/standalone_tq_consumer_implementation.md create mode 100644 docs/standalone_tq_foundation_implementation.md create mode 100644 docs/standalone_vllm_tq_dspark_training_plan.md create mode 100644 tests/unit/test_tq_consumer.py create mode 100644 tools/run_dspark_tq_consumer.sh create mode 100644 tools/run_dspark_tq_consumer_test.sh rename {examples => tools}/run_dspark_tq_owner.sh (83%) create mode 100644 tools/run_tq_connection_smoke.sh rename {examples => tools}/tq_connection_smoke.py (77%) create mode 100644 tools/tq_delayed_test_producer.py create mode 100644 verl_speco/trainer/tq_feature_store.py create mode 100644 verl_speco/trainer/tq_sample_source.py diff --git a/docs/async_vllm_mooncake_dspark_training_plan.md b/docs/async_vllm_mooncake_dspark_training_plan.md new file mode 100644 index 00000000..6cda0e40 --- /dev/null +++ b/docs/async_vllm_mooncake_dspark_training_plan.md @@ -0,0 +1,2896 @@ +# verl-SpeCo 异步 vLLM → TransferQueue(Mooncake 后端)→ DSpark 流式训练方案 + +## 1. 文档范围 + +本文只讨论当前 `verl-SpeCo-ls` 项目的独立草稿模型训练入口: + +```text +examples/run_qwen3-8b_drafter_separate_training.sh + → python -m verl_speco.draft_train_launcher + → torch.distributed.run + → python -m verl_speco.draft_train + → run_standalone_draft_training() +``` + +目标是训练 `mode=train/offline` 的草稿模型,不是 RL 与草稿模型一起训练,也不是先生成全量 hidden states 再长期保存到磁盘。 + +输入文件已经包含提前生成好的 response。新流水线需要: + +1. Producer 读取 prompt 和预生成 response,构造完整 token 序列; +2. Producer 并行请求一个或多个 vLLM endpoint 做 prefill; +3. hidden states 写入 TransferQueue(简称 TQ),TQ 的数据后端使用 Mooncake; +4. DSpark 的各训练 rank 从 TQ/Mooncake 并行读取自己的 local batch; +5. 所有 rank 完成同一个 optimizer step 后,清理这一批 TQ 数据; +6. 不使用自研 Stream Coordinator;第一版复用 PR #48 已验证的 TQ KV 传输模式,由 TQ 保存 payload 和 tags,rank 0 通过 `kv_list` 发现 ready key 并组成 global batch。 + +本文不照搬 AngelSpec 的进程组织。实现依据是旁边只读参考仓库 `verl-SpeCo` 的 PR #48 分支,尤其是: + +```text +verl_speco/integration/transferqueue_bridge.py +verl_speco/integration/sglang_runtime.py +verl_speco/integration/oldlogprob_runtime.py +verl_speco/workers/speco_worker.py +verl_speco/integration/task_runner.py +``` + +参考仓库只用于理解和复用设计,不修改其代码。真正实现仍放在本项目 `verl-SpeCo-ls`,服务于 standalone drafter training。 + +建议按下面顺序阅读: + +1. 第一部分先介绍 verl/SpeCo 中的进程、`TokenOutput`、`DataProto`、WorkerGroup、replica 和 owner rank; +2. 再完整解释 PR #48 的 SGLang TQ 路径; +3. 再解释 PR #48 的 old-logprob TQ 路径; +4. 再解释 hidden 恢复后如何进入 online buffer 和 optimizer step; +5. 第二部分才讨论如何把相同 TQ 传输能力改造成独立 Producer/DSpark Consumer。 + +## 第一部分:PR #48 原始 TQ 流程 + +这一部分只解释只读参考仓库 `../verl-SpeCo` 的 PR #48。这里的 Producer、Consumer、Ray driver、数据结构和生命周期全部是 PR #48 当前代码的行为,不是本文为 standalone 设计的行为。 + +### 1A. 阅读 PR #48 前必须知道的项目对象 + +#### SGLang server + +SGLang server 是 rollout 推理进程。它接收 prompt,执行 target model 推理并生成 response。SpeCo 的 SGLang patch 还会在推理过程中收集指定层 hidden states。 + +它不是 drafter trainer,也不执行 optimizer step。在 PR #48 中它是 hidden-state Producer。 + +#### TokenOutput + +`TokenOutput` 是一次 rollout request 的返回对象。核心字段可以概括为: + +```python +TokenOutput( + token_ids=list[int], + log_probs=..., + routed_experts=..., + extra_fields={ + "global_steps": int, + "drafter_sample": dict | None, + }, +) +``` + +`token_ids` 是生成结果;`extra_fields` 是 SpeCo 添加的旁路字段。`drafter_sample` 不影响正常 response 返回,它用于把草稿训练所需信息从 rollout 侧带回训练控制层。 + +#### DataProto 和 non_tensor_batch + +verl 将一批 rollout 结果整理成 `DataProto`。它通常分为: + +```python +DataProto( + batch=TensorDict(...), + non_tensor_batch={...}, + meta_info={...}, +) +``` + +- `batch`:规则的批量 tensor,例如 prompts、responses、attention mask; +- `non_tensor_batch`:不能直接组成规则 dense batch 的 Python 对象或 object array; +- `meta_info`:批次级配置和指标。 + +每个 request 的 `TokenOutput.extra_fields["drafter_sample"]` 最终会被 rollout/agent-loop 聚合到: + +```python +gen_batch_output.non_tensor_batch["drafter_sample"] +``` + +所以 SGLang 侧写的是单 request `TokenOutput`,driver 侧拿到的是批量 `DataProto`。 + +#### RayPPOTrainer driver + +`SpecoRayPPOTrainer` 所在进程是控制进程,简称 driver。它负责按训练 step 依次调用 rollout、old-logprob、actor update,以及向各 WorkerGroup 发 RPC。 + +driver 不执行 drafter 模型 forward。PR #48 之前,它会接触包含 hidden tensor 的 `drafter_sample`;PR #48 之后,它只处理中小 tensor、标量 metadata 和 TQ key。 + +#### WorkerGroup + +WorkerGroup 是 verl 对一组 Ray worker 的调用封装。代码: + +```python +self.drafter_wg.collect_rollout_features(buckets) +``` + +不是普通本地函数调用,而是根据注册的 dispatch rule,将参数分发到多张卡上的 `SpecoWorker.collect_rollout_features()`。 + +#### Rollout replica + +rollout replica 是一组共同承载一个 rollout 模型副本的进程/GPU。一个任务可能存在多个 data-parallel rollout replica。`replica_rank` 标识当前 sample 是哪个 rollout replica 生成的: + +```text +replica_rank = 0, 1, 2, ... +``` + +#### Drafter training replica、DP rank 和 SP rank + +drafter 训练也可能按 data parallel 和 sequence parallel 组织: + +```text +drafter replica / DP rank 0 + ├─ SP rank 0 + └─ SP rank 1 + +drafter replica / DP rank 1 + ├─ SP rank 0 + └─ SP rank 1 +``` + +同一个 drafter DP replica 内的 SP ranks 共同执行一个模型副本的训练。`replica_rank` 用来把 rollout replica 产生的数据路由到对应 drafter DP replica。 + +#### Owner rank + +`collect_rollout_features` 注册了: + +```python +@register( + dispatch_mode=make_nd_compute_dispatch_fn( + mesh_name="drafter_owner_route" + ) +) +``` + +每个 drafter DP replica 的 `SP rank 0` 被标记为 collect leader/owner。driver 传入的是按 replica 分好的 bucket,dispatch 层负责把 bucket 发到对应训练组。一个 replica 内可能有多个 rank 需要共同训练,但只有指定 leader 负责汇总 RPC 返回。 + +这里的 owner route 是 PR #48 为什么不能简单让任意 worker 随机拿 sample 的原因:sample 必须进入与当前 drafter device mesh 一致的训练组。 + +## 2. PR #48 改造前的 online 特征流程 + +PR #48 改造的是 SpeCo 的 online drafter feature transport。hidden states 有两条主要来源。 + +### 2.1 SGLang rollout hidden 路径 + +改造前: + +```text +SGLang server + → 生成 drafter_sample,其中直接包含 hidden_states CPU tensor + → TokenOutput.extra_fields + → RayPPOTrainer driver 收集 drafter_sample + → driver 按 drafter replica/owner 分桶 + → Ray dispatch / object store + → SpecoWorker.collect_rollout_features(samples) + → _store_rollout_sample() + → online drafter buffer/train +``` + +此时 `drafter_sample` 类似: + +```python +{ + "input_ids": int64[1, total_len], + "prompts": int64[1, prompt_len], + "responses": int64[1, response_len], + "hidden_states": bf16[1, hidden_rows, hidden_dim], + "hidden_positions": int64[1, hidden_rows], + "target_logprobs": tensor | None, + "global_step": 42, + "replica_rank": 1, +} +``` + +问题是整个字典经过 driver,而 `hidden_states` 是其中最大的字段。driver 的 host memory 和 Ray object store 都会承载这些 tensor。 + +### 2.2 old-logprob hook hidden 路径 + +另一条路径在 actor old-logprob forward 中捕获 hidden states。改造前,大 chunk 通过: + +```python +chunk_ref = ray.put(hidden_chunk) +``` + +sample 不直接带 tensor,而是带: + +```python +{ + "hidden_states_ref_chunks": [ + { + "ref": ray_object_ref, + "start": 0, + "length": 512, + }, + ], +} +``` + +drafter worker 再 `ray.get(ref)`,根据 `start/length` 切出每条 sample 所需行。 + +### 2.3 PR #48 要改变的边界 + +PR #48 没有改变: + +- rollout 什么时候产生 sample; +- driver 如何触发 drafter worker; +- drafter worker 如何调用 `_store_rollout_sample()`; +- drafter model 的训练逻辑; +- drafter 权重发布。 + +它只改变大 tensor 的跨进程介质: + +```text +改造前:Producer → Ray driver/object store → Consumer +改造后:Producer → TQ storage → Consumer + key 仍走原 Ray 控制路径 +``` + +## 3. PR #48 改造后的完整 TQ 流程 + +### 3.0 总览 + +PR #48 的目标不是让 TQ 自己产生训练 batch,而是把原来经过 Ray driver/Ray object store 的大 hidden tensor 攁到 TQ。原来的控制路径继续存在,只是控制路径上从“大 tensor”变成“小 key”。 + +整体结构是: + +```text + 原 Ray 控制路径 + drafter_sample / chunk ref +Producer ───────────────────── key ───────────────────▶ Consumer + │ │ + │ kv_put(large tensor) │ kv_batch_get(key) + ▼ ▼ +TransferQueue storage ─────────────────────────────────────┘ +``` + +因此 PR #48 同时保留两条通道: + +```text +控制通道:Producer → Ray driver → drafter worker +数据通道:Producer → TQ storage → drafter worker +``` + +控制通道负责告诉 Consumer “本次应该处理哪个样本、对应哪个 TQ key”;数据通道负责传输 hidden states 等大 tensor。 + +#### 3.0.1 配置放在哪里 + +PR #48 在 drafter training 配置下增加: + +```yaml +actor_rollout_ref: + rollout: + drafter: + training: + transfer_queue: + enable: false + backend: + storage_backend: SimpleStorage + SimpleStorage: + total_storage_size: 100000 + num_data_storage_units: 8 +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/config/speco_base.yaml +``` + +这里的 `transfer_queue` 是 SpeCo drafter 自己的配置,不是 standalone `feature_store.type`,也不是简单把上游 verl 的 `transfer_queue.enable` 打开。 + +#### 3.0.2 TaskRunner 创建整套 TQ + +RL 任务启动时,`SpecoTaskRunner.run()` 在创建 workers 之前调用: + +```python +from verl_speco.integration.transferqueue_bridge import ( + close_transfer_queue, + init_transfer_queue, +) + +transfer_queue_started = init_transfer_queue(config) +try: + trainer.init_workers() + trainer.fit() +finally: + if transfer_queue_started: + close_transfer_queue() +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/integration/task_runner.py:319 +``` + +`init_transfer_queue(config)` 内部读取: + +```python +config.actor_rollout_ref.rollout.drafter.training.transfer_queue +``` + +然后执行: + +```python +tq.init(_to_plain_dict(tq_cfg)) +``` + +并记录: + +```python +_state["initialized"] = True +_state["owner"] = True +``` + +这里的 owner 是“创建/拥有 TQ 生命周期的进程”。只有 owner 在任务结束时执行 `tq.close()`。 + +关键顺序是: + +```text +SpecoTaskRunner +→ tq.init(完整配置) +→ trainer.init_workers() +→ Ray workers 启动 +``` + +也就是说,PR #48 假设 TQ Controller/storage 已经由 TaskRunner 在 Ray 集群环境中建立,后启动的 worker 只需要连接。 + +#### 3.0.3 每个 Producer/Consumer 进程怎么连接 TQ + +bridge 中的 `_ensure_initialized()` 是进程级懒初始化: + +```python +def _ensure_initialized(): + if _state["initialized"]: + return + + with _state_lock: + if _state["initialized"]: + return + + tq.init() + _state["initialized"] = True +``` + +注意这里是: + +```python +tq.init() +``` + +不是: + +```python +tq.init(config) +``` + +无参初始化的含义是连接 TaskRunner 已创建的同一套 TQ。它依赖 PR #48 所处的 Ray 运行环境完成服务发现。 + +因此 PR #48 不是“每个 worker 各创建一套 TQ”,而是: + +```text +TaskRunner:tq.init(config),创建一次 +SGLang producer:tq.init(),连接 +actor producer:tq.init(),连接 +drafter consumer:tq.init(),连接 +``` + +#### 3.0.4 SGLang Producer 怎么写 hidden states + +SGLang 已经完成 rollout,并组装出 `drafter_sample` 后,PR #48 执行: + +```python +configure_transfer_queue(training_cfg) + +if is_transfer_queue_enabled(): + tq_key = make_sample_key( + collection_global_steps, + self.replica_rank, + request_id, + ) + + tq_payload = { + "hidden_states": hidden_states.unsqueeze(0).cpu(), + } + + if target_logprobs is not None: + tq_payload["target_logprobs"] = ( + target_logprobs.unsqueeze(0).cpu() + ) + + put_sample( + tq_key, + tq_payload, + tag={ + "global_step": collection_global_steps, + "replica_rank": self.replica_rank, + }, + ) + + drafter_sample["hidden_states_tq_key"] = tq_key + drafter_sample["hidden_states"] = None +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/integration/sglang_runtime.py:2020 +``` + +这里发生了两条不同的数据流: + +```text +大 tensor:SGLang → TQ +小字典/key:SGLang → 原 Ray/driver 路径 → drafter worker +``` + +写进 TQ 后将: + +```python +drafter_sample["hidden_states"] = None +``` + +是为了避免 hidden states 继续经过 driver/Ray object store。driver 仍收到 sample,但大 tensor 已替换成: + +```python +drafter_sample["hidden_states_tq_key"] +``` + +#### 3.0.5 `put_sample()` 实际怎么写 + +bridge 中: + +```python +def put_sample(key, tensor_dict, *, tag=None): + payload = { + k: v + for k, v in tensor_dict.items() + if torch.is_tensor(v) + } + + _ensure_initialized() + + tq.kv_put( + key=key, + partition_id="speco_drafter_features", + fields=payload, + tag=tag or {}, + ) +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/integration/transferqueue_bridge.py:206 +``` + +这里可以明确看到: + +- PR #48 使用 TQ 高层 KV API; +- 一个 key 对应一个 sample; +- `fields` 是 tensor 字典; +- `tag` 是小 metadata; +- partition 当前写死为 `speco_drafter_features`; +- 写入前 tensor 已 `.cpu()`; +- 写入失败直接抛异常,不静默回退。 + +key 的生成代码是: + +```python +def make_sample_key(global_step, replica_rank, request_id): + return f"speco:{global_step}:{replica_rank}:{request_id}" +``` + +这个 key 对 RL rollout 是合理的,因为它用 step、rollout replica 和 request ID 标识一次在线采样。 + +#### 3.0.5.1 PR #48 中“数据”和“元数据”分别长什么样 + +PR #48 一条 SGLang 样本在写 TQ 之前,`drafter_sample` 同时包含大 tensor 和控制信息。简化后类似: + +```python +drafter_sample = { + # 普通训练输入,仍走原 sample/Ray 控制路径 + "input_ids": int64[1, prompt_len + response_len], + "prompts": int64[1, prompt_len], + "responses": int64[1, response_len], + + # 大 tensor,开启 TQ 后从这个字典移除 + "hidden_states": bf16[1, hidden_rows, hidden_dim], + "target_logprobs": fp32[1, rows, topk_or_vocab] | None, + + # hidden 与 token 对齐所需的小字段 + "hidden_positions": int64[1, hidden_rows] | None, + "hidden_position_start": int, + "hidden_position_end": int, + "hidden_window_start": int, + "hidden_window_end": int, + + # 控制信息 + "global_step": int, + "replica_rank": int, +} +``` + +执行 `put_sample()` 时,并不是把整个 `drafter_sample` 放进 TQ。PR #48 只抽取占用大的 tensor: + +```python +tq_payload = { + "hidden_states": bf16[1, hidden_rows, hidden_dim], + "target_logprobs": fp32[1, rows, ...], # 可选 + "hidden_raw_target_logprobs": ..., # 可选 + "hidden_raw_target_logprobs_positions": ..., # 可选 +} +``` + +这就是 TQ 的 data payload。它被传给: + +```python +tq.kv_put(fields=tq_payload) +``` + +另外还有 TQ tag: + +```python +tag = { + "global_step": 42, + "replica_rank": 1, +} +``` + +tag 是 TQ 侧轻量 metadata,用于描述/检索对象,不承载 hidden tensor。 + +写入完成后,仍经 Ray 传递的轻量 `drafter_sample` 变成: + +```python +drafter_sample = { + "input_ids": int64[1, total_len], + "prompts": int64[1, prompt_len], + "responses": int64[1, response_len], + + "hidden_states": None, + "target_logprobs": None, + "hidden_states_tq_key": "speco:42:1:req-007", + + "hidden_positions": int64[1, hidden_rows] | None, + "hidden_position_start": 128, + "hidden_position_end": 640, + "global_step": 42, + "replica_rank": 1, +} +``` + +因此 PR #48 实际存在三类对象: + +| 对象 | 内容 | 传输路径 | 作用 | +|---|---|---|---| +| TQ fields/payload | `hidden_states` 等大 tensor | Producer → TQ storage → Consumer | 避免 driver 搬运大 tensor | +| TQ tag | `global_step`、`replica_rank` | TQ control/index metadata | 描述该 key | +| 轻量 `drafter_sample` | tokens、位置、标量 metadata、`hidden_states_tq_key` | 原 Ray driver 路径 | 告诉 Consumer 取哪个 key,以及如何解释取回的 tensor | + +代码实现解耦的关键不是“所有内容都进 TQ”,而是: + +```python +drafter_sample["hidden_states_tq_key"] = tq_key +drafter_sample["hidden_states"] = None +``` + +第一行给 Consumer 留下寻址信息;第二行阻止大 tensor 继续沿旧路径传输。 + +#### 3.0.5.2 Consumer 如何把两部分重新合成一个训练样本 + +Consumer 最初拿到的是轻量 sample: + +```python +sample["hidden_states"] is None +sample["hidden_states_tq_key"] == "speco:42:1:req-007" +``` + +它执行: + +```python +payload = get_sample(sample["hidden_states_tq_key"]) +sample["hidden_states"] = payload["hidden_states"] +``` + +合并后: + +```python +sample = { + "input_ids": ..., + "prompts": ..., + "responses": ..., + "hidden_positions": ..., + "hidden_states": bf16[1, hidden_rows, hidden_dim], + "hidden_states_tq_key": "speco:42:1:req-007", + ... +} +``` + +后面的 `_store_rollout_sample()` 看到的结构与关闭 TQ 时基本一致,所以训练主体不需要增加 TQ 分支。TQ bridge 只改变“大 tensor 从哪里恢复”,不改变 drafter trainer 的输入语义。 + +#### 3.0.6 old-logprob Producer 怎么写 chunk + +PR #48 的另一条 Producer 路径来自 actor old-logprob hidden hook。原来是: + +```python +chunk_ref = ray.put(hidden_chunk) +``` + +开启 TQ 后改成: + +```python +tq_key = make_sample_key( + global_step, + owner, + f"chunk{len(chunk_refs)}", +) + +put_sample( + tq_key, + {"hidden": hidden_chunk}, + tag={ + "global_step": global_step, + "owner": owner, + }, +) + +chunk_ref = tq_key +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/integration/oldlogprob_runtime.py:533 +``` + +后面的 driver 不需要知道 ref 是 Ray ObjectRef 还是 TQ string key,它只把 ref 当不透明 token 继续传递。 + +#### 3.0.7 Consumer 怎么根据 key 读取 + +drafter worker 收到原来的 sample 小字典后: + +```python +tq_key = sample.get("hidden_states_tq_key") + +if tq_key is not None and self._speco_tq_enabled: + payload = get_sample(tq_key) + + for field in ( + "hidden_states", + "target_logprobs", + "hidden_raw_target_logprobs", + "hidden_raw_target_logprobs_positions", + ): + if payload.get(field) is not None: + sample[field] = payload[field] + + if sample.get("hidden_states") is None: + raise RuntimeError( + "TQ key exists but hidden_states payload is missing" + ) +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/workers/speco_worker.py:848 +``` + +恢复 tensor 后,后面的逻辑仍使用原 `sample`/`batch`,drafter trainer 不需要知道 tensor 来自 Ray 还是 TQ。 + +#### 3.0.8 `get_sample()` 实际怎么读 + +```python +def get_sample(key): + _ensure_initialized() + + result = tq.kv_batch_get( + keys=[key], + partition_id="speco_drafter_features", + ) + + value = _extract_value(result, key) + return _tensordict_to_dict(value) +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/integration/transferqueue_bridge.py:237 +``` + +`_extract_value()` 兼容三种返回形态: + +```python +if isinstance(result, dict): + return result.get(key) +if isinstance(result, (list, tuple)): + return result[0] +return result +``` + +这是因为不同 TQ 版本/后端返回包装可能不同。 + +#### 3.0.9 为什么需要 `_densify_tq_tensor()` + +PR #48 后续修复发现,TQ 把单样本 tensor 放进 TensorDict 后,`kv_batch_get` 可能返回 NestedTensor,并额外带 batch 维。旧代码要执行: + +```python +tensor[start:start + length] +``` + +但 NestedTensor 不支持在 jagged dim 直接 slice。因此加入: + +```python +def _densify_tq_tensor(tensor): + if tensor.is_nested: + parts = [ + part + for part in tensor.unbind() + if part.numel() > 0 + ] + tensor = torch.cat(parts, dim=0) + + if tensor.dim() == 3: + tensor = tensor.squeeze(0) + elif tensor.dim() == 1: + tensor = tensor.unsqueeze(0) + + return tensor.contiguous() +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/workers/speco_worker.py:72 +``` + +standalone Consumer 同样必须做这个转换,不能假定 `kv_batch_get` 返回普通 dense `[seq, hidden]`。 + +#### 3.0.10 为什么需要 per-step cache + +old-logprob 路径里,一个 owner hidden chunk 可能被约 16 个 sample 共同引用。如果每个 sample 都: + +```python +get_sample(same_tq_key) +``` + +就会重复传输同一个数百 MB chunk。PR #48 在每次 `collect_rollout_features()` 开始时创建: + +```python +self._tq_chunk_cache = {} +``` + +解析 ref 时: + +```python +cache_key = ref if isinstance(ref, str) else id(ref) + +if cache_key not in cache: + cache[cache_key] = _resolve_tq_or_ray_ref(ref) + +tensor = cache[cache_key] +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/workers/speco_worker.py:98 +../verl-SpeCo/verl_speco/workers/speco_worker.py:854 +``` + +独立训练如果每个 sample 都是独立 TQ key,主要使用 `kv_batch_get(keys=[...])` 批量读取,不一定需要跨 sample chunk cache;但 prefetch 重试或共享 packed object 时仍应保留 key cache。 + +#### 3.0.11 PR #48 什么时候删除数据 + +PR #48 没有在 `get_sample()` 后删除。其注释明确说明:同一 drafter replica 的多个 TP/SP rank 可能读取同一 key,第一次读取后立刻删除会让剩余 rank 失败。 + +当前策略是任务结束时由 owner: + +```python +tq.close() +``` + +统一结束 TQ 生命周期。也就是说,PR #48 当前没有实现精细的逐 step `kv_clear`。 + +#### 3.0.12 PR #48 的完整时序 + +```text +SpecoTaskRunner + → tq.init(config) + → 启动 Ray workers + +SGLang/actor Producer process + → configure_transfer_queue() + → 第一次 put 时 tq.init() + → kv_put(key, tensor fields, tag) + → 把 key 塞回原 sample/ref + +Ray driver + → 只中转小 sample/key + +drafter worker Consumer process + → 第一次 get 时 tq.init() + → kv_batch_get([key]) + → 解包 TensorDict/NestedTensor + → 恢复 sample["hidden_states"] + → 原 drafter collect/train 逻辑 + +任务结束 + → TaskRunner owner tq.close() +``` + +### 3.1 已经实现的可复用能力 + +PR #48 新增 `verl_speco/integration/transferqueue_bridge.py`,锁定思路是把 TQ 当成独立传输库,不修改上游 verl。它提供: + +```python +configure_transfer_queue(training_cfg) +init_transfer_queue(config) +make_sample_key(global_step, replica_rank, request_id) +put_sample(key, tensor_dict, tag=...) +get_sample(key) +close_transfer_queue() +``` + +实际写入调用是: + +```python +tq.kv_put( + key=key, + partition_id="speco_drafter_features", + fields=payload, + tag=tag, +) +``` + +实际读取调用是: + +```python +result = tq.kv_batch_get( + keys=[key], + partition_id="speco_drafter_features", +) +``` + +另外,PR #48 已经处理了多项 standalone 方案也需要的问题: + +1. TQ 返回值可能是 direct value、`{key: value}` 或 list,需要统一解包; +2. TQ 返回的 tensor 可能是 NestedTensor,必须通过 `unbind + cat` 恢复为生产端写入的 dense tensor; +3. 多个样本引用同一 hidden chunk 时,需要 per-step cache,避免重复 `kv_batch_get` 同一个大对象; +4. TQ key 存在但 payload 缺失时 fail loud,不能静默丢样本; +5. `enable=false` 时保留原传输路径。 + +这些逻辑应直接作为本项目 TQ adapter 的参考。 + +### 3.2 PR #48 的数据流 + +PR #48 优化的是 RL online 路径: + +```text +SGLang/actor worker + → kv_put(hidden states) + → 把 hidden_states_tq_key 塞进原 drafter_sample + → 原 Ray driver 继续传递小 sample/key + → drafter worker collect_rollout_features() + → kv_batch_get(key) +``` + +它没有让 consumer 自己从 TQ 发现下一批 key;key 仍沿原来的 Ray 控制路径到达 drafter worker。 + +### 3.3 PR #48 没有提供的 standalone 能力 + +PR #48 当前没有实现: + +- 从预生成 response 文件读取数据的独立 Producer; +- Producer 并行请求外部 vLLM endpoint; +- standalone DSpark trainer 主动发现 ready key; +- global batch 到各 torchrun rank 的分片; +- 每个 optimizer step 后精确 `kv_clear`; +- EOS; +- standalone 无 Ray 的 TQ bootstrap; +- MooncakeStore 的实际运行验证。 + +PR #48 当前配置是: + +```yaml +transfer_queue: + enable: false + backend: + storage_backend: SimpleStorage + SimpleStorage: + total_storage_size: 100000 + num_data_storage_units: 8 +``` + +并且 bridge 注释针对 `TransferQueue==0.1.7`。旁边当前 verl 主线已使用 `0.1.8` 文案并包含 `MooncakeStore` 配置。当前机器没有安装 `transfer_queue` 包,因此正式实现前必须锁定版本并实机验证 API 签名,不能把 0.1.7 和 0.1.8 混用。 + +### 3.4 standalone 方案对 PR #48 的扩展 + +不再自行实现 `publish_ready/claim/lease/ack` HTTP 服务。第一版在 PR #48 KV 模式上补四个操作: + +```python +tq.kv_batch_put(...) # Producer 批量写 +tq.kv_list(...) # rank 0 列出 key + tag +tq.kv_batch_get(...) # 各 rank 并行读 +tq.kv_clear(...) # optimizer step 成功后删 +``` + +第一版不依赖 TQ `Sampler/StreamingDataLoader/get_meta`,因为 PR #48 并未使用或验证这些接口。等 KV 流程稳定后再升级为 TQ StreamingDataLoader。 + +### 3.5 PR #48 与 standalone 独立训练逐项映射 + +| PR #48 online RL | standalone drafter training | +|---|---| +| `SpecoTaskRunner` 调用 `tq.init(config)` | 新增独立 TQ owner/bootstrap 进程,或在确认安全后由 Producer owner 调用 `tq.init(config)` | +| SGLang/actor worker 是 Producer | `feature_producer.py` 是独立 Producer | +| rollout 过程中已经得到 hidden states | Producer 读取预生成 response,再请求外部 vLLM prefill | +| `make_sample_key(global_step, replica_rank, request_id)` | `make_replay_sample_key(dataset,row,tokens,target fingerprint)` | +| `put_sample()` / `kv_put()` | 继续复用同一写入模式,可扩展为 `kv_batch_put()` | +| key 通过 Ray driver/sample 传给 consumer | 没有 driver;rank 0 用 `kv_list()` 主动发现 ready keys | +| drafter worker `get_sample(key)` | 每个 torchrun rank `kv_batch_get(local_keys)` | +| `collect_rollout_features()` 恢复 sample | 转换为 `DraftFeatureSample` 后调用现有 `prepare_training_batch_from_samples()` | +| 多 TP/SP rank 可能读同一 key,所以不立即删除 | data-parallel rank 读取互不重叠的 key;全 rank step 成功后统一 `kv_clear(global_keys)` | +| TaskRunner 结束时 `tq.close()` | 每 step clear;输入 drain 完成后 owner 最后 `tq.close()` | + +standalone 需要新增的控制流是: + +```text +Producer DSpark rank 0 其他 ranks + │ │ │ + │ kv_put(sample key, fields, tag) │ │ + ├────────────────────────────────────▶│ │ + │ │ kv_list READY keys │ + │ │ │ + │ │ broadcast selected_keys ──▶│ + │ │ │ + │ │ kv_batch_get(local keys) │ kv_batch_get(local keys) + │ │ │ + │ ├──── DSpark synchronized step ────┤ + │ │ │ + │ │ kv_clear(global keys) │ +``` + +这个映射中,TQ 同时承担: + +- 大 tensor 存储/传输; +- key、tag 和 partition 的轻量索引。 + +但第一版 global batch 的选择仍由单个 DSpark job 的 rank 0 完成。这样最接近 PR #48 的 KV API,避免在同一次改造中再引入未经该 PR 验证的 Sampler/StreamingDataLoader。 + +### 3.6 PR #48 的 SGLang 路径:逐函数、逐对象完整流程 + +下面从一次生成请求开始,不省略中间层。 + +#### 阶段 1:SGLang完成生成并收集 hidden states + +执行进程:SGLang rollout server。 + +输入是一次 request 对应的 prompt 和生成配置。生成结束时,代码已经持有: + +```python +prompt_tensor: int64[prompt_len] +response_tensor: int64[response_len] +hidden_states: bf16[hidden_rows, hidden_dim] +hidden_positions: int64[hidden_rows] | None +target_logprobs: tensor | None +request_id: str +collection_global_steps: int +self.replica_rank: int +``` + +这些变量的语义: + +- `prompt_tensor`:输入 prompt token IDs; +- `response_tensor`:SGLang生成的 response token IDs; +- `hidden_states`:target model 指定层在部分 token positions 上的输出; +- `hidden_positions`:每个 hidden row 对应完整 `prompt+response` 序列中的哪个 token position; +- `target_logprobs`:可选的目标概率监督; +- `request_id`:当前 rollout request 标识; +- `replica_rank`:执行该 request 的 rollout replica。 + +SGLang 先构造完整 sample: + +```python +drafter_sample = { + "input_ids": torch.cat( + [prompt_tensor, response_tensor], dim=0 + ).unsqueeze(0), + "prompts": prompt_tensor.unsqueeze(0), + "responses": response_tensor.unsqueeze(0), + "hidden_states": hidden_states.unsqueeze(0).cpu(), + "hidden_positions": hidden_positions.unsqueeze(0).cpu(), + "target_logprobs": ( + target_logprobs.unsqueeze(0).cpu() + if target_logprobs is not None + else None + ), + "global_step": collection_global_steps, + "replica_rank": self.replica_rank, + # 还有 hidden window/alignment metadata +} +``` + +前导维 `1` 表示这是一个单样本 batch。`.cpu()` 表示跨进程传输前把大 tensor 放到 CPU 内存。 + +#### 阶段 2:PR #48 将大 fields 写入 TQ + +同一个 SGLang进程执行: + +```python +tq_key = make_sample_key( + collection_global_steps, + self.replica_rank, + request_id, +) + +tq_payload = { + "hidden_states": hidden_states.unsqueeze(0).cpu(), +} + +put_sample( + tq_key, + tq_payload, + tag={ + "global_step": collection_global_steps, + "replica_rank": self.replica_rank, + }, +) +``` + +调用展开后是: + +```python +tq.init() # 当前进程第一次使用时 +tq.kv_put( + key=tq_key, + partition_id="speco_drafter_features", + fields=tq_payload, + tag=tag, +) +``` + +效果是 TQ 中增加一行: + +```text +partition = speco_drafter_features +key = speco:42:1:req-007 +fields = {hidden_states: bf16[1, H, D], ...} +tag = {global_step: 42, replica_rank: 1} +``` + +`kv_put` 返回后,SGLang侧将旧 sample 改成: + +```python +drafter_sample["hidden_states_tq_key"] = tq_key +drafter_sample["hidden_states"] = None +``` + +此后该 Python 字典不再携带 hidden tensor,只携带定位它的 key。 + +#### 阶段 3:把 drafter_sample 放进 TokenOutput.extra_fields + +SGLang返回: + +```python +TokenOutput( + token_ids=token_ids, + log_probs=log_probs, + routed_experts=routed_experts, + extra_fields={ + "global_steps": collection_global_steps, + "drafter_sample": drafter_sample, + }, +) +``` + +此时 `TokenOutput` 中有两类输出: + +- 正常 rollout 输出:`token_ids/log_probs`; +- SpeCo 训练旁路输出:`extra_fields.drafter_sample`。 + +TQ payload 不在 `TokenOutput` 中,只有 `hidden_states_tq_key` 在其中。 + +#### 阶段 4:多个 TokenOutput 聚合成 gen_batch_output + +rollout/agent-loop 层将多个 request 的输出合并为批量 `DataProto`: + +```python +gen_batch_output.non_tensor_batch["drafter_sample"] +``` + +可能是 object array: + +```python +array([ + {"hidden_states_tq_key": "speco:42:0:req-A", ...}, + {"hidden_states_tq_key": "speco:42:1:req-B", ...}, +], dtype=object) +``` + +之所以进入 `non_tensor_batch`,是因为每条 sample 的 hidden window、Python metadata 和可选字段不一定具有统一 dense shape。 + +#### 阶段 5:driver 从 DataProto 取出 drafter samples + +`generate_sequences_with_speco()` 包装原 rollout 调用: + +```python +gen_batch_output = original_generate_sequences(...) +collected = self._speco_collect_generation_samples(gen_batch_output) +``` + +`_speco_collect_generation_samples()` 调用: + +```python +samples = pop_drafter_samples(gen_batch_output) +``` + +`pop_drafter_samples()` 实际执行: + +```python +non_tensor_batch = gen_batch_output.non_tensor_batch +samples_array = non_tensor_batch.pop("drafter_sample", None) +samples = normalize_drafter_samples(samples_array) +``` + +这里 `pop` 有两个作用: + +1. 取得 SpeCo drafter side-channel samples; +2. 从正常 PPO 的 `gen_batch_output` 中移除该旁路字段,避免后续 PPO batch 继续携带它。 + +`normalize_drafter_samples()` 将 dict、object array 或 list 统一成: + +```python +samples: list[dict] +``` + +#### 阶段 6:driver 按 replica_rank 分桶 + +假设有两个 rollout/drafter replicas,收到: + +```python +samples = [ + {"replica_rank": 1, "hidden_states_tq_key": "k1", ...}, + {"replica_rank": 0, "hidden_states_tq_key": "k2", ...}, + {"replica_rank": 1, "hidden_states_tq_key": "k3", ...}, +] +``` + +执行: + +```python +buckets = bucket_drafter_samples_by_replica( + samples, + num_replicas=2, +) +``` + +结果: + +```python +buckets = [ + [sample_k2], # bucket 0 + [sample_k1, sample_k3], # bucket 1 +] +``` + +分桶依据只有: + +```python +owner_rank = int(sample["replica_rank"]) +buckets[owner_rank].append(sample) +``` + +这一步没有读取 TQ,也没有处理 hidden tensor;只对小字典做路由。 + +#### 阶段 7:driver 通过 WorkerGroup RPC 分发 buckets + +driver 调用: + +```python +self._speco_set_drafter_global_step() +self._speco_collect_rollout_features_rpc( + "rollout", + buckets, +) +``` + +RPC 内部调用: + +```python +self.drafter_wg.collect_rollout_features(buckets) +``` + +因为 worker 方法注册了 `drafter_owner_route` dispatch,WorkerGroup 将 `buckets[0]` 发给 drafter DP replica 0,将 `buckets[1]` 发给 drafter DP replica 1。一个训练 replica 内的 SP ranks 根据 mesh dispatch 规则参与对应调用。 + +这里传输的对象仍是: + +```python +list[dict] +``` + +其中包含 tokens、position metadata 和 TQ key,不包含被置空的 hidden states。 + +#### 阶段 8:SpecoWorker 根据 key 从 TQ 恢复 tensor + +目标 worker 执行: + +```python +def collect_rollout_features(self, samples): + for sample in samples: + tq_key = sample.get("hidden_states_tq_key") + payload = get_sample(tq_key) + sample["hidden_states"] = payload["hidden_states"] +``` + +`get_sample()` 展开为: + +```python +tq.init() # 此 Consumer 进程第一次使用时 +result = tq.kv_batch_get( + keys=[tq_key], + partition_id="speco_drafter_features", +) +payload = _extract_value(result, tq_key) +payload = _tensordict_to_dict(payload) +``` + +现在 `sample` 再次包含: + +```python +{ + "input_ids": ..., + "hidden_positions": ..., + "hidden_states": bf16[1, hidden_rows, hidden_dim], + "hidden_states_tq_key": "...", +} +``` + +这与关闭 TQ 时 worker 收到的逻辑内容一致。 + +#### 阶段 9:worker 构造 DrafterBaseTrainer 所需 batch dict + +worker 先保留 token fields: + +```python +batch = { + "input_ids": sample["input_ids"], + "prompts": sample["prompts"], + "responses": sample["responses"], +} +``` + +再复制 hidden alignment metadata,例如: + +```python +batch["hidden_positions"] +batch["hidden_position_start"] +batch["hidden_position_end"] +batch["hidden_states_layout"] +batch["global_step"] +``` + +hidden tensor 单独作为参数: + +```python +self._store_rollout_sample( + batch=batch, + hidden_states=hidden, + target_logprobs=target_logprobs, +) +``` + +#### 阶段 10:样本进入在线 buffer 或落盘 + +`_store_rollout_sample()` 根据 training mode 分支: + +```python +if mode == "collect_only": + self._write_rollout_feature_sample( + batch, + hidden_states, + target_logprobs, + ) +else: + self.trainer.collect_online_data( + batch, + hidden_states, + target_logprobs, + ) +``` + +`collect_only` 会转换成 `DraftFeatureSample` 并写 `TorchShardFeatureStore`。online 模式则进入 `DrafterBaseTrainer.collect_online_data()`。 + +`collect_online_data()` 做: + +1. 将 `input_ids/hidden_states/positions/logprobs` 规范化到 CPU; +2. 按 batch 维拆成逐样本; +3. 根据 `hidden_positions` 校验 hidden row 与 token position; +4. 截取可训练窗口; +5. 构造内部 training item; +6. 保存到当前 step 的 `collected_data`,或在启用 data buffer 时保存到跨 step buffer。 + +因此 TQ get 完成不代表立即 optimizer step。它先恢复 online training sample,再进入现有数据准备逻辑。 + +#### 阶段 11:driver 在 actor update 周期触发 drafter 训练 + +driver 包装了 `update_actor()`: + +```python +should_train_drafter = ( + self._speco_should_attempt_drafter_train_this_step() +) + +actor_output = original_update_actor(...) + +if should_train_drafter: + drafter_trained, metrics = self._speco_train_drafter() +``` + +`_speco_train_drafter()` 再向 WorkerGroup 发: + +```python +self.drafter_wg.train_drafter() +``` + +每个 `SpecoWorker.train_drafter()`: + +1. 检查是否属于 drafter training group; +2. 检查 `training_interval_steps`; +3. 激活 drafter training model; +4. 循环 `train_steps_per_trigger` 次; +5. 每次调用 `self.trainer.training_step(global_step)`; +6. 成功时准备需要发布的 drafter state dict; +7. 清理训练期间临时状态。 + +`training_step()` 从刚才的 online `collected_data/DataBuffer` 组成 batch,执行 drafter forward、loss、backward 和 optimizer step。 + +所以 SGLang TQ 路径的最终效果是: + +```text +TQ 只替换 hidden tensor 跨进程传输 +→ sample 收集逻辑不变 +→ online buffer 不变 +→ drafter training trigger 不变 +→ loss/optimizer 不变 +``` + +### 3.7 PR #48 old-logprob 路径的完整差异 + +old-logprob 路径没有 `TokenOutput.extra_fields`。它从 PPO 的 actor old-logprob forward 开始。 + +#### 阶段 1:driver 构造 collect plan + +driver 根据 batch、collect interval 和 drafter owner 数量决定: + +```python +collect_mask: bool[batch] +hidden_positions: list/tensor per sample +owner_rank: int64[batch] +prompt_lens: int64[batch] +response_lens: int64[batch] +``` + +并把 `global_step` 等控制字段放入 old-logprob micro-batch。 + +#### 阶段 2:actor forward hook 选择 hidden rows + +actor worker 在 old-logprob forward 中捕获指定层 hidden states,根据 `collect_mask/position_mask` 只保留需要训练的样本和 token rows。 + +输出可以是 dense selected tensor,也可以是 sparse rows。随后 `_put_oldlogprob_hidden_refs()` 将同一个 owner 的多条 sample rows 拼成一个较大的 `hidden_chunk`。 + +#### 阶段 3:hidden chunk 写入 TQ + +改造前: + +```python +chunk_ref = ray.put(hidden_chunk) +``` + +PR #48: + +```python +tq_key = make_sample_key( + global_step, + owner, + f"chunk{chunk_index}", +) + +put_sample( + tq_key, + {"hidden": hidden_chunk}, + tag={"global_step": global_step, "owner": owner}, +) + +chunk_ref = tq_key +``` + +TQ fields: + +```python +{"hidden": bf16[total_owner_rows, hidden_dim]} +``` + +控制路径中的 chunk metadata: + +```python +chunk_info = { + "sample_indices": [0, 3, 5], + "starts": [0, 128, 384], + "lengths": [128, 256, 96], + "row_indices": [...], + "dtype": "bfloat16", + "shape": [480, hidden_dim], +} +``` + +`starts/lengths` 描述每条 sample 在共享 chunk 中对应的行区间。 + +#### 阶段 4:driver 将 chunk ref 还原为逐样本引用 + +driver 的 `_speco_collect_oldlogprob_features()` 读取: + +```python +chunk_refs = ["speco:42:0:chunk0", ...] +chunk_meta = [chunk_info, ...] +``` + +然后为每个 batch sample 构造: + +```python +sample["hidden_states_ref_chunks"] = [ + { + "ref": "speco:42:0:chunk0", + "chunk_start": 128, + "chunk_length": 256, + "chunk_row_indices": ..., + "dtype": "bfloat16", + "shape": [480, hidden_dim], + } +] +``` + +同时构造该 sample 的: + +```python +input_ids +prompts +responses +hidden_positions +hidden_states_layout +replica_rank=owner +``` + +再按 owner 放入 `buckets[owner]`,通过同一个 `collect_rollout_features()` RPC 发给 drafter worker。 + +#### 阶段 5:Consumer 获取共享 chunk 并切片 + +drafter worker 发现: + +```python +sample.get("hidden_states") is None +sample.get("hidden_states_ref_chunks") is not None +``` + +于是调用 `_resolve_hidden_state_chunks()`。对字符串 ref: + +```python +if ref.startswith("speco:"): + full_chunk = get_sample(ref)["hidden"] + full_chunk = _densify_tq_tensor(full_chunk) +``` + +然后按 sample metadata 取行: + +```python +sample_hidden = full_chunk[ + chunk_start : chunk_start + chunk_length +] +``` + +同一个 chunk 被多个 sample 复用,所以使用: + +```python +self._tq_chunk_cache[ref] = full_chunk +``` + +保证一次 `collect_rollout_features()` 中同一个 TQ key 只 get 一次。 + +得到逐样本 hidden 后,后续 `_store_rollout_sample → collect_online_data → train_drafter` 与 SGLang 路径相同。 + +### 3.8 PR #48 数据生命周期和清理 + +PR #48 的 TQ row 生命周期是: + +```text +TaskRunner tq.init(config) +→ Producer kv_put +→ key 经 Ray 控制路径传递 +→ 一个或多个 drafter TP/SP rank kv_batch_get +→ online drafter 收集/训练继续执行 +→ 整个 trainer.fit() 结束 +→ TaskRunner finally 调用 tq.close() +``` + +当前没有: + +```python +tq.kv_clear(key) +``` + +原因是同一 sample/key 可能被一个 drafter replica 的多个 rank 读取。若第一个 rank get 后删除,其余 rank 可能 get 失败。 + +因此 PR #48 采用任务级生命周期,而不是样本级确认和回收。这简化了并发正确性,但意味着长任务中的 TQ storage 占用可能持续增长;代码注释也把“leader 在 barrier 后精细 clear”留作后续工作。 + +### 3.9 PR #48 开启与关闭时的行为差异 + +`configure_transfer_queue()` 返回: + +```python +enabled_in_config and transfer_queue_importable +``` + +关闭时: + +```text +SGLang drafter_sample 继续内联 hidden_states +old-logprob 继续 ray.put(hidden_chunk) +Consumer 继续 ray.get/ref resolve +``` + +开启时: + +```text +SGLang hidden fields → TQ,sample 只带 key +old-logprob hidden chunk → TQ,ref 变成字符串 key +Consumer 根据 key 类型走 TQ get +``` + +如果配置开启但 `transfer_queue` 包未安装,bridge 会记录 warning,并让 `is_transfer_queue_enabled()` 返回 false,保留旧 Ray 路径。若已经进入 `put_sample/get_sample` 却发生 TQ 错误,则抛异常,不静默丢 hidden data。 + +## 第二部分:基于 PR #48 的 standalone drafter training 适配 + +从这一部分开始才讨论 `verl-SpeCo-ls` 的独立训练。下面的 `kv_list`、rank 0 选 global keys、逐 step `kv_clear`、独立 vLLM Producer 都是需要在本项目新增的逻辑,不是 PR #48 已有逻辑。 + +### 当前 standalone 基线 + +当前独立训练是: + +```text +draft_train_launcher +→ torch.distributed.run +→ 每个 rank 创建 DraftFeatureDataLoader +→ 每个 rank 在训练循环内同步 TargetFeatureReplayer.materialize() +→ vLLM/file hidden payload +→ prepare_training_batch_from_samples() +→ training_step_from_batch() +``` + +新方案要把 `TargetFeatureReplayer.materialize()` 从训练 rank 的同步取数路径移到独立 Producer,同时保持后两步训练接口不变。 + +## 4. TQ metadata 到底记录什么 + +### 4.1 Partition + +一次训练运行使用一个独立 partition: + +```python +partition_id = f"speco:{run_id}:dspark_train" +``` + +partition 用来隔离: + +- 不同训练 run; +- train 和 validation; +- 不同 target checkpoint 生成的 hidden states。 + +不能让两个 target 模型共用同一 partition,否则训练侧可能消费错误的 hidden states。 + +### 4.2 Sample key + +每条输入样本使用稳定 key: + +```python +sample_key = sha256( + dataset_id + + row_id + + prompt_token_ids + + response_token_ids + + tokenizer_fingerprint + + target_model_fingerprint + + target_layer_ids + + hidden_states_layout +).hexdigest() +``` + +稳定 key 用于: + +- vLLM HTTP 请求重试时不生成不同对象; +- Producer 重启后识别相同样本; +- 检查 hidden states 是否属于正确模型和正确层; +- TQ/Mooncake 清理时准确定位对象。 + +### 4.3 Fields 与 READY 约定 + +每个样本包含固定字段: + +```python +{ + "input_ids": int64[seq], + "loss_mask": float32[seq], + "position_ids": int64[seq], + "hidden_states": bf16[seq, aux_hidden_dim or aux_hidden_dim+hidden_size], +} +``` + +这里必须保持当前 `DraftFeatureSample` 的契约:`TargetFeatureReplayer._feature_from_vllm_payload()` 会把多个 aux layers flatten;当 layout 是 `dflash_aux_plus_last` 时,还会把 final hidden 拼到同一个 `hidden_states` tensor 尾部。`DSparkTrainerBackend.preprocess_individual_items()` 再根据 metadata 中的 `hidden_states_layout` 将 final hidden 切出来,生成训练 batch 的 `target_last_hidden_states`。 + +因此 TQ 不需要新增一个独立 `target_last_hidden_states` field。开启 DSpark L1 时,要求: + +```text +metadata.hidden_states_layout = dflash_aux_plus_last +hidden_states.shape[-1] = num_context_layers * hidden_size + hidden_size +``` + +完整 TQ native metadata API 可以做字段级 ready 判定,但 PR #48 当前走的是高层 KV API:一次 `kv_put` 把一个样本的多个 tensor fields 一起写入。因此第一版采用更直接的约定: + +```python +required_fields = [ + "input_ids", + "loss_mask", + "position_ids", + "hidden_states", +] +``` + +Producer 先在内存中验证所有必需字段,再进行一次 `kv_put`。只有 `kv_put` 成功返回,key 才会出现在 `kv_list` 结果中,并带有: + +```python +tag={ + "status": "ready", + "run_id": run_id, + "sample_id": sample_key, +} +``` + +Consumer 只选择 `status=ready` 且 `run_id` 匹配的 key。不要先 put `input_ids`、再单独 put `hidden_states`,否则 Consumer 可能观察到半成品。 + +### 4.4 Tags + +tags 是轻量 metadata,不放大 tensor: + +```python +tags = { + "sample_id": sample_key, + "source_row": row_id, + "seq_len": seq_len, + "payload_bytes": payload_bytes, + "target_model_fp": target_model_fingerprint, + "target_layers": "8,16,24", + "hidden_layout": "dflash_aux_plus_last", + "producer_status": "success", +} +``` + +tags 可用于过滤、监控、背压统计和错误排查。它不能替代 tensor shape/dtype 校验。 + +### 4.5 Run ID,而不是先依赖 task_name + +PR #48 的 `kv_put/kv_batch_get` 路径没有使用 `task_name` 或 Sampler consumption history。第一版按它的已验证接口,在 partition 和 tags 中放 `run_id`: + +```python +partition_id = "speco_drafter_features" +tag = { + "run_id": run_id, + "status": "ready", +} +``` + +不同 run 最好直接使用不同 partition: + +```python +partition_id = f"speco_drafter_features_{run_id}" +``` + +这样 job 重启和清理更简单。`task_name=dspark_train` 留到后续迁移 TQ native metadata/Sampler 时再使用。 + +### 4.6 standalone 中一条样本的完整对象形态 + +standalone 没有 PR #48 的轻量 `drafter_sample → Ray driver` 路径,因此需要让 TQ tag 承担“如何找到和解释 payload”的 metadata 作用。 + +#### Producer 读到的原始记录 + +```python +source_record = { + "dataset_id": "math-train", + "row_id": 12345, + "prompt": "...", + "response": "已经提前生成的 response", +} +``` + +#### Token replay 样本 + +分词和对齐后: + +```python +replay_sample = DraftReplaySample( + input_ids=int64[full_seq], + loss_mask=float32[full_seq], + position_ids=int64[full_seq], + feature_positions=int64[feature_rows], + draft_position_ids=int64[feature_rows], + metadata={ + "dataset_id": "math-train", + "row_id": 12345, + }, +) +``` + +这里的 `full_seq` 是 prompt 与预生成 response 拼接后的长度;`feature_positions` 指明哪些 token 位置最终进入 drafter 监督样本。 + +#### vLLM 返回的原始 hidden payload + +当前文件协议要求 safetensors 至少包含: + +```python +vllm_payload = { + "token_ids": int64[prefill_rows], + "hidden_states": bf16[prefill_rows, returned_layers, hidden_size], +} +``` + +这还不能直接给 DSpark。Producer 应复用当前: + +```python +TargetFeatureReplayer._feature_from_vllm_payload(...) +``` + +完成 token 校验、position 对齐、选层和 flatten。 + +#### Producer 最终得到的 DraftFeatureSample + +```python +feature = DraftFeatureSample( + algorithm="DSpark", + input_ids=int64[feature_rows], + loss_mask=float32[feature_rows], + position_ids=int64[feature_rows], + hidden_states=bf16[feature_rows, feature_hidden_dim], + metadata={ + "hidden_states_layout": "dflash_aux_plus_last", + "target_layer_ids": [8, 16, 24], + "target_model_path": "...", + "target_config_fingerprint": "...", + "feature_start": 128, + "feature_end": 640, + "sequence_length": 512, + }, +) +``` + +若有 3 个 aux layer、target hidden size 为 4096,并包含 final hidden: + +```text +feature_hidden_dim = 3 * 4096 + 4096 = 16384 +hidden_states.shape = [feature_rows, 16384] +``` + +前 `12288` 维是 aux/context hidden,最后 `4096` 维是 DSpark L1 所需 final hidden。现有 DSpark backend 会根据 `dflash_aux_plus_last` 自动拆分。 + +#### 写入 TQ 的 data fields + +第一版建议一个 key 对应一个已经规范化完成的 `DraftFeatureSample`: + +```python +tq_fields = { + "input_ids": feature.input_ids.cpu(), + "loss_mask": feature.loss_mask.cpu(), + "position_ids": feature.position_ids.cpu(), + "hidden_states": feature.hidden_states.cpu(), +} +``` + +这里只放 tensor,因为 PR #48 的 `put_sample()` 会过滤非 tensor: + +```python +payload = { + key: value + for key, value in tensor_dict.items() + if torch.is_tensor(value) +} +``` + +#### 写入 TQ 的 tag metadata + +```python +tq_tag = { + "run_id": "run-20260818-001", + "status": "ready", + "sample_id": sample_key, + "sequence_no": 12345, + "algorithm": "DSpark", + "hidden_states_layout": "dflash_aux_plus_last", + "target_model_fingerprint": "sha256:...", + "target_layer_ids": "8,16,24", + "feature_rows": 512, + "hidden_dim": 16384, + "payload_bytes": 16777216, +} +``` + +tag 中只放 TQ 版本支持序列化的小标量/字符串。列表等复杂对象可以编码为稳定字符串或 JSON。`payload_bytes` 用于背压统计。 + +#### TQ 中逻辑上保存的 row + +```text +partition: speco_drafter_features_run-20260818-001 +key: 86a4...ef2 + +fields: + input_ids → int64[512] + loss_mask → float32[512] + position_ids → int64[512] + hidden_states → bf16[512, 16384] + +tag: + status → ready + sequence_no → 12345 + hidden_layout → dflash_aux_plus_last + target_model_fp → sha256:... +``` + +#### Consumer 恢复出的对象 + +rank 0 通过 `kv_list` 同时取得 key 和 tag;每个 rank 用 local keys 调用 `kv_batch_get` 取得 fields,然后组合: + +```python +feature = DraftFeatureSample( + algorithm=tag["algorithm"], + input_ids=densify(fields["input_ids"]).reshape(-1), + loss_mask=densify(fields["loss_mask"]).reshape(-1), + position_ids=densify(fields["position_ids"]).reshape(-1), + hidden_states=densify(fields["hidden_states"]), + metadata={ + "hidden_states_layout": tag["hidden_states_layout"], + "target_model_fingerprint": tag["target_model_fingerprint"], + }, +) + +feature.validate(strict=True) +``` + +这样传给: + +```python +trainer.prepare_training_batch_from_samples([feature, ...]) +``` + +的数据结构,与现有文件 feature store 读出的 `DraftFeatureSample` 一致。也就是说,TQ 替换的是存取介质和调度方式,不改变 DSpark backend 的样本契约。 + +## 5. 新的整体架构 + +```text + 小 metadata + ┌──────────────────────────┐ + │ TransferQueueController │ + │ KV metadata / key / tags │ + │ partition / storage map │ + └────────────┬─────────────┘ + │ +JSONL/token replay │ + │ │ + ▼ │ +Feature Producer │ + ├─ tokenizer/window │ + ├─ asyncio bounded concurrency │ + ├─ vLLM endpoint pool │ + ├─ validate/pack │ + └─ TQ put ─────────────────────┤ + ▼ + TQ Mooncake backend + hidden-state tensors + │ + ┌───────────────────┼───────────────────┐ + ▼ ▼ ▼ + DSpark rank 0 DSpark rank 1 DSpark rank N + TQ get TQ get TQ get + └───────────────────┼───────────────────┘ + ▼ + synchronized optimizer step + │ + ▼ + TQ clear after success +``` + +大 tensor 的路径是: + +```text +vLLM/Producer memory → TQ Mooncake backend → each training rank +``` + +不会走: + +```text +Mooncake → rank 0 → rank 1/2/3 +``` + +rank 0 最多只广播 `kv_list` 得到的 key 字符串列表;tags 只在选 batch 时由 rank 0 使用。 + +### 5.1 standalone 每一步为什么能实现推理和训练异步 + +#### 步骤 A:Producer 独立推进输入 cursor + +Producer 自己维护: + +```python +reader_cursor = 12346 +``` + +它不等待 Trainer 请求样本。只要 TQ ready bytes 没超过背压上限,就继续读取文件并创建 vLLM task。 + +效果是 Producer 的执行进度与 `optimizer_step` 解耦: + +```text +Producer sequence_no: 1200,1201,1202,... +Trainer optimizer_step: 87 +``` + +两者通过 TQ 中的 ready rows 衔接,不互相直接调用。 + +#### 步骤 B:并发 vLLM task 完成顺序可以乱序 + +例如 Producer 同时提交: + +```text +sequence_no 100 → endpoint 0 +sequence_no 101 → endpoint 1 +sequence_no 102 → endpoint 0 +``` + +完成顺序可能是: + +```text +101 → 100 → 102 +``` + +每个 task 完成后独立执行 `kv_put`,所以慢请求不会阻塞已经完成的请求写入。tag 中的 `sequence_no` 保留原数据顺序。 + +#### 步骤 C:`kv_put` 成功是 READY 可见性的边界 + +Producer 在调用前已经得到完整 `DraftFeatureSample`。一次 `kv_put` 写入该 sample 的全部 tensor fields,并在 tag 中标记 `status=ready`。 + +因此 Consumer 的判断规则是: + +```text +kv_list 能列出该 key +且 tag.run_id 匹配 +且 tag.status == ready +→ 可以尝试 kv_batch_get +``` + +Consumer 仍需对取回 fields 做完整性校验;tag 是调度 metadata,不是正确性证明。 + +#### 步骤 D:rank 0 只负责选 key + +rank 0 执行: + +```python +entries = list_ready_keys() +selected = sorted(entries, key=sequence_no)[:global_batch_size] +``` + +这一步处理的数据只是: + +```python +[ + {"key": "k100", "sequence_no": 100, ...}, + {"key": "k101", "sequence_no": 101, ...}, +] +``` + +不包含 `[seq, hidden_dim]` hidden tensor,所以 rank 0 不成为大数据中转瓶颈。 + +#### 步骤 E:广播保证所有 rank 对同一个 global step 达成一致 + +所有 rank 调用同一次: + +```python +dist.broadcast_object_list(holder, src=0) +``` + +广播结束后,每个 rank 看到完全相同的 global key list。然后按确定性区间切分: + +```text +rank 0: keys[0:per_rank] +rank 1: keys[per_rank:2*per_rank] +... +``` + +这样不会出现 rank 0 训练 batch A、rank 1 训练 batch B 数量不同,或者某个 rank 没进入 backward 的情况。 + +#### 步骤 F:各 rank 直接读取 Mooncake 后端 + +每个 rank 执行: + +```python +tq.kv_batch_get(keys=local_keys, partition_id=partition_id) +``` + +TQ 根据 key 找到 storage backend 中的数据。使用 MooncakeStore 时,大 tensor 数据路径是 Mooncake → 本 rank;rank 0 不读取其他 rank 的 local payload。 + +因此: + +```text +控制面:rank 0 → broadcast small keys +数据面:Mooncake → each rank directly +``` + +#### 步骤 G:恢复现有 DraftFeatureSample 契约 + +每个 rank 将 TQ fields 与 tag 合并、densify、校验,得到 `list[DraftFeatureSample]`。从这里开始继续执行项目现有代码: + +```python +batch = trainer.prepare_training_batch_from_samples( + materialized_samples, + step=optimizer_step, +) + +ok = await trainer.training_step_from_batch( + batch, + optimizer_step, +) +``` + +所以 TQ 不进入 DSpark model/backend 内部,训练数学逻辑不变。 + +#### 步骤 H:全 rank 成功以后才能清理 + +每个 rank 的 `ok` 通过现有 `_all_ranks_true()` 聚合: + +```text +rank 0 ok = true +rank 1 ok = true +rank 2 ok = true +rank 3 ok = true +→ global_ok = true +``` + +只有此时 rank 0 执行: + +```python +tq.kv_clear(keys=global_keys, partition_id=partition_id) +``` + +这样保证被删除的数据已经参与完成的 optimizer step。若任何 rank get/OOM/backward 失败,不执行 clear,便于作业失败后的诊断或恢复。 + +#### 步骤 I:异步重叠如何形成 + +时间线上: + +```text +时间 ─────────────────────────────────────────▶ + +Producer: vLLM(batch N+1) ─ put ─ vLLM(batch N+2) ─ put +Trainer: get(batch N) ─ train(batch N) ─ get/train(batch N+1) +``` + +Producer 和 Trainer 是不同进程,互相不调用;TQ ready rows 是缓冲区。因此 vLLM prefill、网络传输和 DSpark GPU 训练可以重叠。背压只在缓冲区达到容量上限时暂停 Producer。 + +## 6. Producer:读取预生成 response 并并行请求 vLLM + +### 6.1 输入处理 + +Producer 从现有 JSONL/token replay 数据源读取: + +```python +sample = { + "row_id": "12345", + "prompt": "...", + "response": "提前生成好的文本", +} +``` + +构造: + +```python +prompt_ids = tokenizer.encode(sample["prompt"]) +response_ids = tokenizer.encode(sample["response"]) +input_ids = prompt_ids + response_ids +``` + +同时产生: + +```python +loss_mask +position_ids +feature_positions +sample_key +``` + +### 6.2 有界并发 + +不能按样本串行请求: + +```python +for sample in samples: + result = request_vllm(sample) +``` + +改成: + +```python +async def run_producer(samples): + semaphore = asyncio.Semaphore(max_inflight_requests) + + async def run_one(sample): + async with semaphore: + result = await vllm_pool.prefill(sample) + feature = validate_and_pack(sample, result) + await tq_transport.put(feature) + + async with asyncio.TaskGroup() as group: + for sample in samples: + group.create_task(run_one(sample)) +``` + +`max_inflight_requests` 是 Producer 同时未完成的请求数量,不是 vLLM batch size。vLLM 服务端仍会对同时到达的请求做自己的 continuous batching。 + +### 6.3 多 endpoint + +多个 endpoint 例如: + +```yaml +vllm_endpoints: + - http://node0:8000/v1 + - http://node1:8000/v1 + - http://node2:8000/v1 +``` + +调度器维护每个 endpoint 的 inflight 数: + +```python +endpoint = min( + endpoints, + key=lambda item: item.inflight, +) +``` + +请求前 `inflight += 1`,在 `finally` 中 `inflight -= 1`。失败只重试对应样本,不阻塞全部 Producer。 + +### 6.4 当前 vLLM 文件桥接 + +当前客户端协议期望: + +```python +response.kv_transfer_params["hidden_states_path"] +``` + +所以第一阶段仍然是: + +```text +vLLM 写临时 safetensors +→ Producer load_file +→ 校验 token_ids/hidden_states +→ TQ put 到 Mooncake backend +→ TQ put 成功后删除临时文件 +``` + +删除必须发生在 TQ put 成功之后: + +```python +path = request_vllm_hidden_file(sample) +try: + feature = load_and_validate(path) + await tq_transport.put(feature) +finally: + if put_succeeded: + Path(path).unlink(missing_ok=True) +``` + +### 6.5 目标版本:vLLM 直接写 TQ/Mooncake + +目标响应可改成: + +```json +{ + "kv_transfer_params": { + "backend": "transfer_queue", + "partition_id": "speco:run-1:dspark_train", + "sample_key": "abc123" + } +} +``` + +服务端顺序必须是: + +```text +prefill +→ 捕获指定层 hidden states +→ TQ/Mooncake put 完成 +→ 返回 HTTP success 和 sample key +``` + +这需要定制 vLLM exporter;当前 `verl-SpeCo-ls` 中没有服务端 writer 实现。 + +## 7. 按 PR #48 扩展 TQ bridge + +不要重新发明一套 transport。将 PR #48 的 bridge 设计移植到本项目并增加 standalone 所需方法: + +```python +class StandaloneTQTransport: + def put_sample(self, key, tensor_dict, tag): ... + def list_ready_keys(self, run_id): ... + def get_samples(self, keys, fields=None): ... + def clear_samples(self, keys): ... + def put_control(self, key, tag): ... + def close(self): ... +``` + +写入延续 PR #48 的真实形式: + +```python +tq.kv_put( + key=key, + partition_id=partition_id, + fields={ + "input_ids": feature.input_ids.cpu(), + "loss_mask": feature.loss_mask.cpu(), + "position_ids": feature.position_ids.cpu(), + "hidden_states": feature.hidden_states.cpu(), + }, + tag={ + "run_id": run_id, + "status": "ready", + "sequence_no": sequence_no, + "payload_bytes": payload_bytes, + }, +) +``` + +批量读取延续 PR #48 的 `kv_batch_get`: + +```python +result = tq.kv_batch_get( + keys=keys, + partition_id=partition_id, + fields=required_fields, # 0.1.7 是否支持该参数需实机确认 +) +``` + +新增发现和清理: + +```python +items = tq.kv_list(partition_id=partition_id) +tq.kv_clear(keys=keys, partition_id=partition_id) +``` + +这里的 `kv_list/kv_clear` 参数名必须根据锁定的 TQ 版本验证。当前参考仓库只实际调用了 `kv_put/kv_batch_get`,没有为这两个接口提供运行证据。 + +读取结果继续复用 PR #48 的两个适配函数: + +```python +value = _extract_value(result, key) +row = _tensordict_to_dict(value) +row["hidden_states"] = _densify_tq_tensor(row["hidden_states"]) +``` + +## 8. DSpark 多 rank 如何消费 + +### 8.1 第一版:rank 0 用 kv_list 发现 READY keys + +PR #48 中 key 由 Ray driver 传给 drafter worker;standalone 没有这条控制路径,所以 rank 0 需要主动列出 key: + +```python +rank = dist.get_rank() +world_size = dist.get_world_size() +global_batch_size = batch_size_per_gpu * world_size + +if rank == 0: + entries = tq_transport.list_ready_keys(run_id=run_id) + entries.sort(key=lambda x: (x.tag["sequence_no"], x.key)) + selected_keys = [x.key for x in entries[:global_batch_size]] +else: + selected_keys = None + +holder = [selected_keys] +dist.broadcast_object_list(holder, src=0) +selected_keys = holder[0] +``` + +`sequence_no` 是输入文件顺序。它让多个 Producer 并发完成顺序不同的情况下,Trainer 仍能确定性地组成 batch。 + +rank 0 广播的是字符串 key 列表,不是 hidden-state tensor。 + +### 8.2 各 rank 切自己的 keys + +例如 global batch keys: + +```text +[s0, s1, s2, s3, s4, s5, s6, s7] +``` + +world size 为 4、每卡 batch size 为 2: + +```text +rank 0 → [s0, s1] +rank 1 → [s2, s3] +rank 2 → [s4, s5] +rank 3 → [s6, s7] +``` + +代码: + +```python +def shard_keys(keys, rank, world_size): + assert len(keys) % world_size == 0 + per_rank = len(keys) // world_size + start = rank * per_rank + end = start + per_rank + return keys[start:end] +``` + +### 8.3 每个 rank 并行 get + +所有进程执行: + +```python +local_keys = shard_keys( + selected_keys, + rank=rank, + world_size=world_size, +) + +local_payloads = tq_transport.get_samples(local_keys) +``` + +数据路径: + +```text +rank 0 ← Mooncake(s0,s1) +rank 1 ← Mooncake(s2,s3) +rank 2 ← Mooncake(s4,s5) +rank 3 ← Mooncake(s6,s7) +``` + +不是 rank 0 get 全部后再 scatter。 + +### 8.4 转成当前训练格式 + +TQ 返回的数据需要先按 PR #48 的规则解包、densify,再转换成现有 `DraftFeatureSample`: + +```python +def tq_row_to_feature(row, tag): + return DraftFeatureSample( + algorithm="DSpark", + input_ids=row["input_ids"], + loss_mask=row["loss_mask"], + position_ids=row["position_ids"], + hidden_states=row["hidden_states"], + metadata={ + "hidden_states_layout": tag["hidden_states_layout"], + **row.get("metadata", {}), + }, + ) +``` + +`hidden_states_layout` 不能只放在无法恢复的临时 Python 对象里;应随 TQ tag 保存。rank 0 从 `kv_list` 取得 key/tag 后,需要把每个 local key 对应的 tag 一起交给本 rank 的转换逻辑。Consumer 必须把 layout 放回 `DraftFeatureSample.metadata`。DSpark L1 所需的 `target_last_hidden_states` 随后由现有 backend 从拼接的 `hidden_states` 中切出,transport 层不计算 loss,也不拆 layout。 + +## 9. 修改当前训练循环 + +在 `run_standalone_draft_training()` 中增加数据源分支: + +```python +feature_store_type = str(feature_store_cfg.get("type", "torch_shard")) + +if feature_store_type == "transfer_queue": + tq_stream = build_transfer_queue_stream( + config=config, + rank=rank, + world_size=world_size, + ) + store = None + loader = None + feature_replayer = None +else: + store = build_feature_store_from_config( + feature_store_cfg, + read_only=True, + ) + loader = DraftFeatureDataLoader(...) +``` + +流式训练循环: + +```python +while successful_steps < max_steps: + global_keys, materialized_samples = tq_stream.next_local_batch() + + batch = trainer.prepare_training_batch_from_samples( + materialized_samples, + step=optimizer_step, + ) + + has_batch = batch is not None + if not _all_ranks_true(has_batch, trainer.runtime_device): + raise RuntimeError("at least one rank failed to fetch its TQ batch") + + ok = await trainer.training_step_from_batch( + batch, + optimizer_step, + ) + + if not _all_ranks_true(ok, trainer.runtime_device): + raise RuntimeError("DSpark step failed on at least one rank") + + dist.barrier() + if rank == 0: + tq_stream.clear_global_batch(global_keys) + dist.barrier() +``` + +TQ 接在数据输入层,不放进 `DSparkTrainerBackend`。Backend 只负责模型、forward、loss、backward 和 optimizer。 + +## 10. READY key、inflight key 和训练提交 + +### 10.1 Ready + +在本方案中 ready 表示: + +> `dspark_train` 所需的所有 fields 已经成功写入 TQ storage,训练侧可以读取。 + +第一版不是通过 PR #48 尚未验证的 Sampler 自动判定,而是 Producer 完整校验 payload 后一次 `kv_put`,并设置 `tag.status=ready`。rank 0 的 `kv_list` 只选择这些 key。 + +### 10.2 Inflight key + +rank 0 选出一个 global batch 后,在本地保存: + +```python +inflight_global_keys = selected_keys +``` + +其他 step 不应再次选中这批 key。因为 standalone 只有一个同步 DSpark job,第一版由 rank 0 进程内 `set` 排除 inflight key 即可: + +```python +ready = [x for x in listed if x.key not in inflight_keys] +``` + +若以后允许多个独立 Trainer job 同时消费同一 partition,进程内 set 就不够,届时必须使用 TQ Controller/Sampler 的原子消费分配能力。 + +### 10.3 Optimizer committed + +optimizer committed 表示所有 DSpark rank 已经完成: + +```text +forward → backward → gradient synchronization → optimizer.step +``` + +它比 `kv_list` 发现 key、甚至 `kv_batch_get` 取出 tensor 都更晚。 + +第一版推荐简单语义: + +```text +TQ 负责 key/tag 和 tensor 传输 +rank 0 负责单 Trainer job 的 batch 选择和 inflight set +训练失败 → 整个作业 fail-fast +训练成功 → kv_clear payload,并从 inflight set 移除 +恢复 → 从最近 checkpoint + 输入 cursor 重新启动 +``` + +这样不需要重新实现 lease/ack 状态机,也没有夸大 PR #48 当前尚未使用的 Sampler 能力。 + +## 11. 为什么训练完一个 step 才清理 + +不能在 `kv_batch_get()` 后立即 clear: + +```text +get 成功 +→ clear +→ forward OOM +→ 数据已不存在,无法重试 +``` + +正确顺序: + +```text +rank 0..N get +→ 所有 rank 确认 batch 有效 +→ training_step_from_batch +→ _all_ranks_true(ok) +→ rank 0 kv_clear global keys +``` + +当前 `_all_ranks_true()` 已经是项目现有的跨 rank 同步工具,可以继续复用。 + +## 12. 背压 + +背压表示 Producer 生成速度高于 Trainer 消费速度时,Producer 必须暂停继续提交,以免 Mooncake 内存无限增长。 + +建议限制: + +```yaml +max_vllm_inflight_requests: 32 +max_pending_put_bytes: 8589934592 +max_tq_ready_samples: 256 +max_tq_ready_bytes: 68719476736 +``` + +Producer 在 tags 中写: + +```python +{"payload_bytes": payload_bytes} +``` + +周期性通过 TQ list/metadata 统计当前 partition 尚未消费的数据量: + +```python +while ready_bytes >= max_tq_ready_bytes: + await asyncio.sleep(backpressure_poll_interval) +``` + +如果锁定的 TQ 版本提供容量/ready 统计接口,应直接使用,避免全量扫描 keys。 + +## 13. Stable ID、幂等和孤儿数据 + +### 13.1 幂等 + +幂等表示同一操作重复执行,最终逻辑结果仍只有一份。 + +Producer 对同一样本重试时必须使用相同 `sample_key`: + +```python +await tq.put(key="abc123", ...) +await tq.put(key="abc123", ...) +``` + +不能每次生成随机 key: + +```text +abc123-retry-1 +abc123-retry-2 +``` + +否则一个输入可能训练多次并持续占用 Mooncake。 + +### 13.2 孤儿数据 + +孤儿数据表示 tensor 已经写进 storage,但由于 Producer 崩溃或 metadata 更新失败,没有进入正常消费路径。 + +使用 TQ 后,metadata 和 storage 由同一套系统管理,可以减少“Mooncake有对象、自研 Coordinator 没记录”的双系统窗口,但仍要配置 TTL/partition cleanup: + +```text +训练正常结束 → clear partition +训练异常退出 → 下次启动检查旧 partition +超过 TTL → 清理未消费数据 +``` + +## 14. EOS 和 drop-last + +EOS 表示 Producer 已经读完输入,并且所有 vLLM/TQ put 都已完成。 + +TQ 中需要一种结束条件,具体使用 Controller API、特殊 metadata 或单独的运行状态记录取决于锁定版本。不能把“当前暂时没有 ready sample”当成 EOS,因为 Producer 可能仍在请求 vLLM。 + +最后不足一个 global batch 时: + +```python +global_batch_size = batch_size_per_gpu * world_size +``` + +第一版使用 `drop_last=true`,避免某些 rank 有数据、某些 rank 没数据导致分布式训练不同步。 + +结束条件: + +```text +producer_done == true +and ready_samples < global_batch_size +and inflight_requests == 0 +and pending_puts == 0 +``` + +## 15. 双缓冲预取 + +训练 batch N 时,CPU 后台线程预取 batch N+1: + +```python +next_future = executor.submit(tq_stream.next_local_batch) + +current_batch = first_batch +while current_batch is not None: + next_batch = next_future.result() + next_future = executor.submit(tq_stream.next_local_batch) + + train(current_batch) + current_batch = next_batch +``` + +实际顺序应调整为避免等待 future 后才训练。推荐: + +```python +current = tq_stream.next_local_batch() + +while current is not None: + future = executor.submit(tq_stream.next_local_batch) + train_and_clear(current) + current = future.result() +``` + +第一版只预取一个 global batch,避免 Trainer 崩溃时大量数据已被采样但未训练。 + +如果 Mooncake/TQ Python get 是阻塞函数,使用专用 `ThreadPoolExecutor`,不要阻塞 Producer 的 asyncio event loop。 + +## 16. 建议代码结构 + +```text +verl_speco/ + trainer/ + tq_transport.py # TQ client、put/get/meta/clear 封装 + tq_feature_stream.py # kv_list、global keys、rank shard、decode/prefetch + feature_producer.py # JSONL → 并发 vLLM → TQ + draft_training_loop.py # 增加 transfer_queue 数据源分支 + target_feature_replay.py # 复用 vLLM payload 校验/转换逻辑 +``` + +不要新增: + +```text +coordinator.py +coordinator_client.py +``` + +建议抽象: + +```python +class StreamingFeatureSource(Protocol): + def next_local_batch(self) -> tuple[list[str], list[DraftFeatureSample]]: ... + def clear_global_batch(self, keys: list[str]) -> None: ... + def close(self) -> None: ... +``` + +这样训练循环不依赖 TQ 的具体类型。 + +## 17. 配置草案 + +```yaml +actor_rollout_ref: + rollout: + drafter: + training: + mode: offline + backend: dspark + batch_size_per_gpu: 2 + max_steps: 1000 + + feature_store: + type: transfer_queue + partition_id: speco_drafter_features_${run_id} + drop_last: true + prefetch_steps: 1 + + transfer_queue: + # 与 PR #48 的配置层级和 init 方式保持一致。 + enable: true + package_version: 0.1.8 # 最终以实测版本为准 + backend: + storage_backend: MooncakeStore + MooncakeStore: + auto_init: false + metadata_server: localhost:50123 + master_server_address: localhost:50124 + local_hostname: localhost + protocol: tcp + global_segment_size: 4294967296 + local_buffer_size: 1073741824 + device_name: "" + + required_fields: + - input_ids + - loss_mask + - position_ids + - hidden_states + + producer: + input_path: /path/to/generated_responses.jsonl + vllm_endpoints: + - http://node0:8000/v1 + - http://node1:8000/v1 + max_inflight_requests: 32 + max_pending_put_bytes: 8589934592 + max_ready_samples: 256 + max_ready_bytes: 68719476736 +``` + +当前 examples 中的: + +```bash +transfer_queue.enable=False +``` + +属于 verl RL 主入口配置,当前 standalone `draft_train_launcher` 不读取它。不能只改成 `True`;必须实现上述 `feature_store.type=transfer_queue` 分支。 + +## 18. 启动顺序 + +逻辑顺序: + +```text +1. 启动 Mooncake metadata/master 服务; +2. 启动 standalone TQ owner 进程,调用一次 `tq.init(tq_config)` 创建/连接 TQ Controller 和 MooncakeStore backend; +3. 启动一个或多个定制 vLLM server +4. 启动 Feature Producer +5. Producer 和所有 Trainer rank 调用无参 `tq.init()`,连接 owner 创建的同一套 TQ; +6. 启动 verl_speco.draft_train_launcher +7. torchrun 启动所有 DSpark rank +8. 各 rank 连接 TQ +9. rank 0 用 `kv_list` 选择 global keys,各 rank 并行 `kv_batch_get`/train; +10. 输入耗尽后 Producer 发布 done 状态 +11. Trainer drain 完整 global batches 后退出 +12. 清理 partition,停止 Producer、vLLM、TQ、Mooncake +``` + +PR #48 的 `init_transfer_queue()` 在 Ray `SpecoTaskRunner` 中调用 `tq.init(config)`,worker 的无参 `tq.init()`依靠同一个 Ray 集群发现 named Controller;这部分不能原样复制到 standalone torchrun。 + +本项目要求一开始不使用 Ray,因此 Phase 0 必须先证明锁定的 TQ 版本支持独立 owner/controller 进程,以及 Producer/torchrun rank 如何获得连接信息。若 TQ 0.1.7/0.1.8 实际只能通过 Ray named actor 完成发现,那么有两个选择: + +1. 接受仅用 Ray 承载 TQ 控制面的最小方案; +2. 给 TQ 增加或使用其已有的独立 ZMQ/server-info bootstrap。 + +在这项验证完成前,文档不能声称 PR #48 已经提供“无 Ray TQ standalone 启动”。 + +## 19. 故障处理 + +### vLLM 请求失败 + +- 对单个 sample 按稳定 key 重试; +- 指数退避; +- 超过次数记录失败,并根据配置 fail-fast 或跳过; +- 不写不完整 TQ fields。 + +### vLLM 文件读取成功,但 TQ put 失败 + +- 暂时保留临时文件; +- 重试 TQ put; +- put 成功后再删除; +- 不把样本视为 ready。 + +### 某个训练 rank get 失败 + +- 该 rank 报告 `local_ok=false`; +- `_all_ranks_true()` 使全部 rank 得到一致失败结果; +- 第一版整个训练 fail-fast; +- 不 clear global batch。 + +### OOM/optimizer step 失败 + +- 不 clear; +- 所有 rank 一致退出; +- 从最近训练 checkpoint 恢复; +- 根据 TQ 消费提交语义决定是否重放当前 batch。 + +### clear 失败 + +- optimizer 已成功,不能再次训练这批; +- 将 batch keys 写入本地小型 `gc_pending` 日志; +- 后台重试 clear; +- checkpoint 保存最近 committed sample IDs,避免恢复时重复消费。 + +## 20. 观测指标 + +Producer: + +```text +producer/vllm_inflight +producer/vllm_requests_per_sec +producer/vllm_prefill_tokens_per_sec +producer/vllm_p50_latency +producer/vllm_p95_latency +producer/tq_put_bytes_per_sec +producer/tq_put_failures +producer/pending_put_bytes +``` + +TQ/Mooncake: + +```text +tq/ready_samples +tq/ready_bytes +tq/consumed_samples +tq/storage_bytes +tq/clear_failures +mooncake/put_bandwidth +mooncake/get_bandwidth +``` + +Trainer: + +```text +trainer/tq_wait_seconds +trainer/tq_get_seconds +trainer/tq_get_bytes_per_sec +trainer/decode_seconds +trainer/h2d_seconds +trainer/step_seconds +trainer/data_stall_ratio +trainer/successful_steps +``` + +## 21. 实施阶段 + +### Phase 0:锁定依赖和契约 + +- 从 PR #48 的 `TransferQueue==0.1.7` 起验证,同时对比当前 verl 文档使用的 0.1.8; +- 实测 `tq.init(config)`、worker `tq.init()`、`kv_put`、`kv_batch_get`、`kv_list`、`kv_clear`; +- 验证该版本是否支持无 Ray owner/controller 部署以及连接信息传递; +- 先用 PR #48 已配置的 `SimpleStorage` 做最小闭环; +- 再把 backend 换成 `MooncakeStore`,验证 tcp,最后再验证 rdma; +- 固定 fields、tags、partition 和 `run_id/sequence_no/status`; +- 写 fake TQ 单元测试。 + +### Phase 1:文件桥接 + TQ KV 模式 + +- 新增独立 Producer; +- 32 个有界并发 vLLM 请求; +- 读取 vLLM 临时 safetensors; +- TQ put 成功后删除文件; +- standalone trainer 由 rank 0 `kv_list` 并广播 global keys; +- 各 rank 并行 `kv_batch_get`; +- 复用 PR #48 的返回值解包、NestedTensor densify 和 per-step cache; +- optimizer 成功后 `kv_clear`。 + +验收:连续训练 1000 step,临时文件数量、TQ ready bytes 和 Mooncake占用均保持有界。 + +### Phase 2:双缓冲与多 endpoint + +- 增加多 endpoint 最少 inflight 调度; +- 增加一个 global batch 预取; +- 动态背压; +- 注入单 rank get 失败,确认所有 rank 一致退出而非死锁。 + +### Phase 3:vLLM 直接写 TQ/Mooncake + +- 修改外部定制 vLLM exporter; +- 去掉 `hidden_states_path` 临时文件; +- HTTP 响应返回 partition/sample key; +- 验证 HTTP 重试的幂等性。 + +### Phase 4:可选升级到 TQ StreamingDataLoader + +- 在当前保守方案稳定后再引入 RankAwareSampler; +- 让每个 rank 自动取得 local micro-batch; +- 去掉 rank 0 手工 key-list 广播; +- 验证与 torchrun/DSpark 的 global step 对齐。 + +## 22. 最终推荐 + +针对当前 `verl-SpeCo-ls`,推荐的第一版不是 AngelSpec 的完整架构,也不是单独写 Coordinator,而是: + +```text +当前预生成 response 文件 +→ 独立 asyncio Producer +→ 并行访问多个 vLLM endpoint +→ 读取并校验临时 hidden-state 文件 +→ TransferQueue put +→ Mooncake storage backend +→ rank 0 kv_list 获取 READY global keys +→ broadcast key list +→ 各 DSpark rank 并行 kv_batch_get +→ 现有 prepare_training_batch_from_samples() +→ 现有 training_step_from_batch() +→ 全 rank 成功 +→ TQ clear +``` + +这套方案保留当前 standalone DSpark 训练主体,只替换 `DraftFeatureDataLoader + TargetFeatureReplayer.materialize()` 所在的数据输入路径。它真正建立在 PR #48 已实现的 KV transport 之上,而不是假设 PR #48 已经实现了 standalone Sampler/StreamingDataLoader。 + +## 23. 参考 + +- verl TransferQueue: +- TransferQueue: +- Mooncake Store: +- 本地只读参考:`../verl-SpeCo/verl_speco/integration/transferqueue_bridge.py`(PR #48) +- 本地只读参考:`../verl-SpeCo/docs/transferqueue_integration_plan.md`(PR #48) diff --git a/docs/standalone_tq_consumer_implementation.md b/docs/standalone_tq_consumer_implementation.md new file mode 100644 index 00000000..456afd40 --- /dev/null +++ b/docs/standalone_tq_consumer_implementation.md @@ -0,0 +1,725 @@ +# 独立 DSpark 训练 TQ Consumer 实现说明 + +## 1. 文档范围和当前结论 + +本文只说明当前仓库中已经实现的独立训练 Consumer。这里的 Consumer 是由 `torchrun` 启动的 DSpark 草稿模型训练任务:它持续从 TransferQueue(下文简称 TQ)发现样本,各训练 rank 分别取得自己负责的 Tensor,复用原有 DSpark 训练逻辑完成一次 optimizer step,然后由 rank 0 删除这一整个 global batch 对应的 TQ 记录。 + +当前已完成的能力是: + +1. `feature_store.type=tq` 可以作为独立训练的数据源,不要求磁盘 `path`。 +2. 每个训练 rank 都连接同一个 Ray 集群、同一个 TQ Controller 和同一个 partition。 +3. 只有 rank 0 调用 `kv_list` 发现 ready key,并把 key/tag 分配给各 rank。 +4. key 和 tag 通过 `torch.distributed.broadcast_object_list` 传输;hidden states 等 Tensor 不经过该广播。 +5. 每个 rank 根据分配到的 key,直接调用 TQ `kv_batch_get` 获取自己的 Tensor。 +6. TQ Tensor 被解码成原训练代码已经认识的 `DraftFeatureSample`,然后复用 `DrafterBaseTrainer` 的 batch 构造、DSpark loss、反向传播和 optimizer step。 +7. 只有当所有 rank 都成功完成该 step 后,rank 0 才调用 `kv_clear` 删除整个 global batch。 +8. Producer 发布 EOS 后,如果剩余样本不足一个 global batch,当前第一版会丢弃并清理这部分尾样本,然后正常结束训练迭代。 + +本文不会把尚未实现的 Producer 写成现有能力。Producer 后续需要复用本文第 6 节所述的公共协议,调用 `encode_sample()` 生成 fields,再使用 bridge 写入相同 TQ。 + +## 2. 本次涉及的文件 + +### 2.1 本次新增的 Consumer 核心文件 + +| 文件 | 实现的组件 | 作用 | +|---|---|---| +| `verl_speco/trainer/tq_feature_store.py` | `TQFeatureStore`、`ReadyEntry`、`EosMetadata` | 将公共 TQ bridge 包装成 Consumer 数据访问层,负责连接、发现、批量读取、解码、删除和读取 EOS | +| `verl_speco/trainer/tq_sample_source.py` | `TQFeatureDataLoader`、`TQLocalBatch`、`build_assignments()` | 实现多 rank 流式取数:rank 0 发现样本并分配 key,各 rank 自己从 TQ 取 Tensor | +| `tools/run_dspark_tq_consumer.sh` | Consumer 测试启动工具 | 给出一套完整的 DSpark、offline、TQ 配置和 `torchrun` 启动方式 | +| `tests/unit/test_tq_consumer.py` | Consumer 单元测试 | 覆盖 store 构建、ready 过滤排序、协议解码、rank 分配、EOS、清理和非 rank 0 行为 | + +### 2.2 本次修改的既有文件 + +| 文件 | 修改内容 | 为什么要改 | +|---|---|---| +| `verl_speco/trainer/feature_store.py` | factory 新增 `type=tq` 分支 | 让既有独立训练入口能够像选择磁盘 feature store 一样选择流式 TQ 数据源 | +| `verl_speco/trainer/draft_training_loop.py` | 接入 TQ store/loader、跨 rank 连接检查、训练成功后清理 | 将流式取数接入原训练循环,同时保留原 DSpark trainer、loss、optimizer、metric 和 checkpoint 逻辑 | +| `verl_speco/draft_train_launcher.py` | 增加 TQ 启动参数的 fail-fast 检查 | 在启动多个 torchrun 子进程前检查 `enable`、Ray address 和 `run_id`,避免各 rank 启动后才失败 | +| `verl_speco/config/speco_base.yaml` | 标注 `feature_store.type=tq` 为无路径流式数据源 | 保留统一 Hydra 配置入口;TQ 的公共配置仍位于 sibling `training.transfer_queue` | +| `tools/tq_connection_smoke.py` | 将原连接 smoke 扩展为真实 Consumer 路径测试 | 验证 owner、真实 TQ、Consumer 读取、EOS、清理和仅关闭本地 client | +| `tests/unit/test_draft_train_launcher.py` | 增加 TQ 参数检查测试 | 验证必要配置缺失时 launcher 直接拒绝启动 | +| `tests/unit/test_draft_training_loop.py` | 增加连接和 clear 时序测试 | 验证只由 rank 0 清理、clear 失败会报告、连接失败会传播 | + +### 2.3 直接复用的公共基础 + +以下文件不是这次 Consumer 才创造的概念,但 Consumer 直接使用它们: + +| 文件 | 被复用的能力 | +|---|---| +| `verl_speco/integration/transferqueue_bridge.py` | 屏蔽 TQ 0.1.7 API 细节,提供连接、`kv_list`、`kv_batch_get`、`kv_clear` 和本地关闭接口 | +| `verl_speco/transport/drafter_sample_protocol.py` | 定义 key、tag、fields、metadata 格式,以及 `encode_sample()` / `decode_sample()` | +| `verl_speco/trainer/feature_store.py` | 复用 `DraftFeatureSample`,使 TQ 数据进入训练侧后与磁盘 feature sample 类型一致 | +| `verl_speco/trainer/base_trainer.py` 及既有 backend | 复用 `DrafterBaseTrainer.prepare_training_batch_from_samples()` 和 `training_step_from_batch()` 等训练实现 | + +## 3. 运行时角色 + +### 3.1 TQ Owner + +TQ Owner 是单独的普通 Python 进程。它连接指定 Ray 集群,并以带配置的 `tq.init(config)` 创建任务级 named Controller 和 storage actors。Owner 持有全局 TQ 生命周期;Consumer 结束时不能关闭它。 + +Owner 不是训练 rank,也不执行 DSpark 模型。它的主要作用是让 Producer 和 Consumer 能通过同一个 Ray actor registry 找到同一个 TQ Controller。 + +### 3.2 Producer + +Producer 是后续需要实现的独立推理进程。它应并行调用 vLLM hidden-state 接口,构造一条条 `DraftFeatureSample` 和 `SampleMetadata`,再写入 TQ。 + +Producer 与 Consumer 不通过 Ray RPC 互相调用,也不通过 HTTP 直接传 Tensor。二者只需满足: + +- 连接同一个 Ray address; +- 使用同一个 Ray namespace; +- 使用同一个 TQ partition; +- 使用同一个 `run_id` 和协议版本。 + +### 3.3 Consumer launcher + +`python -m verl_speco.draft_train_launcher` 是父进程。它检查命令行 override,构造 `python -m torch.distributed.run ...` 命令,然后启动训练子进程。 + +launcher 自己不连接 TQ、不取样本、也不持有 GPU 模型。 + +### 3.4 Consumer training rank + +`torchrun --nproc_per_node=N` 会启动 N 个训练 OS 进程。每个进程有独立的: + +- global rank; +- local rank; +- GPU; +- `DrafterBaseTrainer`; +- `TQFeatureStore` 和本地 TQ client; +- DSpark 模型分片及 optimizer 状态。 + +这些 rank 共同执行一个分布式草稿模型训练任务。rank 0 额外负责发现和删除 TQ key;但所有 rank 都会取得各自的训练 Tensor,并参加模型 collective、梯度同步和 optimizer step。 + +### 3.5 Ray 和 torch.distributed 的职责不同 + +本方案仍然使用 Ray,但只因为 TQ 0.1.7 通过 Ray named actor 找 Controller。Consumer 不创建用于训练的 Ray actor,训练本身仍由 `torchrun` 和 `torch.distributed` 执行。 + +两种通信分别是: + +- Ray/TQ:Owner、Producer、每个 Consumer rank 连接共享 TQ;大 Tensor 通过 TQ backend 传输。 +- `torch.distributed`:训练 rank 之间广播小型 key/tag 命令、同步成功状态、训练模型 collective。 + +## 4. 共同配置以及“连接同一个 TQ”的实现 + +关键配置位于: + +```yaml +actor_rollout_ref: + rollout: + drafter: + training: + mode: offline + feature_store: + type: tq + path: null + transfer_queue: + enable: true + ray: + address: 127.0.0.1:6379 + namespace: speco-drafter + partition_id: speco_drafter_features + run_id: dspark-standalone-run + schema_version: 1 + poll_interval_seconds: 0.5 + drop_last: true +``` + +这些字段的含义如下: + +| 字段 | 使用者 | 含义 | +|---|---|---| +| `feature_store.type=tq` | Consumer | 选择流式 TQ source,而不是磁盘 shard/replay source | +| `feature_store.path=null` | Consumer | TQ 不从本地路径读文件,因此无需 path | +| `transfer_queue.enable` | Owner、Producer、Consumer | 开启 bridge 的 TQ 路径 | +| `ray.address` | 三端 | 连接同一个 Ray 集群 | +| `ray.namespace` | 三端 | 在同一 actor namespace 查找 named Controller | +| `partition_id` | 三端 | 对同一个 TQ KV 分区执行 put/list/get/clear | +| `run_id` | Producer、Consumer | 在共享 partition 中区分本次训练数据;Consumer 只接收匹配的样本 | +| `schema_version` | Producer、Consumer | 共同使用的数据协议版本 | +| `poll_interval_seconds` | Consumer rank 0 | ready 数量不足时的轮询间隔 | +| `drop_last` | Consumer | 第一版必须为 true;EOS 后不足 global batch 的尾样本被清理 | + +三端并不是通过共享 Python 对象得到这些配置。每个进程都各自读取相同取值,然后执行: + +```python +configure_transfer_queue(config) +connect_ray_cluster(ray_address, ray_namespace) +connect_transfer_queue_client() +``` + +`connect_ray_cluster()` 内部调用 `ray.init(address=..., namespace=...)`。`connect_transfer_queue_client()` 再调用无参数 `tq.init()`;TQ 由此在当前 Ray namespace 查找 Owner 创建的 named Controller。之后所有 KV 操作都显式携带相同的 `partition_id`。 + +因此,“连接同一个 TQ”实际由三层身份共同决定:同一 Ray 集群、同一 namespace 下的同一 named Controller、同一 `partition_id`。 + +## 5. 一条样本在 TQ 中的实际格式 + +### 5.1 一条 key 对应一个 sample + +本协议没有把一个训练 batch 存成一个 TQ key。一条 key 对应一条独立训练样本。假设: + +```text +run_id = dspark-run-001 +sequence_no = 17 +sample_id = prompt-000017 +``` + +则 key 为: + +```text +drafter:v1:dspark-run-001:000000000017:prompt-000017 +``` + +`sequence_no` 是本次 run 内的样本顺序号,不是 batch 编号,也不是 optimizer step。Consumer 用它稳定排序,之后每次从有序 ready 列表前部取一个 global batch。 + +### 5.2 tag:用于轻量发现和过滤 + +该 key 的 tag 是普通小字典: + +```python +{ + "record_type": "sample", + "status": "ready", + "schema_version": 1, + "run_id": "dspark-run-001", + "sequence_no": 17, + "sample_id": "prompt-000017", + "algorithm": "DSPARK", +} +``` + +tag 存在 TQ 的 KV 元信息中。`kv_list(partition_id=...)` 返回 `key -> tag`,不需要先加载 hidden states。rank 0 正是依靠 tag 筛选当前 run、当前 schema、DSPARK 且状态为 ready 的记录。 + +### 5.3 fields:真正的 Tensor payload + +同一 key 的 fields 是一个 Tensor 字典: + +```python +{ + "input_ids": Tensor[int64, shape=[L]], + "loss_mask": Tensor[float32, shape=[L]], + "position_ids": Tensor[int64, shape=[L]], + "hidden_states": Tensor[dtype, shape=[L, D]], + "metadata_json": Tensor[uint8, shape=[M]], + # 以下是可选字段: + "last_hidden_states": Tensor[..., ...], + "target": Tensor[..., ...], + "target_logprobs": Tensor[..., ...], +} +``` + +这里 `L` 是 feature window 的 token 数,`D` 是目标模型 hidden size,`M` 是 metadata JSON 序列化后的 UTF-8 字节数。 + +`hidden_states` 等大 Tensor 只存于 fields,通过 TQ `kv_batch_get` 传输;不会放进 tag,也不会通过训练 rank 的 object broadcast。 + +### 5.4 metadata_json:内容丰富但仍随 fields 读取 + +TQ fields 只能承载 Tensor,因此结构化 metadata 被编码为 `uint8` Tensor。解码后的字典格式是: + +```python +{ + "schema_version": 1, + "run_id": "dspark-run-001", + "sample_id": "prompt-000017", + "sequence_no": 17, + "algorithm": "DSPARK", + "target_model_id": "/models/Qwen3-8B", + "target_model_revision": "main", + "tokenizer_fingerprint": "...", + "target_layer_ids": [35], + "hidden_states_layout": "token_major", + "hidden_dtype": "bfloat16", + "hidden_shape": [L, D], + "feature_length": L, + "full_sequence_length": 256, + "feature_start": 64, + "feature_end": 64 + L, + "use_logits": False, +} +``` + +字段分工是: + +- tag:只放发现、过滤、排序所需的小字段;`kv_list` 可直接得到。 +- fields:放训练 Tensor 和完整 metadata;只有被某个 rank 选中后才 `kv_batch_get`。 +- key:把 tag 和 fields 重新关联起来,也是清理记录时传给 `kv_clear` 的标识。 + +### 5.5 控制记录 + +控制记录与 sample 放在同一 partition,但通过 tag 的 `record_type=control` 区分。 + +Owner readiness key: + +```text +control:v1::owner-ready +``` + +EOS key: + +```text +control:v1::eos +``` + +EOS tag 包含 `status=eos` 和 `total_samples`。EOS 表示 Producer 不会再为本次 run 增加新样本;它不是一条训练样本。 + +## 6. 公共协议如何把 Producer 输出还原为训练对象 + +Producer 应调用: + +```python +fields = encode_sample(sample, metadata) +key = make_sample_key(metadata) +tag = make_ready_tag(metadata) +put_sample(key, fields, tag=tag) +``` + +`encode_sample()` 会将所有 Tensor detach、转到 CPU、整理为 contiguous,并统一 `input_ids/position_ids` 为 int64、`loss_mask` 为 float32。随后校验 token 长度、hidden shape 和 metadata 一致,再把 metadata JSON 编成 uint8 Tensor。 + +Consumer 的逆过程位于 `TQFeatureStore.get_many()`: + +```python +records = get_samples([entry.key for entry in entries]) +sample = decode_sample( + key=key, + tag=entry.tag, + fields=fields, + expected_config=self.expected_config, +) +``` + +`get_samples()` 最终调用一次 TQ `kv_batch_get(keys=[...], partition_id=...)`。bridge 将 TQ 返回的 batched TensorDict 或 mapping 拆成与请求 key 顺序一致的普通 fields 字典。 + +`decode_sample()` 随后: + +1. 检查必需 fields 是否存在。 +2. 将 `metadata_json` 从 uint8 Tensor 还原为字典和 `SampleMetadata`。 +3. 根据 metadata 重新计算 key,并与实际 key 比较。 +4. 比较 tag 与 metadata 的公共身份字段。 +5. 检查 Consumer 的 expected config。 +6. 将 Tensor detach 到 CPU,统一基础 dtype/shape。 +7. 检查所有主 Tensor 第一维等于 `feature_length`,hidden shape/dtype 与 metadata 一致。 +8. 构造 `DraftFeatureSample`。 + +输出不再是 TQ 专用对象,而是既有训练代码使用的: + +```python +DraftFeatureSample( + input_ids=..., + loss_mask=..., + position_ids=..., + hidden_states=..., + metadata=..., + ..., +) +``` + +这是能够复用原训练逻辑的关键边界:TQ 只负责上游存储和传输,`decode_sample()` 后的数据类型与磁盘 feature store 读取结果一致。 + +## 7. Consumer 从启动到结束的完整执行流程 + +### 阶段 1:launcher 检查配置并启动 torchrun + +执行者是 launcher 父进程。入口是 `verl_speco.draft_train_launcher.main()`。 + +当 override 中出现 `feature_store.type=tq`,`validate_tq_launch_config()` 会要求: + +- `training.transfer_queue.enable=true`; +- `training.transfer_queue.ray.address` 非空; +- `training.transfer_queue.run_id` 非空。 + +检查成功后构造: + +```text +python -m torch.distributed.run + --nnodes=... + --nproc_per_node=... + -m verl_speco.draft_train + <全部 Hydra overrides> +``` + +配置参数是普通子进程命令行参数。此阶段没有 TQ Tensor 传输。 + +### 阶段 2:每个 rank 初始化训练运行时 + +每个 torchrun 子进程进入 `run_standalone_draft_training()`,调用 `_init_distributed()` 得到 `rank/local_rank/world_size`,绑定本 rank GPU,然后构造原有 `DrafterBaseTrainer` 和 DSpark backend。 + +`speculative_algorithm=DSPARK` 决定 backend 和 DSpark 模型训练实现;`feature_store.type=tq` 只改变数据来源,不替换 trainer。 + +### 阶段 3:factory 创建 TQFeatureStore + +训练循环调用: + +```python +store = build_feature_store_from_config( + feature_store_cfg, + read_only=True, + transfer_queue_cfg=training_cfg.get("transfer_queue"), +) +``` + +factory 在 `type=tq` 时不读取 `feature_store.path`,而是把 sibling `training.transfer_queue` 交给 `TQFeatureStore.from_config()`。 + +TQ store 被限定为 `read_only=True`,意思是它是训练 Consumer source。这里的“read only”不表示永不修改 TQ;成功消费后仍可通过明确的 `clear_many()` 删除记录,但不会把它当作通用 feature writer。 + +### 阶段 4:所有 rank 分别连接同一个 TQ + +训练循环调用 `_connect_tq_store_across_ranks()`。每个 rank 都独立执行 `store.connect()`: + +```text +configure_transfer_queue +→ ray.init(address, namespace) +→ tq.init() 连接 named Controller +→ 本 rank 设置 _connected=True +``` + +之后 `_all_ranks_true()` 使用 `dist.all_reduce(MIN)` 汇总连接结果。只要一个 rank 连接失败,所有 rank 都停止,不允许部分 rank 进入后续 broadcast 或 FSDP collective。 + +这里没有“rank 0 建一个 client 给其他 rank 共用”。TQ client 是进程本地对象,N 个 rank 有 N 个 client,但它们指向同一 Controller/partition。 + +### 阶段 5:创建 TQFeatureDataLoader + +每个 rank 构造自己的 loader,参数包括相同的 `batch_size_per_gpu`、`world_size`、轮询间隔和 drop-last,以及不同的 `rank`。 + +假设: + +```text +world_size = 2 +batch_size_per_gpu = 2 +global_batch_size = 4 +``` + +那么只有 ready 数量至少为 4,rank 0 才发布一个 batch 命令。 + +### 阶段 6:rank 0 发现 ready key + +rank 0 首先检查 `owner_ready()`。Owner 尚未发布 readiness marker 时,rank 0 sleep 后继续轮询,不会让其他 rank 开始取数。 + +Owner ready 后,rank 0 调用 `list_ready()`,其底层是: + +```text +tq.kv_list(partition_id) +→ key -> tag +→ 按 record_type/status/run/schema/algorithm 过滤 +→ 按 (sequence_no, key) 排序 +``` + +此阶段没有读取 fields,因此 hidden states 尚未传到训练进程。 + +### 阶段 7:rank 0 切分 global batch + +若排序后的前四条是 `k0、k1、k2、k3`,`build_assignments()` 产生: + +```python +assignments = [ + [ReadyEntry(k0, tag0), ReadyEntry(k1, tag1)], # rank 0 + [ReadyEntry(k2, tag2), ReadyEntry(k3, tag3)], # rank 1 +] +``` + +每条样本只出现在一个 rank 的 assignment 中,因此各 rank 不会取得同一训练样本。这里采用连续、不重叠的切片。 + +rank 0 随后构造普通 Python 命令字典: + +```python +{ + "kind": "batch", + "global_keys": [k0, k1, k2, k3], + "assignments": [ + [{"key": k0, "tag": tag0}, {"key": k1, "tag": tag1}], + [{"key": k2, "tag": tag2}, {"key": k3, "tag": tag3}], + ], +} +``` + +`global_keys` 只用于 rank 0 在训练完成后一次清理整个 batch;`assignments` 用于每个 rank 知道自己应该 get 哪些 key。 + +### 阶段 8:小型命令通过 torch.distributed 广播 + +各 rank 同时进入: + +```python +dist.broadcast_object_list(payload, src=0) +``` + +rank 0 的 payload 中是上述字典,其他 rank 的初始值是 `None`。PyTorch 会序列化这个普通 Python 对象并广播给所有 rank。 + +这条边界只传输字符串、整数和小字典 tag。`hidden_states`、`input_ids` 等 fields 不在命令中,所以不会经 rank 0 中转,也不会随 broadcast 复制完整 global batch Tensor。 + +### 阶段 9:每个 rank 直接从 TQ 取本地 payload + +每个 rank 从 `assignments[self.rank]` 还原自己的 `ReadyEntry`: + +```python +local_entries = [_entry_from_wire(item) for item in assignments[self.rank]] +samples = self.store.get_many(local_entries) +``` + +在上述例子中: + +- rank 0 调用 `kv_batch_get(keys=[k0, k1], partition_id=...)`; +- rank 1 调用 `kv_batch_get(keys=[k2, k3], partition_id=...)`。 + +大 Tensor 的数据面因此是 TQ storage 到目标训练 rank,不经过训练 rank 0 的 Python 内存。每个 rank 得到两个 CPU `DraftFeatureSample`。 + +loader yield: + +```python +TQLocalBatch( + local_keys=[本 rank 的 key], + local_samples=[本 rank 的 DraftFeatureSample], + global_keys=[完整 global batch key] if rank == 0 else None, +) +``` + +非 rank 0 不保存 `global_keys`,避免多个 rank 都尝试 clear。 + +### 阶段 10:复用已有训练 batch 构造 + +训练循环识别 `TQLocalBatch` 后,只取: + +```python +samples = tq_local_batch.local_samples +``` + +然后调用原有接口: + +```python +batch = trainer.prepare_training_batch_from_samples( + materialized_samples, + step=optimizer_step, +) +``` + +TQ 路径禁止同时开启 `target_feature_pipeline`,因为样本已经包含目标模型 hidden states,不需要训练侧再访问 vLLM materialize 一次。 + +此时 TQ 专用的 key/tag 已不参与 DSpark 数学计算;训练接口看到的是普通 `DraftFeatureSample`,并按原逻辑整理 input ids、hidden states、mask、position ids 和 DSpark 训练所需输入。 + +### 阶段 11:所有 rank 同步 batch 是否可训练 + +每个 rank 判断 `batch is not None`,再通过 `_all_ranks_true()` 做 `all_reduce(MIN)`。 + +只有所有 rank 都成功构造 batch,才能进入训练。如果任一 rank 解码或 batch 构造失败,TQ 路径直接报错,而且这些 key 不会被删除。 + +### 阶段 12:执行原有 DSpark training step + +每个 rank 调用: + +```python +ok = await trainer.training_step_from_batch(batch, optimizer_step) +``` + +该调用复用既有模型 forward、DSpark loss(包括配置开启时的 L1 loss)、backward、梯度同步和 optimizer step。TQ 新代码没有重新实现 loss 或 optimizer。 + +之后再次以 `_all_ranks_true(ok)` 同步。只有所有 rank 都返回成功,才认为这一个 global batch 已经安全消费。 + +### 阶段 13:训练成功后由 rank 0 删除 global batch + +训练循环调用 `_clear_tq_batch_across_ranks()`: + +1. rank 0 使用 `tq_local_batch.global_keys` 调用 `loader.clear_completed_batch()`。 +2. loader 调用 `store.clear_many(global_keys)`。 +3. bridge 最终调用 `tq.kv_clear(keys=[k0,k1,k2,k3], partition_id=...)`。 +4. 所有 rank 通过 `all_reduce(MAX)` 同步 clear 是否失败。 + +删除发生在 optimizer step 全 rank 成功之后。不是“某个 rank get 完就删除”,因为 get 完只代表 Tensor 已读取,不能代表训练 step 已成功。 + +clear 成功后才增加 `successful_steps`,然后复用原有 metrics 和 checkpoint 调度。 + +### 阶段 14:EOS 和尾 batch + +当 ready 样本少于一个 global batch时,rank 0 查询 EOS: + +- 没有 EOS:说明 Producer 以后仍可能写入更多样本,sleep 后继续轮询。 +- 已有 EOS 且 ready 为空:广播 `{"kind": "stop"}`,所有 rank 结束迭代。 +- 已有 EOS 且存在不足一个 global batch 的尾样本:rank 0 先 clear 这些尾 key,再广播 stop。 + +第一版强制 `drop_last=true`,因此不会构造各 rank batch size 不一致的最后一步。 + +### 阶段 15:checkpoint 和退出清理 + +正常 step 完成后仍按原 `save_interval_steps` 保存 checkpoint;循环结束后按 `save_final_checkpoint` 决定是否保存最终 checkpoint。 + +`finally` 中每个 rank 调用 `store.close()`。对 `TQFeatureStore` 而言,这只是: + +```text +关闭本进程 TQ client +→ 如果本进程自行 ray.init,则 ray.shutdown() +``` + +它不会调用全局 `tq.close()`,不会杀死 Owner 创建的 Controller,也不会影响仍在运行的 Producer 或其他 rank。 + +## 8. 控制面和数据面的完整边界 + +| 数据 | 从哪里到哪里 | 传输机制 | 是否经过 rank 0 | +|---|---|---|---| +| 启动配置 | launcher 到 torchrun 子进程 | 命令行 Hydra overrides | 每个 rank 都收到 | +| ready key/tag | TQ Controller 到 rank 0 | `tq.kv_list` | 是,只有 rank 0 list | +| batch assignment | rank 0 到全部 rank | `dist.broadcast_object_list` | 由 rank 0 发出 | +| hidden states 等 fields | TQ storage 到被分配的 rank | `tq.kv_batch_get` | rank 1 的 Tensor 不经过 rank 0 | +| batch 准备/训练成功状态 | 全部 rank 之间 | Tensor `all_reduce` | collective,无单点 payload relay | +| clear 请求 | rank 0 到 TQ | `tq.kv_clear(global_keys)` | 只有 rank 0 发起 | +| 梯度和模型 collective | 训练 rank 之间 | 既有 PyTorch distributed/FSDP 路径 | 与 TQ 无关 | + +## 9. 当前“最简单校验”具体简单在哪里 + +`TQFeatureStore` 构造的 expected config 只固定: + +```python +ExpectedFeatureConfig( + run_id=<当前训练 run_id>, + schema_version=<当前 schema>, + algorithm="DSPARK", +) +``` + +因此当前不会拿 Consumer 配置额外比较: + +- target model ID/revision; +- tokenizer fingerprint; +- target layer IDs; +- hidden layout; +- hidden dtype 的外部预期值。 + +但这不等于完全不校验。`decode_sample()` 仍然强制检查: + +- 必需 fields 存在; +- key、tag、metadata 三者身份一致; +- schema/run/algorithm 符合 Consumer; +- Tensor 类型正确; +- input/mask/position/hidden 长度一致; +- hidden 实际 shape/dtype 与该样本 metadata 一致; +- feature window 合法。 + +这满足“第一版少做外部模型身份检查”,同时避免把结构损坏或错 run 的数据送入训练。 + +## 10. 失败、删除和重复消费语义 + +当前实现遵循以下规则: + +1. 连接失败:所有 rank 同步停止。 +2. rank 0 list/EOS 失败:rank 0 广播 error 命令,其他 rank 不会永久等待 batch broadcast。 +3. 某 rank get/decode 失败:`_next_batch_across_ranks()` 将失败同步给全部 rank,不进入模型训练 collective。 +4. 某 rank 无法构造 batch:报错,global keys 保留在 TQ。 +5. 某 rank training step 失败:报错,global keys 保留在 TQ。 +6. 全 rank training step 成功:rank 0 clear 整个 global batch。 +7. clear 失败:错误传播到全部 rank,训练停止;不会把该 step 继续当成已正常完成。 +8. 达到 `max_steps`:循环停止;尚未选择的 ready 样本保留在 TQ。 + +第一版尚未实现完整的崩溃恢复协议。尤其是“optimizer step 已成功,但进程在 clear 前崩溃”时,key 仍存在;重新启动 Consumer 可能再次读取它。要实现严格 exactly-once,需要把 checkpoint step、已消费 sequence 或事务状态纳入协议。该能力应作为后续增强,而不是当前已实现能力。 + +## 11. 如何启动和检查 + +推荐启动顺序: + +1. 启动 Ray head。 +2. 启动 TQ Owner,并保持该进程存活。 +3. 启动 Producer,使用相同 Ray address、namespace、partition 和 run ID。 +4. 启动 `tools/run_dspark_tq_consumer.sh`。 +5. Producer 完成全部样本后发布 EOS。 +6. Consumer 消费完成并退出后,再停止 Owner/Ray。 + +示例: + +```bash +MODEL_PATH=/models/Qwen3-8B \ +DRAFTER_PATH=/models/dspark-drafter \ +DRAFT_CKPTS_DIR=/checkpoints/dspark-tq \ +TRAIN_DEVICES=0,1,2,3 \ +TRAIN_GPUS=4 \ +RAY_ADDRESS=127.0.0.1:6379 \ +SPECO_TQ_RUN_ID=dspark-run-001 \ +bash tools/run_dspark_tq_consumer.sh +``` + +脚本中的 namespace 固定为 `speco-drafter`,默认 partition 来自公共配置 `speco_drafter_features`。Owner 和 Producer 必须使用相同值。 + +## 12. 已完成的测试 + +### 12.1 Consumer/factory/协议/launcher 单元测试 + +已执行: + +```text +python -m pytest \ + tests/unit/test_tq_consumer.py \ + tests/unit/test_draft_train_launcher.py \ + tests/unit/test_transferqueue_bridge.py \ + tests/unit/test_drafter_sample_protocol.py \ + tests/unit/test_draft_feature_store.py \ + -q +``` + +结果:`44 passed`。 + +### 12.2 真实 TQ 0.1.7 跨进程 smoke + +`tools/tq_connection_smoke.py` 使用真实 Ray + TQ Owner 和另一个 Consumer 进程验证了: + +- Owner 发布 owner-ready; +- 两条 sample 写入 TQ; +- Consumer 经 `TQFeatureStore` 和 `TQFeatureDataLoader` 读到两条样本; +- hidden shape 正确; +- Consumer clear 已完成 batch; +- EOS 后迭代停止; +- Consumer 只关闭本地 client,Owner 仍能继续观察完成标记并正常关闭。 + +实际 smoke 输出包含: + +```text +CLIENT_OK samples=2 shape=(3,4) +CLIENT_CLOSED_LOCAL_ONLY +OWNER_OBSERVED_SAMPLES_CLEARED +OWNER_CLOSED +``` + +### 12.3 当前环境未覆盖的部分 + +完整 `tests/unit/test_draft_training_loop.py` 在当前 Windows 环境无法完整收集,因为上游 `verl/ray` 依赖不齐;新增训练循环测试代码已通过 Python 编译检查,连接/clear helper 也通过针对性单元逻辑验证。真实多 GPU DSpark 训练仍需要在目标 Linux GPU 环境执行集成测试。 + +## 13. 当前限制和后续建议 + +当前第一版有意不实现以下复杂能力: + +1. Producer 本身尚未在本次 Consumer 改动中实现。 +2. 只支持 `algorithm=DSPARK`。 +3. 只支持 `drop_last=true`。 +4. 不支持 TQ 与 `target_feature_pipeline.enabled=true` 同时开启。 +5. 不提供严格的 crash exactly-once 或 checkpoint/queue 联合恢复。 +6. rank 0 仍通过 `kv_list` 轮询整个 partition;数据量很大时可考虑 cursor/ready queue 优化。 +7. 当前外部 expected config 校验较简化,后续可把 model revision、tokenizer fingerprint、layer/layout/dtype 预期接入 Hydra 配置。 +8. 尚需在真实多机、多 GPU、Mooncake backend 环境验证吞吐、背压、Owner 生命周期和网络故障行为。 + +建议下一阶段优先完成 Producer,并严格复用 `drafter_sample_protocol.py`,不要在 Producer 另造一套 key/tag/fields 格式。完成 Producer 后,首先跑 world size 1 的端到端训练,再跑多 rank 验证每条 key 只分配给一个 rank、训练成功后只由 rank 0 clear。 + +## 14. 最终路径摘要 + +```text +Producer(待实现) + vLLM 并行 prefill + → DraftFeatureSample + SampleMetadata + → encode_sample 得到 Tensor fields + → TQ kv_put(key, fields, tag) + +Consumer rank 0 + kv_list 只取 key/tag + → 过滤并按 sequence_no 排序 + → 切出 global batch + → broadcast 每个 rank 的 key/tag assignment + +每个 Consumer rank + 取 assignments[rank] + → kv_batch_get 本 rank keys + → decode_sample 得到 DraftFeatureSample + → 原 prepare_training_batch_from_samples + → 原 DSpark training_step_from_batch + +全部 rank + 同步确认 optimizer step 成功 + → rank 0 kv_clear(global_keys) + → 原 metrics/checkpoint + → 下一批 + +Producer 发布 EOS + → rank 0 确认没有完整 global batch + → 清理不足一批的尾样本 + → broadcast stop + → 各 rank 关闭本地 TQ client 并退出 +``` diff --git a/docs/standalone_tq_foundation_implementation.md b/docs/standalone_tq_foundation_implementation.md new file mode 100644 index 00000000..e8f95581 --- /dev/null +++ b/docs/standalone_tq_foundation_implementation.md @@ -0,0 +1,1037 @@ +# Standalone TQ 公共基础层实现说明 + +## 1. 文档范围和已验证结论 + +本文解释当前仓库已经实现并测试通过的 TQ 公共基础层: + +```text +verl_speco/transport/drafter_sample_protocol.py +verl_speco/integration/transferqueue_bridge.py +verl_speco/config/speco_base.yaml +verl_speco/tq_owner.py +tools/run_dspark_tq_owner.sh +tools/tq_connection_smoke.py +tests/unit/test_drafter_sample_protocol.py +tests/unit/test_transferqueue_bridge.py +pyproject.toml +``` + +当前已经实现: + +1. Producer 和 Consumer 共用的样本 key、tag、fields 和 metadata 协议; +2. 普通进程连接 Ray 集群; +3. TQ Owner 创建 named `TransferQueueController`; +4. 独立 Client 发现并连接同一个 Controller; +5. 单样本 put、元数据 list、批量 get 和批量 clear; +6. Owner 与 Client 不同的关闭边界; +7. 独立 Owner 入口和共享 Hydra 配置; +8. mock 单元测试和真实双进程 TQ 0.1.7 smoke test。 + +当前还没有实现: + +1. Producer 读取 JSONL、并发访问多个 vLLM endpoint 的完整 pipeline; +2. `feature_store.type=tq` 工厂分支; +3. `TQFeatureStore` 和 `TQFeatureDataLoader`; +4. rank 0 选择 global keys、各 rank 读取 local keys; +5. TQ batch 接入 DSpark optimizer step; +6. optimizer step 成功后的 rank 0 clear。 + +因此当前代码已经证明“两个独立进程能通过 Ray 连接同一个 TQ,并按共享协议批量传输多个样本”,但尚未接到正式 Producer 和 DSpark Consumer 主循环。 + +## 2. 运行时角色和术语 + +### 2.1 Ray head + +Ray head 是 Ray 集群的控制节点。TQ 0.1.7 使用 Ray 管理 named Controller 和 storage actors。 + +Ray 不传输 standalone hidden-state payload。当前代码没有对这些样本调用: + +```python +ray.put(hidden_states) +``` + +### 2.2 TQ Owner + +TQ Owner 是普通 Python OS 进程,入口为: + +```text +python -m verl_speco.tq_owner +``` + +它不是 Ray actor。它先调用 `ray.init(address=...)` 加入 Ray,再调用带完整配置的 `tq.init(config)` 创建 TQ Controller 和 storage。 + +Owner 是唯一允许调用全局 `tq.close()` 的进程。 + +### 2.3 Named TransferQueueController + +TQ 0.1.7 内部创建: + +```python +TransferQueueController.options( + name="TransferQueueController" +).remote(...) +``` + +`name` 将 Controller actor 注册到 Ray actor registry。其他进程加入同一 Ray address 和 namespace 后,通过: + +```python +ray.get_actor("TransferQueueController") +``` + +取得 actor handle,再读取 TQ backend 配置。 + +Controller 保存控制信息和数据位置;使用 MooncakeStore 时,大 tensor 本身保存在 MooncakeStore。 + +### 2.4 TQ Client + +Owner、Producer、每个 Consumer rank 都在各自 OS 进程内拥有独立 TQ Client。 + +普通 Client 通过无参: + +```python +tq.init() +``` + +发现 named Controller并初始化本进程的 storage manager。Client 不是 Ray actor,Producer 和 torchrun rank 也不需要改成 Ray actor。 + +### 2.5 Partition、key、tag 和 fields + +当前固定 partition: + +```text +speco_drafter_features +``` + +TQ 中一条记录逻辑上是: + +```text +partition_id +└── key + ├── tag:轻量 dict,由 kv_list 发现 + └── fields:Tensor payload,由 kv_batch_get 读取 +``` + +## 3. 共享配置如何工作 + +共享配置定义在 `verl_speco/config/speco_base.yaml`: + +```yaml +transfer_queue: + enable: false + package_version: "0.1.7" + ray: + address: null + namespace: speco-drafter + partition_id: speco_drafter_features + run_id: null + schema_version: 1 + connect_timeout_seconds: 120 + poll_interval_seconds: 0.5 + drop_last: true + controller: + polling_mode: true + backend: + storage_backend: SimpleStorage + SimpleStorage: + total_storage_size: 100000 + num_data_storage_units: 8 + MooncakeStore: + auto_init: false + metadata_server: localhost:50050 + master_server_address: localhost:50051 + local_hostname: "" + protocol: tcp + global_segment_size: 4294967296 + local_buffer_size: 1073741824 + device_name: "" +``` + +### 3.1 Ray 连接字段 + +```yaml +ray: + address: 10.0.0.1:6379 + namespace: speco-drafter +``` + +它们决定当前进程连接哪个 Ray 集群,以及在哪个 namespace 查找 `TransferQueueController`。Owner、Producer 和所有 Consumer ranks 必须使用相同值。 + +### 3.2 SPECO 协议字段 + +```yaml +partition_id: speco_drafter_features +run_id: dspark-20260819-a +schema_version: 1 +``` + +这些字段用于路由和校验,不属于 TQ 原生配置。`run_id` 用于隔离不同 pipeline 的记录。 + +### 3.3 TQ 原生字段 + +```yaml +controller: ... +backend: ... +``` + +只有这些字段应传给 `tq.init(full_config)`。bridge 的 `_native_tq_config()` 会删除: + +```text +enable +package_version +ray +partition_id +run_id +schema_version +connect_timeout_seconds +poll_interval_seconds +drop_last +``` + +对象变化为: + +```text +完整 SPECO transfer_queue dict +→ _native_tq_config() +→ controller/backend等TQ字段 +→ OmegaConf DictConfig +→ tq.init() +``` + +## 4. Bridge 的进程内状态 + +`verl_speco/integration/transferqueue_bridge.py` 在每个 OS 进程内分别维护: + +```python +_state = { + "enabled": False, + "configured": False, + "initialized": False, + "config": None, + "owner": False, + "ray_initialized_here": False, + "ray_address": None, + "ray_namespace": None, +} +``` + +该 dict 不跨进程共享。Owner、Producer、每个 rank 分别拥有自己的 `_state`。 + +| 字段 | 含义 | +|---|---| +| `enabled` | 当前进程配置是否开启 TQ | +| `configured` | 是否调用过 `configure_transfer_queue()` | +| `initialized` | 当前进程是否执行过 `tq.init()` | +| `config` | 当前进程保存的普通 dict 配置 | +| `owner` | 当前进程是否创建了全局 Controller/Storage | +| `ray_initialized_here` | bridge 是否负责调用了本进程的 `ray.init()` | +| `ray_address/namespace` | 本进程的 Ray 连接信息 | + +`_state_lock` 只保护同一进程内多个线程同时初始化,不是分布式锁。 + +## 5. Owner 的完整启动数据流 + +Owner 入口是 `verl_speco/tq_owner.py`,直接加载 Hydra 主配置 +`verl_speco/config/speco_base.yaml`。`run_owner()` 会复制其中的 +`transfer_queue` 子配置,并仅在 Owner 进程的副本中强制设置 +`enable=true`;普通训练任务看到的共享默认值仍是 `false`。 + +### 阶段 1:读取配置 + +执行者:Owner OS 进程。 + +入口: + +```python +run_owner(config) +``` + +取得: + +```python +training_cfg = config.actor_rollout_ref.rollout.drafter.training +tq_cfg = training_cfg.transfer_queue +``` + +然后调用: + +```python +configure_transfer_queue(training_cfg) +``` + +该函数只把 OmegaConf 转成普通 dict并更新当前进程 `_state`,不会连接 Ray,也不会创建 TQ。 + +### 阶段 2:连接 Ray + +Owner 调用: + +```python +connect_ray_cluster(ray_address, namespace) +``` + +内部执行: + +```python +if not ray.is_initialized(): + ray.init(address=ray_address, namespace=namespace) +``` + +边界类型是 Ray control-plane connection。此时还没有传输训练 tensor。 + +### 阶段 3:创建 Controller 和 Storage + +Owner 调用: + +```python +start_transfer_queue_owner(tq_cfg) +``` + +执行顺序: + +1. `_extract_tq_config()` 得到普通 dict; +2. 检查 `enable=true`; +3. 检查 `TransferQueue` 包可用; +4. 防止本进程重复初始化; +5. `_native_tq_config()` 删除 SPECO 字段; +6. `_as_tq_config()` 转 OmegaConf; +7. `tq.init(native_config)` 创建 Controller、Storage 和 Owner Client; +8. 设置 `_state.owner=True`、`initialized=True`。 + +Ray 中形成: + +```text +Ray cluster / namespace +├── named actor: TransferQueueController +└── storage backend + ├── SimpleStorage actors + └── 或 MooncakeStore connection/process +``` + +### 阶段 4:发布 owner-ready + +调用: + +```python +publish_owner_ready(run_id, schema_version) +``` + +生成: + +```python +key = "control:v1::owner-ready" +fields = {"marker": torch.tensor([1], dtype=torch.uint8)} +tag = { + "record_type": "control", + "status": "owner_ready", + "schema_version": 1, + "run_id": run_id, +} +``` + +这是一条控制记录,不进入训练 batch。 + +### 阶段 5:常驻和关闭 + +Owner 安装 `SIGINT/SIGTERM` handler,并等待: + +```python +stop_event.wait() +``` + +收到信号后调用: + +```python +close_transfer_queue_owner() +``` + +执行: + +```text +验证owner身份 +→ tq.close() +→ 清理Controller/Storage +→ ray.shutdown() +``` + +Owner 必须在 Producer 和 Consumer 退出后才能关闭。 + +## 6. 普通 Client 如何连接同一个 TQ + +Producer 和每个 Consumer rank 后续使用相同顺序: + +```python +configure_transfer_queue(training_cfg) +connect_ray_cluster(ray_address, namespace) +connect_transfer_queue_client() +``` + +`connect_transfer_queue_client()` 最终调用无参: + +```python +tq.init() +``` + +TQ 0.1.7 内部通过: + +```python +ray.get_actor("TransferQueueController") +``` + +找到 Owner 创建的 Controller,读取 backend 配置,然后创建当前进程的 TQ Client。 + +对象和边界变化: + +```text +actor名称字符串 +→ Ray actor registry +→ Controller actor handle +→ Controller.get_config.remote() +→ TQ DictConfig +→ 当前进程TransferQueueClient +→ 同一个SimpleStorage/MooncakeStore +``` + +## 7. 一条具体样本的初始对象 + +真实 smoke test使用: + +```python +sample = DraftFeatureSample( + algorithm="DSPARK", + input_ids=torch.tensor([1, 2, 3]), # int64[3], CPU + loss_mask=torch.tensor([0.0, 1.0, 1.0]), # float32[3], CPU + position_ids=torch.tensor([0, 1, 2]), # int64[3], CPU + hidden_states=torch.arange( + 12, dtype=torch.float32 + ).reshape(3, 4), # float32[3,4], CPU +) +``` + +同时构造: + +```python +meta = SampleMetadata( + schema_version=1, + run_id="codex-batch-smoke", + sample_id="smoke-0000", + sequence_no=0, + algorithm="DSPARK", + target_model_id="smoke-target", + target_model_revision="smoke-revision", + tokenizer_fingerprint="smoke-tokenizer", + target_layer_ids=[0], + hidden_states_layout="dflash_aux", + hidden_dtype="float32", + hidden_shape=[3, 4], + feature_length=3, + full_sequence_length=3, + feature_start=0, + feature_end=3, + use_logits=False, +) +``` + +`DraftFeatureSample` 和 `SampleMetadata` 都是进程内 Python 对象,不直接经过 TQ。 + +## 8. Key 的生成和两个同名函数 + +共享协议调用: + +```python +make_sample_key(meta) +``` + +输出: + +```text +drafter:v1:codex-batch-smoke:000000000000:smoke-0000 +``` + +字段顺序: + +```text +drafter / schema version / run_id / 12位sequence_no / sample_id +``` + +`sequence_no` 是输入文件 record 序号,不是训练 batch 编号。 + +bridge 为兼容 PR #48 还保留另一个: + +```python +transferqueue_bridge.make_sample_key( + global_step, + replica_rank, + request_id, +) +``` + +它生成: + +```text +speco::: +``` + +standalone Producer 必须从 `verl_speco.transport.drafter_sample_protocol` import `make_sample_key`,不能使用 bridge 中的 PR #48 旧函数。 + +## 9. Tag 如何生成 + +```python +tag = make_ready_tag(meta) +``` + +输出: + +```python +{ + "record_type": "sample", + "status": "ready", + "schema_version": 1, + "run_id": "codex-batch-smoke", + "sequence_no": 0, + "sample_id": "smoke-0000", + "algorithm": "DSPARK", +} +``` + +tag 只包含发现、过滤和排序需要的标量/字符串。`list_samples()` 只返回 key/tag,不读取 hidden states。 + +## 10. `encode_sample()` 如何生成 fields + +调用: + +```python +fields = encode_sample(sample, meta) +``` + +### 10.1 校验 + +执行: + +```text +SampleMetadata.validate() +DraftFeatureSample.validate(strict=True) +``` + +随后检查: + +1. hidden states 是一个 dense tensor; +2. ids/mask/position 长度等于 `feature_length`; +3. hidden 第一维等于 `feature_length`; +4. hidden shape 等于 metadata; +5. hidden dtype 等于 metadata; +6. feature window 长度正确。 + +### 10.2 Tensor 规范化 + +```text +input_ids → CPU contiguous int64[L] +loss_mask → CPU contiguous float32[L] +position_ids → CPU contiguous int64[L] +hidden_states → CPU contiguous,保持模型dtype +``` + +没有 `position_ids` 时生成 `torch.arange(L, dtype=int64)`。 + +### 10.3 Metadata JSON 编码 + +```text +SampleMetadata dataclass +→ dict +→ JSON UTF-8 bytes +→ torch.uint8[M] +``` + +实现等价于: + +```python +raw = json.dumps(metadata).encode("utf-8") +metadata_json = torch.tensor(list(raw), dtype=torch.uint8) +``` + +### 10.4 最终 fields + +```python +fields = { + "input_ids": int64[3], + "loss_mask": float32[3], + "position_ids": int64[3], + "hidden_states": float32[3,4], + "metadata_json": uint8[M], +} +``` + +如果存在,协议也保留 `last_hidden_states`、`target` 和 `target_logprobs`。 + +## 11. Bridge 如何写入 TQ + +调用: + +```python +put_sample(key, fields, tag=tag) +``` + +bridge 执行: + +1. 检查 TQ 已启用; +2. 丢弃 fields 中非 tensor 值; +3. 确保本进程已经 `tq.init()`; +4. 取得配置中的 partition; +5. 调用: + +```python +tq.kv_put( + key=key, + partition_id="speco_drafter_features", + fields=fields, + tag=tag, +) +``` + +使用 MooncakeStore 时,大 tensor 路径是: + +```text +Producer CPU tensor +→ Producer TQ Client +→ MooncakeStore +``` + +不是 Ray `ObjectRef`。 + +## 12. Consumer 如何发现 key + +调用: + +```python +records = list_samples() +``` + +内部调用: + +```python +tq.kv_list(partition_id="speco_drafter_features") +``` + +标准化返回类型: + +```python +dict[str, dict[str, Any]] +``` + +示例: + +```python +{ + "drafter:v1:...:smoke-0000": { + "record_type": "sample", + "status": "ready", + "run_id": "codex-batch-smoke", + "sequence_no": 0, + ... + } +} +``` + +bridge 兼容 `key → tag` 和 `partition → key → tag` 两种 wrapper。这一步不读取 fields。 + +## 13. Consumer 如何批量取样本 + +输入: + +```python +keys = [key0, key1] +``` + +调用: + +```python +records = get_samples(keys) +``` + +bridge 只调用一次: + +```python +result = tq.kv_batch_get( + keys=keys, + partition_id="speco_drafter_features", +) +``` + +TQ 0.1.7 返回带 batch 维的 TensorDict。bridge 检查 `result.batch_size`,再执行: + +```python +rows = [result[index] for index in range(len(keys))] +``` + +每行转成普通 dict,最终返回: + +```python +[ + (key0, fields0), + (key1, fields1), +] +``` + +返回顺序与输入 keys 一致。重复 key 会提前报错。 + +## 14. `decode_sample()` 如何恢复训练对象 + +调用: + +```python +sample = decode_sample( + key, + tag, + fields, + expected_config, +) +``` + +### 14.1 Metadata 解码 + +```text +metadata_json uint8[M] +→ bytes +→ UTF-8 +→ json.loads +→ dict +→ SampleMetadata.from_dict +``` + +### 14.2 身份一致性 + +代码根据 metadata 重新生成 key,要求输入 key 完全相等;然后逐项校验 tag 的: + +```text +record_type/status/schema_version/run_id/sequence_no/sample_id/algorithm +``` + +所以 key、tag 和 payload metadata 不能来自不同样本。 + +### 14.3 Consumer 合同 + +Consumer 提供: + +```python +ExpectedFeatureConfig( + run_id="codex-batch-smoke", + schema_version=1, + algorithm="DSPARK", + target_model_id="smoke-target", + target_model_revision=None, + tokenizer_fingerprint=None, + target_layer_ids=None, + hidden_states_layout="dflash_aux", + hidden_dtype="float32", +) +``` + +值为 `None` 的字段不检查;其他字段必须完全一致。正式 Consumer 应填写 target checkpoint、tokenizer、layers、layout 和 dtype,避免使用错误 target 特征。 + +### 14.4 输出 + +完成 tensor 类型、长度、shape、dtype 校验后,构造: + +```python +DraftFeatureSample.from_dict(payload, strict=True) +``` + +输出可以交给现有 `trainer.prepare_training_batch_from_samples()`。当前 smoke test验证到这里,正式 `TQFeatureDataLoader` 尚未实现。 + +## 15. EOS 控制记录 + +调用: + +```python +key, fields, tag = make_eos_record(run_id, total_samples) +``` + +输出: + +```python +key = "control:v1::eos" +fields = {"marker": torch.tensor([1], dtype=torch.uint8)} +tag = { + "record_type": "control", + "status": "eos", + "schema_version": 1, + "run_id": run_id, + "total_samples": total_samples, +} +``` + +EOS 表示不会再发布新样本。Consumer 应先 drain ready samples 再退出。协议函数已实现,正式 Producer/Consumer 尚未调用。 + +## 16. Clear 和数据生命周期 + +bridge 提供: + +```python +clear_samples(keys) +``` + +内部调用: + +```python +tq.kv_clear(keys=keys, partition_id="speco_drafter_features") +``` + +基础层只执行删除,不决定删除时机。正式 Consumer 必须遵守: + +```text +rank 0选择global keys +→ 各rank读取local keys +→ 所有rank完成同一optimizer step +→ 汇总global success +→ rank 0 clear global keys +``` + +不能在 `get_samples()` 后立即 clear,因为 get 成功不代表训练 step 成功。 + +## 17. Client close 和 Owner close + +### 17.1 Client close + +Producer/rank 调用: + +```python +close_transfer_queue_client() +``` + +执行: + +```text +tq.get_client() +→ 当前进程client.close() +→ 如果bridge负责ray.init,则ray.shutdown() +``` + +它不调用全局 `tq.close()`,不会 kill Controller。调用后本进程不能继续使用 TQ。 + +### 17.2 Owner close + +Owner 调用: + +```python +close_transfer_queue_owner() +``` + +执行: + +```text +tq.close() +→ Controller/Storage全局清理 +→ ray.shutdown() +``` + +Owner 若误用 Client close,bridge 会抛 `RuntimeError`。 + +## 18. PR #48 兼容边界 + +bridge 继续保留: + +```python +init_transfer_queue(config) +get_sample(key) +close_transfer_queue() +``` + +PR #48 的 `SpecoTaskRunner` 已是 Ray actor,所以不调用 standalone 的 `connect_ray_cluster()`。旧 worker 继续使用单 key `get_sample()`。 + +standalone 后续使用新增的: + +```python +list_samples() +get_samples(keys) +clear_samples(keys) +``` + +因此没有修改 PR #48 现有调用点的函数签名。 + +## 19. 依赖和命令入口 + +`pyproject.toml` 新增: + +```toml +[project.optional-dependencies] +transfer-queue = ["TransferQueue==0.1.7"] +``` + +安装: + +```bash +pip install -e ".[transfer-queue]" +``` + +TQ 0.1.7 会安装 Ray,并要求 `numpy<2.0.0`。若现有包要求 NumPy 2,需要隔离环境或重新解决依赖。 + +Owner 命令: + +```text +verl-speco-tq-owner +``` + +也可以使用 `tools/run_dspark_tq_owner.sh`。 + +## 20. 单元测试 + +协议测试 `tests/unit/test_drafter_sample_protocol.py` 覆盖: + +1. encode/decode round trip; +2. key 格式; +3. tag 身份冲突; +4. Consumer contract 冲突; +5. hidden shape 冲突; +6. EOS 格式。 + +bridge 测试 `tests/unit/test_transferqueue_bridge.py` 覆盖: + +1. Ray address/namespace 参数; +2. Owner 只向 TQ 传原生配置; +3. Client 无参 `tq.init()`; +4. put/list/get-many/clear; +5. batch 返回顺序; +6. Client close 不调用全局 close; +7. Owner 不能误用 Client close。 + +运行: + +```bash +python -m pytest \ + tests/unit/test_drafter_sample_protocol.py \ + tests/unit/test_transferqueue_bridge.py \ + -q +``` + +## 21. 真实双进程 smoke test + +程序: + +```text +tools/tq_connection_smoke.py +``` + +它使用真实 `TransferQueue==0.1.7`、Ray、SimpleStorage、两个独立 Python进程和两个 batch samples。 + +Owner 路径: + +```text +连接Ray +→ tq.init(full config) +→ 写sample 0和sample 1 +→ 等待client-done +→ clear done marker +→ 全局关闭 +``` + +Client 路径: + +```text +连接同一个Ray +→ tq.init() +→ kv_list发现两个key +→ 一次kv_batch_get([k0,k1]) +→ 拆成两个fields dict +→ 分别decode_sample +→ clear两个sample keys +→ 写client-done +→ 只关闭本地client +``` + +已验证输出: + +```text +OWNER_READY keys=[k0, k1] +CLIENT_OK samples=2 shape=(3, 4) +CLIENT_CLOSED_LOCAL_ONLY +OWNER_OBSERVED_SAMPLES_CLEARED +OWNER_CLOSED +``` + +这证明: + +1. 两个普通进程能连接同一个 TQ; +2. named Controller 发现有效; +3. 0.1.7 的 `kv_list/kv_batch_get/kv_clear` 参数有效; +4. TensorDict batch 能按 key 顺序拆开; +5. 共享协议能恢复 `DraftFeatureSample`; +6. Client close 不会杀掉 Owner; +7. Owner 能最终统一关闭。 + +## 22. 当前完整路径总结 + +```text +Owner +→ ray.init(address, namespace) +→ tq.init(native config) +→ named TransferQueueController + +普通Client +→ ray.init(same address, same namespace) +→ tq.init() +→ 找到同一个Controller + +DraftFeatureSample + SampleMetadata +→ make_sample_key +→ make_ready_tag +→ encode_sample +→ fields + metadata_json tensor +→ bridge.put_sample +→ tq.kv_put +→ SimpleStorage/MooncakeStore + +Consumer/测试Client +→ bridge.list_samples +→ key + tag +→ bridge.get_samples(keys) +→ tq.kv_batch_get +→ TensorDict batch +→ 每个key对应一个fields dict +→ decode_sample +→ DraftFeatureSample + +正式训练成功后(待实现) +→ bridge.clear_samples(global_keys) + +Client退出 +→ close_transfer_queue_client + +所有业务进程退出 +→ Owner close_transfer_queue_owner +→ tq.close +→ ray.shutdown +``` + +## 23. 下一阶段接入约束 + +后续代码不能重新定义协议或直接访问 TQ 私有对象。 + +Producer 应复用: + +```text +SampleMetadata +make_sample_key +make_ready_tag +encode_sample +bridge.put_sample +make_eos_record +``` + +Consumer 应复用: + +```text +bridge.list_samples +bridge.get_samples +decode_sample +bridge.clear_samples +``` + +下一阶段需要新增: + +```text +verl_speco/trainer/tq_feature_store.py +verl_speco/trainer/tq_sample_source.py +feature_store.py 的 type=tq 分支 +draft_training_loop.py 的流式训练分支 +Producer入口、输入读取和并发vLLM文件 +``` + +这些文件应建立在本文已经实现和真实验证过的连接、协议、KV 和关闭接口之上。 diff --git a/docs/standalone_vllm_tq_dspark_training_plan.md b/docs/standalone_vllm_tq_dspark_training_plan.md new file mode 100644 index 00000000..846c9053 --- /dev/null +++ b/docs/standalone_vllm_tq_dspark_training_plan.md @@ -0,0 +1,1131 @@ +# 独立 vLLM Producer + TQ + DSpark Consumer 第一版方案 + +## 1. 第一版要实现什么 + +只实现下面这条主链路: + +```text +包含 prompt + 预生成 response 的输入文件 +→ Producer 并发请求 vLLM prefill +→ Producer 将每条训练样本写入 TQ +→ Consumer 从同一个 TQ 取样本 +→ 独立 torchrun/FSDP DSpark 训练 +→ 一个 optimizer step 成功后删除该 step 的 TQ 样本 +``` + +第一版允许 **TQ 基础设施内部使用 Ray**,因为实测 `TransferQueue==0.1.7` 通过 Ray named actor 发现 `TransferQueueController`。Producer 和 Consumer 仍是普通 OS 进程,不改成 Ray actor;二者不使用 Ray RPC、Ray `ObjectRef` 或 Ray object store 传输训练样本。hidden-state payload 仍通过 TQ 的 `kv_put/kv_batch_get` 和配置的 MooncakeStore backend 传输。第一版不使用 DataProto/WorkerGroup,不生成长期 hidden-state feature store,也暂不实现复杂重试、自动重启和严格 checkpoint 数据恢复。 + +需要运行的组件: + +| 组件 | 数量 | 作用 | +|---|---:|---| +| vLLM server | 一个或多个 | 加载 target model,执行并行 prefill | +| Ray head | 1 个集群 | 保存 TQ named Controller/Storage actors;不承载 Producer/Consumer 业务 RPC | +| TQ owner | 1 个普通进程 | 连接 Ray,创建并保持 TQ controller/storage,最后统一关闭 TQ | +| Producer | 1 个进程 | 读文件、并发请求 vLLM、写 TQ | +| Consumer | 1 个 torchrun 任务 | 多个 rank 从 TQ 取数并训练 DSpark | + +Producer 和 Consumer 不互相调用,也不通过 HTTP 或 Ray 传样本。二者先连接同一个 Ray 集群,再通过无参 `tq.init()` 找到同一个 TQ,最后使用 TQ KV API 读写样本。 + +## 2. 共同的数据约定 + +这部分由两位开发者共同完成并先合入。建议文件: + +```text +verl_speco/transport/drafter_sample_protocol.py +tests/unit/test_drafter_sample_protocol.py +``` + +### 2.1 一个 key 对应一条样本 + +第一版固定: + +```text +一个输入文件 record +→ 一个 sequence_no +→ 一个 sample_id +→ 一个 TQ sample_key +→ 一个单样本 payload +``` + +`sequence_no` 是输入文件中的样本序号,不是 batch 编号。多个 sample keys 到 Consumer 后才组成训练 batch。 + +例如: + +```python +run_id = "dspark-20260818-a" +sequence_no = 17 +sample_id = "train-000017" + +partition_id = "speco_drafter_features" # 第一版沿用 PR 固定值 +sample_key = ( + "drafter:v1:dspark-20260818-a:" + "000000000017:train-000017" +) +``` + +### 2.2 Partition、key、tag 和 payload 的关系 + +TQ 中逻辑上是: + +```text +TQ 实例 +└── partition_id + └── sample_key + ├── tag + └── fields/payload +``` + +- TQ 实例:Producer 和 Consumer 共同连接的 controller/storage; +- `partition_id`:第一版沿用 PR 固定的 `speco_drafter_features`;每次启动新的 TQ owner,保证该 TQ 实例初始为空; +- `sample_key`:该分区中一条训练样本的地址; +- `tag`:轻量索引,Consumer 用 `kv_list()` 看见; +- `fields/payload`:真正的 Tensor 数据,Consumer 用 `kv_batch_get()` 读取。 + +Producer 写入: + +```python +tq.kv_put( + partition_id=partition_id, + key=sample_key, + fields=fields, + tag=tag, +) +``` + +Consumer 先发现 key: + +```python +all_records = tq.kv_list() +tags_by_key = all_records[partition_id] +``` + +这一步只拿 key 和 tag,不搬运 hidden states。 + +Consumer 再取数据: + +```python +result = tq.kv_batch_get( + partition_id=partition_id, + keys=selected_keys, +) +``` + +`kv_batch_get` 中的 batch 表示“一次读取多个独立 sample keys”,不是这些样本在 Producer 写入时就属于同一个对象。 + +### 2.3 Payload 字段 + +一个 sample key 对应的 `fields`: + +```python +fields = { + "input_ids": input_ids, # CPU int64[L] + "loss_mask": loss_mask, # CPU float32[L] + "position_ids": position_ids, # CPU int64[L] + "hidden_states": hidden_states, # CPU bf16[L,D] + "metadata_json": metadata_bytes, # CPU uint8[M] +} +``` + +| field | 含义 | Consumer 中的用途 | +|---|---|---| +| `input_ids` | 经过 feature window 选择后的 token IDs | DSpark token 输入 | +| `loss_mask` | 每个 token 是否参与 loss | 排除 prompt/padding/无效位置 | +| `position_ids` | 每个 row 对应的序列位置 | 位置编码和对齐校验 | +| `hidden_states` | target 指定层 hidden 沿最后一维拼接 | DSpark target/context 特征 | +| `metadata_json` | 模型、layers、layout、shape、样本身份 | Consumer 校验并恢复 metadata | + +符号: + +- `L`:这条训练 feature 保留的 token row 数; +- `H`:target model hidden size; +- `C`:DSpark context layer 数; +- L1 关闭:`D=C*H`,layout=`dflash_aux`; +- L1 开启:`D=C*H+H`,layout=`dflash_aux_plus_last`。 + +示例:`H=4096,C=5,L=1536`,开启 L1: + +```python +input_ids.shape == [1536] +loss_mask.shape == [1536] +position_ids.shape == [1536] +hidden_states.shape == [1536, 24576] +``` + +`metadata_json` 使用 JSON UTF-8 编码为 Tensor,因为参考 PR 的 bridge 只把 Tensor 放入 TQ fields: + +```python +raw = json.dumps(metadata, sort_keys=True).encode("utf-8") +metadata_bytes = torch.tensor(list(raw), dtype=torch.uint8) +``` + +### 2.4 Tag 字段 + +```python +tag = { + "record_type": "sample", + "status": "ready", + "schema_version": 1, + "run_id": "dspark-20260818-a", + "sequence_no": 17, + "sample_id": "train-000017", + "algorithm": "DSPARK", +} +``` + +tag 只存用于发现和筛选的 flat scalar/string。Consumer 只选择: + +```text +record_type=sample +status=ready +schema_version=1 +run_id=当前 run +algorithm=DSPARK +``` + +### 2.5 Metadata 字段 + +`metadata_json` 解码后至少包含: + +```python +metadata = { + "schema_version": 1, + "run_id": "dspark-20260818-a", + "sample_id": "train-000017", + "sequence_no": 17, + "algorithm": "DSPARK", + "target_model_id": "/models/Qwen3-8B", + "target_model_revision": "revision-or-checksum", + "tokenizer_fingerprint": "sha256:...", + "target_layer_ids": [2, 8, 14, 20, 26, -1], + "hidden_states_layout": "dflash_aux_plus_last", + "hidden_dtype": "bfloat16", + "hidden_shape": [1536, 24576], + "feature_length": 1536, + "full_sequence_length": 1800, + "feature_start": 264, + "feature_end": 1800, + "use_logits": False, +} +``` + +其中: + +- `target_model_id/revision`:hidden states 来自哪个 target checkpoint; +- `tokenizer_fingerprint`:Producer 使用的 tokenizer/template 版本; +- `target_layer_ids`:vLLM 返回和参与拼接的层; +- `hidden_states_layout`:Consumer 应如何拆 hidden 最后一维; +- `feature_length`:payload 中四个主要 Tensor 的第一维; +- `full_sequence_length`:完整 prompt+response 的 token 数; +- `[feature_start,feature_end)`:feature 在完整序列中的范围。 + +### 2.6 共享协议接口 + +Producer 和 Consumer 都 import 同一个模块,不各自手写字段名: + +```python +@dataclass(frozen=True) +class SampleMetadata: + schema_version: int + run_id: str + sample_id: str + sequence_no: int + algorithm: str + target_model_id: str + target_model_revision: str + tokenizer_fingerprint: str + target_layer_ids: list[int] + hidden_states_layout: str + hidden_dtype: str + hidden_shape: list[int] + feature_length: int + full_sequence_length: int + feature_start: int + feature_end: int + use_logits: bool + +def make_sample_key(meta: SampleMetadata) -> str: ... +def make_ready_tag(meta: SampleMetadata) -> dict: ... +def encode_sample(sample, meta: SampleMetadata) -> dict[str, Tensor]: ... +def decode_sample(key, tag, fields, expected_config) -> DraftFeatureSample: ... +def make_eos_record(run_id: str, total_samples: int): ... +``` + +Producer 使用 `SampleMetadata/make_sample_key/make_ready_tag/encode_sample`;Consumer 使用 `decode_sample`。 + +`SampleMetadata` Python 对象本身不经过 TQ: + +```text +Producer SampleMetadata +→ JSON +→ uint8 Tensor +→ TQ metadata_json +→ uint8 Tensor +→ JSON +→ Consumer metadata dict +``` + +`decode_sample()` 负责: + +1. 解码 `metadata_json`; +2. 校验 key、tag、metadata 中的 sample 身份一致; +3. 校验模型、tokenizer、layers 和 layout 与 Consumer 配置一致; +4. 校验 Tensor 必需字段、dtype 和 shape; +5. 返回现有 `DraftFeatureSample`。 + +### 2.7 EOS + +Producer 完成全部输入后写一个控制 record: + +```python +eos_key = f"control:v1:{run_id}:eos" +eos_fields = {"marker": torch.tensor([1], dtype=torch.uint8)} +eos_tag = { + "record_type": "control", + "status": "eos", + "schema_version": 1, + "run_id": run_id, + "total_samples": total_samples, +} +``` + +EOS 不进入训练 batch。Consumer 看到 EOS 后继续处理剩余 ready samples;`EOS 已出现且 ready 为空` 时结束。 + +## 3. TQ 怎么启动,Producer 和 Consumer 怎么连接同一个 TQ + +### 3.1 已验证的 TQ 0.1.7 连接机制 + +`TransferQueue==0.1.7` 没有提供“把 Controller 地址直接传给第二个进程”的高层连接接口。它的无参 `tq.init()` 内部执行: + +```python +_TQ_CONTROLLER = ray.get_actor("TransferQueueController") +conf = ray.get(_TQ_CONTROLLER.get_config.remote()) +_maybe_create_tq_client(conf) +``` + +因此所有进程必须先加入同一个 Ray 集群。Ray 在本方案中只承担 TQ 控制面:保存 named `TransferQueueController`、返回 TQ 配置、管理 TQ 创建的 actor。业务进程不通过 Ray 发送 sample key 或 hidden-state tensor。 + +实际连接链路是: + +```text +TQ owner:ray.init(address) → tq.init(full_tq_config) → 创建 named Controller +Producer:ray.init(address) → tq.init() → ray.get_actor() → 创建本地 TQ client +Consumer rank 0..N:ray.init(address) → tq.init() → ray.get_actor() → 创建各自 TQ client +``` + +### 3.2 直接移植并扩展 PR #48 的 bridge + +参考文件: + +```text +C:/Users/xxyyrr/Desktop/上班/verl-SpeCo/ + verl_speco/integration/transferqueue_bridge.py +``` + +第一版不新增 `standalone_tq.py`,直接把该文件移植到目标项目并扩展。保留 PR 已有的 import 检查、进程内幂等状态、`kv_put`、`kv_batch_get`、TensorDict 解包和 owner-only shutdown;新增 Ray 显式连接、批量 list/get/clear 和本地 client 关闭接口。 + +目标文件必须实现以下函数,而不只是提供一个笼统的 transport class: + +```python +def configure_transfer_queue(config: Mapping[str, Any]) -> bool: ... +def connect_ray_cluster(ray_address: str, namespace: str | None = None) -> None: ... +def start_transfer_queue_owner(tq_config: Mapping[str, Any]) -> None: ... +def connect_transfer_queue_client() -> None: ... +def make_sample_key(run_id: str, sequence_no: int, sample_id: str) -> str: ... +def put_sample(key: str, fields: dict[str, Tensor], tag: dict[str, Any]) -> None: ... +def list_samples() -> dict[str, dict[str, Any]]: ... +def get_samples(keys: list[str]) -> list[tuple[str, dict[str, Tensor]]]: ... +def clear_samples(keys: list[str]) -> None: ... +def close_transfer_queue_client() -> None: ... +def close_transfer_queue_owner() -> None: ... +``` + +逐个函数的责任如下。 + +#### `configure_transfer_queue(config)` + +- 从 Hydra/OmegaConf 中读取 Ray address、可选 namespace、固定 partition、run ID、schema version 和 TQ native backend 配置; +- 转成普通 Python dict,保存在进程内 `_state`; +- 校验 `TransferQueue==0.1.7` 可 import; +- 不连接 Ray,不创建 TQ,不产生跨进程副作用; +- 返回该进程是否启用了 TQ。 + +#### `connect_ray_cluster(ray_address, namespace)` + +- 若 `ray.is_initialized()` 为 false,调用 `ray.init(address=ray_address, namespace=namespace)`; +- 若已经初始化,校验现有 Ray context 指向期望集群,不能悄悄连到另一个本地 Ray; +- Owner、Producer 和所有 torchrun ranks 都调用它; +- 此函数只建立当前 OS 进程到 Ray control plane 的连接。 + +#### `start_transfer_queue_owner(tq_config)` + +- 仅由 `tq_owner.py` 调用; +- 前置条件是 `connect_ray_cluster()` 已成功; +- 调用一次 `tq.init(OmegaConf.create(tq_config))`; +- 将 `_state.owner=True`、`_state.initialized=True`; +- TQ 0.1.7 会创建名为 `TransferQueueController` 的 Ray actor,并创建所选 storage backend; +- 重复调用必须报错,不能启动第二套同名 Controller。 + +#### `connect_transfer_queue_client()` + +- 由 Producer 和每个 Consumer rank 调用; +- 前置条件是当前进程已经连接 Ray; +- 调用无参 `tq.init()`,通过 `ray.get_actor("TransferQueueController")` 发现 owner; +- 只创建当前进程的 TQ client,不创建新的 Controller; +- 成功后设置 `_state.initialized=True`;重复调用直接返回。 + +#### `put_sample/list_samples/get_samples/clear_samples` + +- 全部固定使用 `_SPECO_TQ_PARTITION = "speco_drafter_features"`; +- `put_sample()` 调用单样本 `tq.kv_put()`; +- `list_samples()` 调用 `tq.kv_list(partition_id=...)`,只返回 key/tag 元数据; +- `get_samples()` 一次调用 `tq.kv_batch_get(keys=...)`,再按输入 key 顺序解包为普通 dict; +- `clear_samples()` 调用 `tq.kv_clear(keys=..., partition_id=...)`; +- 这些函数不调用 Ray RPC,不把 tensor 放入 Ray object store。 + +#### `close_transfer_queue_client()` 与 `close_transfer_queue_owner()` + +TQ 0.1.7 的公共 `tq.close()` 会 kill 共享 Controller,所以两者不能写成同一个实现: + +- `close_transfer_queue_client()`:通过 0.1.7 已公开的 `tq.get_client()` 取得当前进程 client并调用它的 `close()`,然后 `ray.shutdown()`;绝不能调用会 kill Controller 的全局 `tq.close()`。这个 client close 只用于进程退出阶段,关闭后本进程不能再次调用 TQ; +- `close_transfer_queue_owner()`:仅当 `_state.owner=True` 时调用 `tq.close()`,清理 Controller/Storage,最后 `ray.shutdown()`; +- Producer 或任一训练 rank 提前退出都不能关闭全局 TQ。 + +### 3.3 共享配置 + +Owner、Producer 和 Consumer 必须使用相同的 Ray address/namespace,并通过同一个 named Controller 获得 backend 配置。第一版 partition 固定,不按任务动态创建: + +```yaml +transfer_queue: + enable: true + package_version: "0.1.7" + ray: + address: "ray-head-node:6379" + namespace: "speco-drafter" + partition_id: "speco_drafter_features" + run_id: "dspark-20260819-a" + schema_version: 1 + backend: + storage_backend: MooncakeStore + MooncakeStore: + auto_init: false + metadata_server: "node0:50050" + master_server_address: "node0:50051" + local_hostname: "" + protocol: tcp + global_segment_size: 4294967296 + local_buffer_size: 1073741824 + device_name: "" +``` + +`partition_id/run_id/schema_version` 用于过滤和校验样本;它们不能帮助进程发现 TQ。真正让三个任务连接到同一 TQ 的是“连接同一个 Ray 集群和 namespace,然后找到同名 Controller”。 + +依赖也必须锁定并单独验证。实机安装 `TransferQueue==0.1.7` 会安装 Ray,并要求 `numpy<2.0.0`;当前测试环境中的 `twinkle-kit` 要求 `numpy>=2.0.0`,两者冲突。开发时应使用专门的 TQ/训练环境或重新确认整套依赖约束,不能直接把 0.1.7 安装进已有生产环境后忽略 resolver warning。 + +### 3.4 `tq_owner.py` 要实现的入口和函数 + +新增: + +```text +verl_speco/tq_owner.py +tools/run_dspark_tq_owner.sh +``` + +`tq_owner.py` 建议明确实现: + +```python +def install_signal_handlers(stop_event: threading.Event) -> None: ... +def publish_owner_ready(run_id: str, schema_version: int) -> None: ... +def wait_until_stopped(stop_event: threading.Event) -> None: ... +def run_owner(config: DictConfig) -> int: ... +def main() -> None: ... +``` + +`run_owner()` 的执行顺序必须是: + +```text +configure_transfer_queue(config) +→ connect_ray_cluster(ray.address, ray.namespace) +→ start_transfer_queue_owner(full TQ native config) +→ put owner_ready 控制 record +→ 安装 SIGINT/SIGTERM handler +→ 保持 owner 进程存活 +→ 收到停止信号 +→ close_transfer_queue_owner() +``` + +Owner 必须常驻。TQ 0.1.7 创建 Controller 时没有设置 `lifetime="detached"`,不能在初始化后立即退出。 + +### 3.5 启动和关闭顺序 + +第一版由外部脚本管理全生命周期: + +```text +1. ray start --head,记录 Ray address +2. 启动 Mooncake metadata/master(若 auto_init=false) +3. 启动 TQ owner;owner 连接 Ray并调用 tq.init(full config) +4. 等待 owner_ready +5. 启动一个或多个 vLLM servers +6. 启动 Consumer;每个 torchrun rank 连接 Ray,然后 tq.init() +7. 启动 Producer;连接 Ray,然后 tq.init() +8. Producer 写 EOS,关闭本地 client并退出 +9. Consumer drain、保存 final checkpoint,所有 ranks 关闭本地 client并退出 +10. 给 TQ owner 发送 SIGTERM;仅 owner 执行 tq.close() +11. 等 owner 退出后执行 ray stop +12. 停止 Mooncake 服务 +``` + +外部脚本要用 `trap` 保证异常退出也按“Producer/Consumer → owner → Ray → Mooncake”的顺序清理。不能在 Consumer rank 的 `finally` 中调用全局 `tq.close()`。 + +## 4. Producer 要实现什么 + +### 4.1 Producer 完整顺序 + +```text +读取共享配置 +→ 连接 TQ并校验 owner_ready +→ 初始化 tokenizer +→ 初始化多个 vLLM endpoint clients +→ 流式读取输入文件 +→ 为每条输入分配 sequence_no/sample_id +→ 拼接 prompt+预生成 response,得到 input_ids/loss_mask +→ 并发请求 vLLM prefill +→ 读取 vLLM hidden-state 临时结果 +→ 转换成 DSpark DraftFeatureSample +→ 构造 SampleMetadata +→ encode_sample 得到 fields/tag/key +→ TQ kv_put 一条 sample +→ 删除该请求临时文件 +→ 所有输入完成后写 EOS +→ close_transfer_queue_client()并退出 +``` + +### 4.2 并发模型 + +Producer 是一个进程,内部并发请求多个 endpoint: + +```text +InputReader +→ bounded asyncio input_queue +→ N 个 RequestWorker +→ bounded publish_queue +→ TQ Publisher +``` + +- `vllm_endpoints` 是列表; +- 每个 endpoint 有独立 semaphore; +- 总并发由 `max_inflight_requests` 限制; +- input/publish queue 必须有上限; +- TQ `kv_put` 如果是同步 API,用 `asyncio.to_thread()` 调用; +- TQ ready 数达到 `max_pending_samples` 时暂停继续请求,形成简单背压。 + +`sequence_no` 在 InputReader 中分配,不按 vLLM 完成顺序分配。因此并发乱序不会改变 sample key。 + +### 4.3 vLLM 结果转换 + +复用现有 `TargetFeatureReplayer` 的: + +- OpenAI-compatible vLLM 请求; +- `prompt_token_ids` 校验; +- `kv_transfer_params.hidden_states_path`; +- safetensors 加载; +- `[seq,layers,hidden]` 校验; +- feature positions 选择; +- aux layers flatten; +- DSpark L1 时拼 final hidden。 + +不要复制 `_feature_from_vllm_payload()`;把纯转换逻辑提取成公共函数,让旧 file backend 和新 Producer 共同调用。 + +临时文件顺序: + +```text +加载 +→ 校验/转换 +→ TQ put 成功 +→ 删除 +``` + +第一版仍允许每个并发请求产生短期临时文件,但不生成全量 hidden-state 数据集。 + +### 4.4 Producer 文件分工 + +| 文件 | 实现内容 | +|---|---| +| `verl_speco/standalone_tq_producer.py` | CLI/Hydra 入口,连接 TQ,启动 asyncio pipeline,写 EOS和汇总指标 | +| `verl_speco/producer/input_reader.py` | 流式读输入、分配 `sequence_no/sample_id`、tokenize、构造 loss mask | +| `verl_speco/producer/vllm_feature_client.py` | endpoint pool、并发控制、HTTP 请求、临时结果读取和删除 | +| `verl_speco/trainer/target_feature_replay.py` | 提取可复用的 vLLM payload → `DraftFeatureSample` 纯转换函数 | +| `examples/run_dspark_tq_producer.sh` | Producer 配置和启动命令 | +| `tests/unit/test_drafter_sample_protocol.py` | 协议编码、字段和 shape 测试(两人共同) | +| `tests/integration/test_tq_producer_smoke.py` | fake/短 vLLM → TQ sample + EOS | + +Producer 开发者同时负责移植/扩展 `integration/transferqueue_bridge.py` 和实现 `tq_owner.py`,因为这部分与 TQ 写入和连接直接相关。 + +### 4.5 Producer 各文件的函数级实现规格 + +#### `verl_speco/standalone_tq_producer.py` + +需要实现: + +```python +@dataclass +class ProducerStats: + input_count: int + published_count: int + failed_count: int + pending_bytes: int + +async def publish_one(result: PreparedFeature, transport) -> str: ... +async def run_producer(config: DictConfig) -> ProducerStats: ... +def validate_producer_config(config: DictConfig) -> None: ... +def main() -> None: ... +``` + +`main()` 只负责 Hydra/日志/退出码。`run_producer()` 是可测试的业务入口,执行:连接 Ray → 连接 TQ client → 创建 InputReader 和 vLLM client pool → 启动有界 asyncio pipeline → 等待所有请求及 `kv_put` 完成 → 发布 EOS → 关闭本地 client。它不能启动 Ray head、不能创建 TQ owner、不能调用全局 `tq.close()`。 + +`publish_one()` 接收已经完成转换的一条 `PreparedFeature`,调用共享协议的 `encode_sample()` 得到 `(key, fields, tag)`,再通过 bridge 的 `put_sample()` 发布。只有 `put_sample()` 成功返回,`published_count` 才增加,vLLM 临时文件才允许删除。 + +#### `verl_speco/producer/input_reader.py` + +需要实现: + +```python +@dataclass(frozen=True) +class InputRecord: + sequence_no: int + sample_id: str + prompt: str + response: str + source_metadata: dict[str, Any] + +def iter_input_records(path: str) -> Iterator[InputRecord]: ... +def tokenize_record(record: InputRecord, tokenizer, config) -> TokenizedRequest: ... +def build_loss_mask(input_ids: Tensor, prompt_length: int) -> Tensor: ... +``` + +`iter_input_records()` 流式读取,不把全文件载入内存;在这里按文件顺序分配稳定的 `sequence_no`。`tokenize_record()` 拼接已经存在的 prompt/response,不调用模型生成 response;输出至少包含 `input_ids:int64[L]`、`position_ids:int64[L]`、`loss_mask:float32[L]` 和请求 vLLM 所需字段。 + +#### `verl_speco/producer/vllm_feature_client.py` + +需要实现: + +```python +@dataclass(frozen=True) +class VllmEndpoint: + base_url: str + max_concurrency: int + +class VllmFeatureClientPool: + async def start(self) -> None: ... + async def prefill(self, request: TokenizedRequest) -> RawVllmFeature: ... + async def close(self) -> None: ... + +async def request_prefill(endpoint, request) -> VllmResponse: ... +def choose_endpoint(endpoints, state) -> VllmEndpoint: ... +def load_hidden_state_result(response) -> RawVllmFeature: ... +def delete_temporary_result(raw: RawVllmFeature) -> None: ... +``` + +`prefill()` 必须允许多个 coroutine 同时运行;全局 semaphore 限制总并发,每个 endpoint 另有独立 semaphore。`request_prefill()` 只负责 HTTP 请求和响应校验;`load_hidden_state_result()` 负责读取 vLLM 返回的临时 safetensors/path。临时结果的删除不放在 `load_hidden_state_result()`,而由 `publish_one()` 成功后触发。 + +#### `verl_speco/trainer/target_feature_replay.py` + +把当前类内部的纯转换部分抽成: + +```python +def feature_from_vllm_payload( + payload: RawVllmFeature, + request: TokenizedRequest, + feature_config: FeatureContract, +) -> DraftFeatureSample: ... +``` + +它不进行 HTTP、不访问 TQ、不删除文件,只完成 shape/layout/layer 校验、feature-position 选择、aux layer flatten 和可选 L1 final hidden 拼接。现有 file replay backend 和新 Producer 都调用这一函数,避免两套转换规则。 + +#### `examples/run_dspark_tq_producer.sh` + +负责提供同一套: + +```text +RAY_ADDRESS / Ray namespace +run_id / schema_version / 固定 partition +Mooncake/TQ backend 配置 +输入文件和 tokenizer/model 配置 +vLLM endpoint 列表 +max_inflight_requests / per_endpoint_concurrency +``` + +脚本只启动 Producer,不启动 TQ owner 或 Consumer,便于两位开发者独立调试。 + +## 5. Consumer 要实现什么 + +### 5.1 不新写另一套训练器 + +继续使用现有入口: + +```text +draft_train_launcher.py +→ draft_train.py +→ trainer/draft_training_loop.py +→ DrafterBaseTrainer +→ DSparkTrainerBackend +``` + +训练模式继续使用现有 standalone `training.mode=offline`,用 `feature_store.type=tq` 选择 TQ 数据源。不复制 FSDP、loss、optimizer 或 checkpoint 代码。 + +当前 `feature_store.type` 的作用是告诉 `build_feature_store_from_config()` 应创建哪一种训练数据来源: + +| 当前 type | 对象 | 数据来源 | +|---|---|---| +| `torch_shard` | `TorchShardFeatureStore` | 本地 `.pt` shard,当前 separate-training 示例使用它 | +| `token_replay` | `TokenReplayFeatureStore` | 保存 token replay 的 shard | +| `vllm_safetensors` / `safetensors` | `VllmSafetensorsFeatureStore` | 已生成的 safetensors feature | +| `jsonl_token_replay` / `jsonl` | `JsonlTokenReplayFeatureStore` | JSONL token/text replay | + +第一版新增: + +```yaml +actor_rollout_ref: + rollout: + drafter: + training: + mode: offline + feature_store: + type: tq + path: null + shuffle: false + repeat: false +``` + +这里 `offline` 的含义是“独立于 RL 的 standalone drafter training”,不等于数据必须来自磁盘。 + +不能只给现有工厂增加一个 `type=tq` 然后继续使用通用 `DraftFeatureDataLoader`。当前 loader 会: + +```python +keys = list(store.iter_keys(...)) +``` + +它假设数据集 keys 是一个静态快照;空 store 会直接结束,并且各 rank 独立枚举时可能在 Producer 持续写入的过程中看到不同快照。TQ 是流式、会新增并删除 keys 的数据源,所以 `type=tq` 必须选择专用的 `TQFeatureDataLoader`,由 rank 0 统一发现和分配 keys。 + +#### 5.1.1 当前磁盘 feature store 为什么每个 rank 都会取 keys + +`draft_train_launcher.py` 通过 `torchrun` 启动多个训练进程。每个进程都是一个 rank,并且每个 rank 都会独立进入 `run_standalone_draft_training()`、创建 `DraftFeatureDataLoader`、执行它的 `__iter__()`。当前 loader 的核心逻辑是: + +```python +keys = list( + self.store.iter_keys( + shuffle=self.shuffle, + seed=self.seed + epoch, + ) +) +rank_keys = keys[rank::world_size] + +for key in rank_keys: + batch.append(self.store.read(key)) +``` + +因此当前并不是 rank 0 枚举 keys 后再发送给其他 rank,而是所有 rank 都访问相同的静态 feature-store 路径: + +```text +rank 0:iter_keys() 得到完整静态列表 → keys[0::world_size] → read 自己的 samples +rank 1:iter_keys() 得到完整静态列表 → keys[1::world_size] → read 自己的 samples +... +``` + +例如 store 中固定存在: + +```python +keys = ["k0", "k1", "k2", "k3", "k4", "k5", "k6", "k7"] +``` + +当 `world_size=2` 时,两个 rank 都先得到上述完整列表,然后分别计算: + +```python +# rank 0 +rank_keys = keys[0::2] # ["k0", "k2", "k4", "k6"] + +# rank 1 +rank_keys = keys[1::2] # ["k1", "k3", "k5", "k7"] +``` + +这种实现能够成立,是因为 `.pt`、safetensors 或 replay 文件在训练期间是静态数据集。只要所有 rank 使用同一路径和相同 shuffle seed,`iter_keys()` 就会产生相同顺序的 key 快照,各 rank 可以无通信地算出互不重叠的子集。 + +#### 5.1.2 TQ 为什么必须改成 rank 0 发现 keys + +TQ 中的 keys 会在训练期间动态变化:Producer 持续 `kv_put`,训练成功后 Consumer 又执行 `kv_clear`。如果所有 rank 仍然各自调用 `kv_list`,不同调用时刻可能看到不同快照: + +```python +# rank 0 较早调用 +rank0_keys = ["k0", "k1", "k2", "k3"] + +# Producer 随后写入 k4、k5,rank 1 较晚调用 +rank1_keys = ["k0", "k1", "k2", "k3", "k4", "k5"] +``` + +各 rank 再独立切片后,可能得到不同数量的 local samples,进而无法保证它们以相同顺序进入 forward、backward 和梯度 collective。为此,TQ 专用 loader 必须把控制面和数据面分开: + +```text +控制面:rank 0 执行 kv_list,固定本 step 的 global_keys,并向各 rank 分发 local_keys +数据面:每个 rank 使用自己的 local_keys 直接执行 kv_batch_get,从 TQ/Mooncake 读取 hidden-state payload +``` + +rank 0 只发送较小的 key/tag 元数据,不读取并转发其他 rank 的 hidden-state tensor。所有 rank 仍然都会连接同一个 TQ,也都会调用批量读取接口;只有动态 key 的发现和本 step 的 global batch 决策集中在 rank 0。 + +因此 `feature_store.type=tq` 不是只替换底层 `read()`:它同时改变了 key 的发现、分配、等待和删除语义,需要专用 `TQFeatureDataLoader`,并要求训练循环在所有 rank 成功完成 optimizer step 后,由 rank 0 对本 step 的 `global_keys` 执行 `kv_clear`。 + +### 5.2 Consumer 完整顺序 + +```text +torchrun 启动多个 ranks +→ 每个 rank 初始化 torch.distributed +→ 每个 rank 连接同一个 TQ +→ rank 0 校验 owner_ready,并 broadcast 结果 +→ 初始化现有 DSpark trainer +→ rank 0 kv_list 查找 ready sample keys +→ rank 0 选一个 global batch并分给各 rank +→ 每个 rank kv_batch_get 自己的 local keys +→ decode_sample 得到 list[DraftFeatureSample] +→ prepare_training_batch_from_samples() +→ training_step_from_batch() +→ 所有 rank 汇总 success +→ 成功后 rank 0 kv_clear 这个 global batch 的 keys +→ 继续下一批 +→ 看到 EOS 且 ready 为空 +→ 保存 final checkpoint +→ 所有 ranks close_transfer_queue_client()并退出 +``` + +### 5.3 多 rank 如何分 key + +例如: + +```text +world_size=2 +batch_size_per_gpu=2 +global batch size=4 +``` + +rank 0 选出: + +```python +global_keys = ["k10", "k11", "k12", "k13"] +assignments = [ + ["k10", "k11"], # rank 0 + ["k12", "k13"], # rank 1 +] +``` + +通过 `dist.scatter_object_list` 或 broadcast 发送短字符串列表。每个 rank 直接从 TQ 读取自己的 payload,不通过 rank 0 转发 hidden states。 + +### 5.4 从 TQ 到训练 batch + +每个 rank: + +```python +records = tq_transport.get_samples(local_keys) + +samples = [ + decode_sample( + key=key, + tag=tags_by_key[key], + fields=fields, + expected_config=expected_contract, + ) + for key, fields in records +] + +batch = trainer.prepare_training_batch_from_samples( + samples, + step=optimizer_step, +) + +ok = await trainer.training_step_from_batch( + batch, + optimizer_step, +) +``` + +`records` 是多个独立单样本 payload;`samples` 是变长 `DraftFeatureSample` 列表。Consumer transport 层不要直接 `torch.stack()`,现有 trainer/backend 负责对齐和组 batch。 + +### 5.5 删除与结束 + +第一版采用简单逻辑: + +```text +所有 rank get/decode/train 都成功 +→ all_reduce(global_success)=True +→ rank 0 kv_clear(global_batch_keys) +``` + +任何 rank 失败都不 clear。程序报错退出,第一版不自动恢复。 + +EOS 后不足一个 global batch 的尾部,第一版使用 `drop_last=True`:记录数量,rank 0 clear 这些尾部 keys,然后正常结束。 + +checkpoint 仍使用现有 `save_interval_steps` 和 final checkpoint;第一版不保证崩溃后已 clear 数据能够严格重放。 + +### 5.6 Consumer 文件分工 + +| 文件 | 实现内容 | +|---|---| +| `verl_speco/trainer/feature_store.py` | 工厂增加 `type=tq`,返回 `TQFeatureStore`;TQ 类型不要求 `path` | +| `verl_speco/trainer/tq_feature_store.py` | 实现固定 partition 上的 list/get-many/clear/EOS,不伪装静态 `iter_keys()` | +| `verl_speco/trainer/tq_sample_source.py` | 实现专用 `TQFeatureDataLoader`:rank 0 动态 list/filter/sort、key 分配、各 rank get/decode、EOS/tail | +| `verl_speco/trainer/draft_training_loop.py` | 保持 `mode=offline`;`type=tq` 时跳过 `feature_store.path` 必填检查并选择专用 loader;调用现有 prepare/train,global success 后 clear,final checkpoint | +| `verl_speco/draft_train_launcher.py` | `feature_store.type=tq` 时不要求 feature-store path,透传 torchrun/TQ 配置 | +| `verl_speco/config/speco_base.yaml` | 新增 TQ connection、run、poll、batch 等默认配置 | +| `tools/run_dspark_tq_consumer.sh` | Consumer GPU、batch、checkpoint 和共享 TQ 配置 | +| `tests/unit/test_tq_sample_source.py` | key 过滤/排序/分 rank、EOS、decode 调用测试 | +| `tests/integration/test_dspark_tq_consumer_smoke.py` | 两 rank 取不同样本,完成训练并 clear | + +### 5.7 Consumer 各文件的函数级实现规格 + +#### `verl_speco/trainer/feature_store.py` + +修改现有工厂: + +```python +def build_feature_store_from_config(feature_store_cfg, read_only=False): + store_type = str(feature_store_cfg.get("type", "torch_shard")).lower() + if store_type == "tq": + return TQFeatureStore.from_config(feature_store_cfg) + ... +``` + +要求: + +- 保留 `torch_shard/token_replay/vllm_safetensors/jsonl_token_replay` 的当前行为; +- `type=tq` 时不读取 `path`; +- 工厂只创建对象,不在 import 阶段连接 Ray/TQ; +- `TQFeatureStore` 是流式数据源适配器,不强行实现有误导性的静态 `iter_keys()` 和逐条 `read()`。 + +#### `verl_speco/trainer/tq_feature_store.py` + +需要实现: + +```python +@dataclass(frozen=True) +class ReadyEntry: + key: str + tag: dict[str, Any] + +class TQFeatureStore: + @classmethod + def from_config(cls, cfg) -> "TQFeatureStore": ... + def connect(self) -> None: ... + def list_ready(self, run_id: str) -> list[ReadyEntry]: ... + def get_many(self, entries: list[ReadyEntry]) -> list[DraftFeatureSample]: ... + def clear_many(self, keys: list[str]) -> None: ... + def read_eos(self, run_id: str) -> EosMetadata | None: ... + def close_local(self) -> None: ... +``` + +`connect()` 调用共享 bridge 的 `connect_ray_cluster()` 和 `connect_transfer_queue_client()`。`list_ready()` 调用 `list_samples()` 后只保留 `tag.status=ready`、匹配 `run_id/schema_version` 的数据 key,并按 `(sequence_no,key)` 排序。`get_many()` 一次批量读取 fields,然后逐条调用共享 `decode_sample()`;返回顺序必须和 entries 相同。`clear_many()` 只允许 rank 0 在全局 step 成功后调用。`close_local()` 只关闭本 rank client,不能关闭 Controller。 + +#### `verl_speco/trainer/tq_sample_source.py` + +需要实现: + +```python +@dataclass +class TQLocalBatch: + local_keys: list[str] + local_samples: list[DraftFeatureSample] + global_keys: list[str] | None + +class TQFeatureDataLoader: + def __iter__(self) -> Iterator[TQLocalBatch]: ... + def _select_global_batch(self) -> tuple[list[str], dict[str, ReadyEntry]]: ... + def _broadcast_assignments(self, assignments) -> list[ReadyEntry]: ... + def _handle_eos_and_tail(self) -> bool: ... + def clear_completed_batch(self, global_keys: list[str] | None) -> None: ... +``` + +执行责任必须明确: + +- 所有 rank 创建 loader 并调用 `store.connect()`; +- 只有 rank 0 执行 `_select_global_batch()` 和 `kv_list`; +- rank 0 生成 `list[list[ReadyEntry]]` assignments,使用 `torch.distributed.scatter_object_list` 或 broadcast 分发小型 key/tag; +- 每个 rank 对自己的 local entries 调用 `store.get_many()`,hidden states 从 TQ/Mooncake 直接进入该 rank; +- loader yield `TQLocalBatch`,不能只 yield samples,因为训练后清理还需要 `global_keys`; +- 暂时为空时轮询等待,不能像静态 loader 一样结束;只有 EOS 已出现并且 ready 数据 drain 完才停止; +- `drop_last=True` 时由 rank 0 记录并清理不足 global batch 的尾部 key。 + +#### `verl_speco/trainer/draft_training_loop.py` + +需要新增或调整: + +```python +def build_training_source(config, rank, world_size): ... +def all_ranks_succeeded(local_ok: bool, device) -> bool: ... +async def run_tq_training_loop(trainer, loader, config) -> dict[str, Any]: ... +``` + +`build_training_source()` 根据 `feature_store.type` 分支:静态类型继续创建 `DraftFeatureDataLoader`;`tq` 创建 `TQFeatureDataLoader`,跳过 `feature_store.path` 必填检查,并禁止 `shuffle/repeat`。`run_tq_training_loop()` 对每个 `TQLocalBatch` 调用现有 `prepare_training_batch_from_samples()` 和 `training_step_from_batch()`;所有 rank 通过 collective 得到 global success 后,才让 rank 0 调用 `clear_completed_batch(global_keys)`。异常路径不 clear,finally 只执行 `store.close_local()`。 + +#### `verl_speco/draft_train_launcher.py` + +保留当前 `torch.distributed.run` 启动方式,只增加配置校验和环境透传: + +```python +def validate_tq_launch_config(overrides, launch_config) -> None: ... +def build_child_env(config) -> dict[str, str]: ... +``` + +它不启动 Ray head、不调用 `tq.init()`。职责是确认 `type=tq` 时提供了 Ray address/namespace,且不要求 feature-store path;随后把同一 Ray connection 配置传给每个 torchrun 子进程。每个子 rank 自己建立 Ray/TQ client,不能在 launcher 父进程建一个 client后期待 fork 继承。 + +#### `verl_speco/config/speco_base.yaml` + +增加默认字段: + +```yaml +feature_store: + type: torch_shard + path: null + shuffle: true + repeat: true + tq: + ray_address: null + ray_namespace: speco-drafter + partition_id: speco_drafter_features + run_id: null + schema_version: 1 + poll_interval_seconds: 0.5 + connect_timeout_seconds: 120 + drop_last: true +``` + +当 `type=tq` 时运行期覆盖为 `shuffle=false/repeat=false`。backend 的完整 owner 配置只需要 TQ owner 使用;Producer/Consumer client 从 named Controller 获取它,不各自重新创建 backend。 + +#### Consumer 测试必须覆盖的函数边界 + +- `test_feature_store_factory_builds_tq_without_path()`; +- `test_rank0_filters_and_sorts_ready_entries()`; +- `test_nonzero_rank_never_calls_kv_list()`; +- `test_assignments_are_disjoint_and_global_batch_complete()`; +- `test_each_rank_gets_only_local_keys()`; +- `test_decode_preserves_hidden_states_layout()`; +- `test_clear_only_after_all_ranks_success()`; +- `test_failure_does_not_clear()`; +- `test_eos_drains_ready_then_stops()`; +- `test_client_close_does_not_kill_owner()`。 + +## 6. 两个人怎么分工 + +### 共同先完成 + +1. `drafter_sample_protocol.py`; +2. Ray/TQ connection 配置字段; +3. 一个小型 golden sample; +4. 启动一个测试 Ray head 和 TQ owner,独立进程 A put、独立进程 B list/get/clear 的 smoke test; +5. 验证 client 退出不会 kill owner,只有 owner shutdown 才销毁 Controller。 + +### 开发者 A:Producer/TQ + +负责: + +```text +integration/transferqueue_bridge.py +tq_owner.py +standalone_tq_producer.py +producer/input_reader.py +producer/vllm_feature_client.py +target_feature_replay.py 的公共转换函数 +owner/producer 启动脚本 +Producer/TQ 测试 +``` + +开发者 A 的可交付接口不是“提供一个 TQ 类”,而是: + +```text +bridge:connect_ray_cluster/start_owner/connect_client/put/list/get_many/clear/client_close/owner_close +owner:run_owner/main/signal handler/owner_ready +producer:run_producer/publish_one/统计与 EOS +input reader:iter_input_records/tokenize_record/build_loss_mask +vLLM client:endpoint pool/request_prefill/load/delete +feature conversion:feature_from_vllm_payload +``` + +开发者 B 可以先针对这些接口写 fake transport,不需要等待真实 vLLM 和 Mooncake 联通。 + +### 开发者 B:Consumer/训练 + +负责: + +```text +feature_store.py 的 type=tq 工厂分支 +tq_feature_store.py +tq_sample_source.py / TQFeatureDataLoader +draft_training_loop.py 的 offline + type=tq 分支 +draft_train_launcher.py 配置适配 +speco_base.yaml Consumer 配置 +Consumer 启动脚本 +Consumer/DSpark 测试 +``` + +开发者 B 的可交付接口是: + +```text +feature-store factory:type=tq 分支 +TQFeatureStore:connect/list_ready/get_many/clear_many/read_eos/close_local +TQFeatureDataLoader:select global batch/distribute local keys/get/yield/EOS-tail +training loop:build source/train/global success/clear/final checkpoint +launcher:TQ 配置校验和 torchrun 子进程环境透传 +``` + +### 联调入口 + +建议再提供: + +```text +examples/run_dspark_tq_pipeline_local.sh +``` + +只用于单机联调,顺序启动: + +```text +ray start --head +→ Mooncake metadata/master +→ TQ owner(ray.init + tq.init(full config)) +→ owner_ready +→ vLLM health check +→ Consumer +→ Producer +→ 等 Producer/Consumer 退出 +→ SIGTERM TQ owner(owner 执行 tq.close) +→ ray stop +→ 停止 Mooncake +``` + +最小联调:Producer 发布 8 条样本,2 个 Consumer ranks、每 rank batch size 2,完成 2 个 optimizer steps,8 个 sample keys 被清理,EOS 后保存 final checkpoint。 + +## 7. 第一版验收标准 + +1. TQ owner、Producer 和所有 Consumer ranks 日志显示相同 Ray address/namespace、TQ Controller、partition 和 run ID。 +2. 在同一 Ray 集群中,独立 owner 创建 TQ 后,独立进程 A put,独立进程 B 能 list/get/clear。 +3. Producer 对多个 vLLM endpoints 并发请求,不串行访问。 +4. 一个输入 record 只生成一个 sample key 和一个 payload。 +5. `kv_list` 只拿 tag;`kv_batch_get` 才拿 hidden states。 +6. Consumer 各 rank 读取不同 local keys,不通过 rank 0 搬运 Tensor。 +7. `decode_sample` 能拒绝模型、layer、layout、dtype 或 shape 不匹配的数据。 +8. 所有 rank 训练成功后才 clear 当前 global batch。 +9. Producer 先完成时,Consumer 能 drain 后再退出。 +10. 不产生长期 hidden-state feature store。 +11. Producer 或任一 Consumer rank 退出不会销毁 TQ Controller;只有 owner 调用全局 `tq.close()`。 +12. Ray object store 中不承载 hidden-state payload,训练 tensor 通过 TQ/Mooncake 路径读取。 + +## 8. 后续建议:第一版跑通后再做 + +以下内容不进入第一版开发: + +- Producer HTTP/TQ 复杂重试和 endpoint 熔断; +- Producer 发布 journal,避免重启后重复生成已 clear 样本; +- 自动重启;第一版失败后人工停止整条 pipeline并使用新 `run_id` 重跑; +- Consumer 从最新 checkpoint 自动恢复; +- checkpoint 成功后再 clear 的严格提交窗口; +- TQ owner/storage 整体丢失后的数据重建; +- 多个独立 Consumer 竞争同一 partition; +- lease、ack、超时回收和 exactly-once; +- 动态扩缩容; +- vLLM server 直接写 TQ。 + +第一版先保证:同一个 TQ 能连通、Producer 能并发生产、Consumer 能正确取数训练、每步成功后能及时清理。 diff --git a/tests/unit/test_draft_train_launcher.py b/tests/unit/test_draft_train_launcher.py index 654a0c1f..6894d165 100644 --- a/tests/unit/test_draft_train_launcher.py +++ b/tests/unit/test_draft_train_launcher.py @@ -19,6 +19,7 @@ build_torch_distributed_command, normalize_training_args, resolve_launch_config, + validate_tq_launch_config, ) @@ -110,3 +111,46 @@ def test_launcher_rejects_standalone_multinode() -> None: "speco.draft_training.standalone=true", ] ) + + +def test_launcher_accepts_complete_tq_consumer_config() -> None: + validate_tq_launch_config( + [ + "actor_rollout_ref.rollout.drafter.training.feature_store.type=tq", + "actor_rollout_ref.rollout.drafter.training.transfer_queue.enable=true", + "actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.address=ray:6379", + "actor_rollout_ref.rollout.drafter.training.transfer_queue.run_id=run-a", + ] + ) + + +@pytest.mark.parametrize( + ("overrides", "message"), + [ + ([], "enable=true"), + ( + [ + "actor_rollout_ref.rollout.drafter.training.transfer_queue.enable=true", + "actor_rollout_ref.rollout.drafter.training.transfer_queue.run_id=run-a", + ], + "ray.address", + ), + ( + [ + "actor_rollout_ref.rollout.drafter.training.transfer_queue.enable=true", + "actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.address=ray:6379", + ], + "run_id", + ), + ], +) +def test_launcher_tq_validation_requires_canonical_connection_overrides( + overrides, message +) -> None: + with pytest.raises(ValueError, match=message): + validate_tq_launch_config( + [ + "actor_rollout_ref.rollout.drafter.training.feature_store.type=tq", + *overrides, + ] + ) diff --git a/tests/unit/test_draft_training_loop.py b/tests/unit/test_draft_training_loop.py index d4abdfd3..684d0aec 100644 --- a/tests/unit/test_draft_training_loop.py +++ b/tests/unit/test_draft_training_loop.py @@ -25,6 +25,8 @@ from verl_speco.trainer.draft_training_loop import ( # noqa: E402 _build_backend, + _clear_tq_batch_across_ranks, + _connect_tq_store_across_ranks, _contains_replay_samples, _is_out_of_memory_error, _next_batch_across_ranks, @@ -51,6 +53,62 @@ def _save_checkpoint_async(self, step: int): return self.future +class _FakeTQLoader: + def __init__(self, error: BaseException | None = None): + self.error = error + self.clear_calls: list[list[str] | None] = [] + + def clear_completed_batch(self, keys): + self.clear_calls.append(keys) + if self.error is not None: + raise self.error + + +class _FakeTQStore: + def __init__(self, error: BaseException | None = None): + self.error = error + self.connect_calls = 0 + + def connect(self): + self.connect_calls += 1 + if self.error is not None: + raise self.error + + +def test_tq_completed_batch_is_cleared_once_on_rank_zero() -> None: + loader = _FakeTQLoader() + _clear_tq_batch_across_ranks( + loader, + ["k0", "k1"], + rank=0, + device=torch.device("cpu"), + ) + assert loader.clear_calls == [["k0", "k1"]] + + +def test_tq_clear_failure_is_reported_and_not_retried() -> None: + loader = _FakeTQLoader(RuntimeError("clear failed")) + with pytest.raises(RuntimeError, match="failed to clear"): + _clear_tq_batch_across_ranks( + loader, + ["k0", "k1"], + rank=0, + device=torch.device("cpu"), + ) + assert loader.clear_calls == [["k0", "k1"]] + + +def test_tq_store_connection_failure_is_reported() -> None: + store = _FakeTQStore(RuntimeError("connect failed")) + with pytest.raises(RuntimeError, match="failed to connect"): + _connect_tq_store_across_ranks( + store, + rank=0, + device=torch.device("cpu"), + ) + assert store.connect_calls == 1 + + def _export_trainer(model_type: str, model_path=None): """Minimal trainer stand-in for the standalone checkpoint export helpers.""" return SimpleNamespace( diff --git a/tests/unit/test_tq_consumer.py b/tests/unit/test_tq_consumer.py new file mode 100644 index 00000000..bce05a96 --- /dev/null +++ b/tests/unit/test_tq_consumer.py @@ -0,0 +1,265 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from __future__ import annotations + +from types import SimpleNamespace + +import pytest +import torch + +from verl_speco.trainer.feature_store import ( + DraftFeatureSample, + build_feature_store_from_config, +) +from verl_speco.trainer.tq_feature_store import ReadyEntry, TQFeatureStore +from verl_speco.trainer.tq_sample_source import ( + TQFeatureDataLoader, + build_assignments, +) +from verl_speco.transport.drafter_sample_protocol import ( + SampleMetadata, + encode_sample, + make_ready_tag, + make_sample_key, +) + + +def _config() -> dict: + return { + "enable": True, + "ray": {"address": "ray-head:6379", "namespace": "speco-drafter"}, + "partition_id": "speco_drafter_features", + "run_id": "run-a", + "schema_version": 1, + } + + +def _metadata(sequence_no: int = 0) -> SampleMetadata: + return SampleMetadata( + schema_version=1, + run_id="run-a", + sample_id=f"sample-{sequence_no}", + sequence_no=sequence_no, + algorithm="DSPARK", + target_model_id="producer-model-is-not-strictly-checked", + target_model_revision="producer-revision", + tokenizer_fingerprint="producer-tokenizer", + target_layer_ids=[2, 8, 14], + hidden_states_layout="dflash_aux_plus_last", + hidden_dtype="float32", + hidden_shape=[3, 4], + feature_length=3, + full_sequence_length=3, + feature_start=0, + feature_end=3, + use_logits=False, + ) + + +def _sample() -> DraftFeatureSample: + return DraftFeatureSample( + algorithm="DSPARK", + input_ids=torch.tensor([10, 11, 12]), + loss_mask=torch.tensor([1.0, 1.0, 1.0]), + position_ids=torch.tensor([0, 1, 2]), + hidden_states=torch.arange(12, dtype=torch.float32).reshape(3, 4), + ) + + +def _entry(sequence_no: int) -> ReadyEntry: + meta = _metadata(sequence_no) + return ReadyEntry(key=make_sample_key(meta), tag=make_ready_tag(meta)) + + +def test_feature_store_factory_builds_tq_without_path() -> None: + store = build_feature_store_from_config( + {"type": "tq", "path": None}, + read_only=True, + transfer_queue_cfg=_config(), + ) + assert isinstance(store, TQFeatureStore) + assert store.run_id == "run-a" + + +def test_tq_store_connect_filter_sort_and_minimal_decode(monkeypatch) -> None: + import verl_speco.trainer.tq_feature_store as module + + calls: list[tuple] = [] + monkeypatch.setattr(module, "configure_transfer_queue", lambda cfg: True) + monkeypatch.setattr( + module, + "connect_ray_cluster", + lambda address, namespace: calls.append(("ray", address, namespace)), + ) + monkeypatch.setattr( + module, + "connect_transfer_queue_client", + lambda: calls.append(("tq",)), + ) + entries = [_entry(2), _entry(1)] + unrelated = ReadyEntry( + key="other", + tag={ + **entries[0].tag, + "run_id": "another-run", + }, + ) + monkeypatch.setattr( + module, + "list_samples", + lambda: { + entries[0].key: entries[0].tag, + unrelated.key: unrelated.tag, + entries[1].key: entries[1].tag, + "control:v1:run-a:eos": { + "record_type": "control", + "status": "eos", + "schema_version": 1, + "run_id": "run-a", + "total_samples": 2, + }, + }, + ) + fields_by_key = { + entry.key: encode_sample(_sample(), _metadata(int(entry.tag["sequence_no"]))) + for entry in entries + } + monkeypatch.setattr( + module, + "get_samples", + lambda keys: [(key, fields_by_key[key]) for key in keys], + ) + + store = TQFeatureStore.from_config(_config()) + store.connect() + ready = store.list_ready() + samples = store.get_many(ready) + + assert calls == [("ray", "ray-head:6379", "speco-drafter"), ("tq",)] + assert [entry.tag["sequence_no"] for entry in ready] == [1, 2] + assert len(samples) == 2 + # Model/revision/tokenizer/layers are intentionally not in the first + # Consumer contract; tensor and run/protocol checks still execute. + assert samples[0].metadata["target_model_revision"] == "producer-revision" + eos = store.read_eos() + assert eos is not None and eos.total_samples == 2 + + +def test_build_assignments_is_disjoint_and_complete() -> None: + entries = [_entry(index) for index in range(4)] + assignments = build_assignments(entries, batch_size=2, world_size=2) + assert [[entry.key for entry in rank] for rank in assignments] == [ + [entries[0].key, entries[1].key], + [entries[2].key, entries[3].key], + ] + + +class _FakeStore: + def __init__(self, entries: list[ReadyEntry], *, eos: bool = True): + self.entries = list(entries) + self.eos = eos + self.connected = False + self.get_calls: list[list[str]] = [] + self.clear_calls: list[list[str]] = [] + + def connect(self) -> None: + self.connected = True + + def owner_ready(self) -> bool: + return True + + def list_ready(self): + return list(self.entries) + + def read_eos(self): + return SimpleNamespace(total_samples=len(self.entries)) if self.eos else None + + def get_many(self, entries): + self.get_calls.append([entry.key for entry in entries]) + return [_sample() for _ in entries] + + def clear_many(self, keys): + normalized = list(keys) + self.clear_calls.append(normalized) + selected = set(normalized) + self.entries = [entry for entry in self.entries if entry.key not in selected] + + +def test_world_size_one_streams_trains_then_clears_and_stops() -> None: + entries = [_entry(0), _entry(1)] + store = _FakeStore(entries) + loader = TQFeatureDataLoader( + store, batch_size=2, rank=0, world_size=1, poll_interval_seconds=0.01 + ) + iterator = iter(loader) + batch = next(iterator) + assert batch.local_keys == [entry.key for entry in entries] + assert batch.global_keys == batch.local_keys + assert store.get_calls == [batch.local_keys] + + loader.clear_completed_batch(batch.global_keys) + with pytest.raises(StopIteration): + next(iterator) + assert store.clear_calls == [batch.local_keys] + + +def test_eos_drops_incomplete_global_tail() -> None: + entry = _entry(0) + store = _FakeStore([entry]) + loader = TQFeatureDataLoader(store, batch_size=2, rank=0, world_size=1) + assert list(loader) == [] + assert store.clear_calls == [[entry.key]] + + +def test_rank0_discovery_failure_is_raised_without_clearing() -> None: + store = _FakeStore([_entry(0)]) + + def _fail_list(): + raise RuntimeError("list failed") + + store.list_ready = _fail_list + loader = TQFeatureDataLoader(store, batch_size=1, rank=0, world_size=1) + with pytest.raises(RuntimeError, match="list failed"): + next(iter(loader)) + assert store.clear_calls == [] + + +def test_nonzero_rank_uses_broadcast_assignment_without_listing(monkeypatch) -> None: + import verl_speco.trainer.tq_sample_source as module + + entries = [_entry(0), _entry(1)] + store = _FakeStore(entries, eos=False) + + def _broadcast(payload, src): + assert src == 0 + payload[0] = { + "kind": "batch", + "global_keys": [entry.key for entry in entries], + "assignments": [ + [{"key": entries[0].key, "tag": entries[0].tag}], + [{"key": entries[1].key, "tag": entries[1].tag}], + ], + } + + monkeypatch.setattr(module.dist, "is_initialized", lambda: True) + monkeypatch.setattr(module.dist, "get_world_size", lambda: 2) + monkeypatch.setattr(module.dist, "broadcast_object_list", _broadcast) + store.list_ready = lambda: pytest.fail("nonzero rank must not list TQ keys") + + loader = TQFeatureDataLoader(store, batch_size=1, rank=1, world_size=2) + batch = next(iter(loader)) + assert batch.local_keys == [entries[1].key] + assert batch.global_keys is None + assert store.get_calls == [[entries[1].key]] diff --git a/tools/run_dspark_tq_consumer.sh b/tools/run_dspark_tq_consumer.sh new file mode 100644 index 00000000..24e1d1d8 --- /dev/null +++ b/tools/run_dspark_tq_consumer.sh @@ -0,0 +1,48 @@ +set -euo pipefail +set -x + +# Standalone DSpark Consumer. Start Ray and verl_speco.tq_owner first, then run +# the Producer with the same Ray address, namespace, partition and run ID. +MODEL_PATH=${MODEL_PATH:-/path/to/Qwen3-8B} +DRAFTER_PATH=${DRAFTER_PATH:-/path/to/dspark-drafter} +DRAFT_CKPTS_DIR=${DRAFT_CKPTS_DIR:-/path/to/dspark-tq-checkpoints} +TRAIN_DEVICES=${TRAIN_DEVICES:-0,1,2,3} +TRAIN_GPUS=${TRAIN_GPUS:-4} +BATCH_SIZE_PER_GPU=${BATCH_SIZE_PER_GPU:-2} +MAX_STEPS=${MAX_STEPS:-1000} +DSPARK_NUM_TARGET_LAYERS=${DSPARK_NUM_TARGET_LAYERS:-5} +RAY_ADDRESS=${RAY_ADDRESS:-127.0.0.1:6379} +TQ_NAMESPACE=${TQ_NAMESPACE:-speco-drafter} +TQ_PARTITION_ID=${TQ_PARTITION_ID:-speco_drafter_features} +SPECO_TQ_RUN_ID=${SPECO_TQ_RUN_ID:-dspark-standalone-run} + +CUDA_VISIBLE_DEVICES=${TRAIN_DEVICES} PYTHONUNBUFFERED=1 \ +python3 -m verl_speco.draft_train_launcher \ + speco.draft_training.nproc_per_node=${TRAIN_GPUS} \ + speco.draft_training.nnodes=1 \ + actor_rollout_ref.model.path=${MODEL_PATH} \ + actor_rollout_ref.actor.strategy=fsdp2 \ + actor_rollout_ref.rollout.drafter.enable=True \ + actor_rollout_ref.rollout.drafter.enable_drafter_training=True \ + actor_rollout_ref.rollout.drafter.model_path=${DRAFTER_PATH} \ + actor_rollout_ref.rollout.drafter.checkpoint_path=${DRAFT_CKPTS_DIR} \ + actor_rollout_ref.rollout.drafter.speculative_algorithm=DSPARK \ + actor_rollout_ref.rollout.drafter.training.mode=offline \ + actor_rollout_ref.rollout.drafter.training.feature_store.type=tq \ + actor_rollout_ref.rollout.drafter.training.feature_store.path=null \ + actor_rollout_ref.rollout.drafter.training.feature_store.shuffle=False \ + actor_rollout_ref.rollout.drafter.training.feature_store.repeat=False \ + actor_rollout_ref.rollout.drafter.training.transfer_queue.enable=True \ + actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.address=${RAY_ADDRESS} \ + actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.namespace=${TQ_NAMESPACE} \ + actor_rollout_ref.rollout.drafter.training.transfer_queue.partition_id=${TQ_PARTITION_ID} \ + actor_rollout_ref.rollout.drafter.training.transfer_queue.run_id=${SPECO_TQ_RUN_ID} \ + actor_rollout_ref.rollout.drafter.training.transfer_queue.drop_last=True \ + actor_rollout_ref.rollout.drafter.training.batch_size_per_gpu=${BATCH_SIZE_PER_GPU} \ + actor_rollout_ref.rollout.drafter.training.max_steps=${MAX_STEPS} \ + actor_rollout_ref.rollout.drafter.training.dspark_num_target_layers=${DSPARK_NUM_TARGET_LAYERS} \ + actor_rollout_ref.rollout.drafter.training.save_interval_steps=100 \ + actor_rollout_ref.rollout.drafter.training.lr=1e-5 \ + actor_rollout_ref.rollout.drafter.training.lr_warmup_steps=50 \ + actor_rollout_ref.rollout.drafter.training.warmup_style=cosine \ + "$@" diff --git a/tools/run_dspark_tq_consumer_test.sh b/tools/run_dspark_tq_consumer_test.sh new file mode 100644 index 00000000..370d0be8 --- /dev/null +++ b/tools/run_dspark_tq_consumer_test.sh @@ -0,0 +1,84 @@ +#!/usr/bin/env bash +# Exercise the real standalone DSpark Consumer with delayed synthetic TQ data. +set -euo pipefail + +: "${MODEL_PATH:?Set MODEL_PATH to the target model directory}" +: "${DRAFTER_PATH:?Set DRAFTER_PATH to the DSpark drafter directory}" + +PYTHON_BIN=${PYTHON_BIN:-python3} +RAY_ADDRESS=${RAY_ADDRESS:-127.0.0.1:6379} +TQ_NAMESPACE=${TQ_NAMESPACE:-speco-drafter} +TQ_PARTITION_ID=${TQ_PARTITION_ID:-speco_drafter_features} +SPECO_TQ_RUN_ID=${SPECO_TQ_RUN_ID:-dspark-consumer-test-$$} +TRAIN_DEVICES=${TRAIN_DEVICES:-0} +TRAIN_GPUS=${TRAIN_GPUS:-1} +BATCH_SIZE_PER_GPU=${BATCH_SIZE_PER_GPU:-1} +NUM_BATCHES=${NUM_BATCHES:-3} +SEQUENCE_LENGTH=${SEQUENCE_LENGTH:-64} +DSPARK_NUM_TARGET_LAYERS=${DSPARK_NUM_TARGET_LAYERS:-5} +INITIAL_DELAY_SECONDS=${INITIAL_DELAY_SECONDS:-5} +BATCH_INTERVAL_SECONDS=${BATCH_INTERVAL_SECONDS:-5} +DRAFT_CKPTS_DIR=${DRAFT_CKPTS_DIR:-/tmp/speco-dspark-tq-consumer-test-${SPECO_TQ_RUN_ID}} + +owner_pid="" +producer_pid="" + +cleanup() { + if [[ -n "${producer_pid}" ]] && kill -0 "${producer_pid}" 2>/dev/null; then + kill "${producer_pid}" 2>/dev/null || true + wait "${producer_pid}" 2>/dev/null || true + fi + if [[ -n "${owner_pid}" ]] && kill -0 "${owner_pid}" 2>/dev/null; then + kill "${owner_pid}" 2>/dev/null || true + wait "${owner_pid}" 2>/dev/null || true + fi +} +trap cleanup EXIT INT TERM + +echo "[1/4] Starting TQ owner run_id=${SPECO_TQ_RUN_ID}" +RAY_ADDRESS="${RAY_ADDRESS}" TQ_NAMESPACE="${TQ_NAMESPACE}" \ +TQ_PARTITION_ID="${TQ_PARTITION_ID}" SPECO_TQ_RUN_ID="${SPECO_TQ_RUN_ID}" \ + bash tools/run_dspark_tq_owner.sh & +owner_pid=$! + +echo "[2/4] Starting delayed synthetic Producer" +"${PYTHON_BIN}" tools/tq_delayed_test_producer.py \ + --model-path "${MODEL_PATH}" \ + --ray-address "${RAY_ADDRESS}" \ + --namespace "${TQ_NAMESPACE}" \ + --partition-id "${TQ_PARTITION_ID}" \ + --run-id "${SPECO_TQ_RUN_ID}" \ + --world-size "${TRAIN_GPUS}" \ + --batch-size-per-gpu "${BATCH_SIZE_PER_GPU}" \ + --num-batches "${NUM_BATCHES}" \ + --sequence-length "${SEQUENCE_LENGTH}" \ + --num-target-layers "${DSPARK_NUM_TARGET_LAYERS}" \ + --initial-delay "${INITIAL_DELAY_SECONDS}" \ + --batch-interval "${BATCH_INTERVAL_SECONDS}" & +producer_pid=$! + +echo "[3/4] Running the real standalone Consumer; it should wait between batches" +MODEL_PATH="${MODEL_PATH}" \ +DRAFTER_PATH="${DRAFTER_PATH}" \ +DRAFT_CKPTS_DIR="${DRAFT_CKPTS_DIR}" \ +TRAIN_DEVICES="${TRAIN_DEVICES}" \ +TRAIN_GPUS="${TRAIN_GPUS}" \ +BATCH_SIZE_PER_GPU="${BATCH_SIZE_PER_GPU}" \ +MAX_STEPS="$((NUM_BATCHES + 1))" \ +DSPARK_NUM_TARGET_LAYERS="${DSPARK_NUM_TARGET_LAYERS}" \ +RAY_ADDRESS="${RAY_ADDRESS}" \ +TQ_NAMESPACE="${TQ_NAMESPACE}" \ +TQ_PARTITION_ID="${TQ_PARTITION_ID}" \ +SPECO_TQ_RUN_ID="${SPECO_TQ_RUN_ID}" \ + bash tools/run_dspark_tq_consumer.sh "$@" + +wait "${producer_pid}" +producer_pid="" + +echo "[4/4] Consumer observed EOS and exited; stopping this test's TQ owner" +kill "${owner_pid}" 2>/dev/null || true +wait "${owner_pid}" 2>/dev/null || true +owner_pid="" +trap - EXIT INT TERM + +echo "DSPARK_TQ_CONSUMER_TEST_OK run_id=${SPECO_TQ_RUN_ID} batches=${NUM_BATCHES}" diff --git a/examples/run_dspark_tq_owner.sh b/tools/run_dspark_tq_owner.sh similarity index 83% rename from examples/run_dspark_tq_owner.sh rename to tools/run_dspark_tq_owner.sh index b94d38b7..d1c0f245 100644 --- a/examples/run_dspark_tq_owner.sh +++ b/tools/run_dspark_tq_owner.sh @@ -16,9 +16,12 @@ set -euo pipefail : "${RAY_ADDRESS:?Set RAY_ADDRESS to the running Ray head, for example 10.0.0.1:6379}" : "${SPECO_TQ_RUN_ID:?Set SPECO_TQ_RUN_ID to a unique pipeline run id}" +TQ_NAMESPACE=${TQ_NAMESPACE:-speco-drafter} +TQ_PARTITION_ID=${TQ_PARTITION_ID:-speco_drafter_features} python -m verl_speco.tq_owner \ actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.address="${RAY_ADDRESS}" \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.namespace=speco-drafter \ + actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.namespace="${TQ_NAMESPACE}" \ + actor_rollout_ref.rollout.drafter.training.transfer_queue.partition_id="${TQ_PARTITION_ID}" \ actor_rollout_ref.rollout.drafter.training.transfer_queue.run_id="${SPECO_TQ_RUN_ID}" \ actor_rollout_ref.rollout.drafter.training.transfer_queue.backend.storage_backend=SimpleStorage diff --git a/tools/run_tq_connection_smoke.sh b/tools/run_tq_connection_smoke.sh new file mode 100644 index 00000000..e7173b21 --- /dev/null +++ b/tools/run_tq_connection_smoke.sh @@ -0,0 +1,53 @@ +#!/usr/bin/env bash +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +set -euo pipefail + +# A Ray head must already be running. This script starts only the two smoke +# roles: an owner that also publishes two synthetic samples, and a client that +# reads, validates, clears, and observes EOS. +PYTHON_BIN=${PYTHON_BIN:-python3} +RAY_ADDRESS=${RAY_ADDRESS:-127.0.0.1:6379} +TQ_NAMESPACE=${TQ_NAMESPACE:-speco-drafter-smoke} +SPECO_TQ_RUN_ID=${SPECO_TQ_RUN_ID:-tq-smoke-$$} +SMOKE_TIMEOUT_SECONDS=${SMOKE_TIMEOUT_SECONDS:-60} + +owner_pid="" + +cleanup() { + if [[ -n "${owner_pid}" ]] && kill -0 "${owner_pid}" 2>/dev/null; then + kill "${owner_pid}" 2>/dev/null || true + wait "${owner_pid}" 2>/dev/null || true + fi +} +trap cleanup EXIT INT TERM + +"${PYTHON_BIN}" tools/tq_connection_smoke.py owner \ + --ray-address "${RAY_ADDRESS}" \ + --namespace "${TQ_NAMESPACE}" \ + --run-id "${SPECO_TQ_RUN_ID}" \ + --timeout "${SMOKE_TIMEOUT_SECONDS}" & +owner_pid=$! + +"${PYTHON_BIN}" tools/tq_connection_smoke.py client \ + --ray-address "${RAY_ADDRESS}" \ + --namespace "${TQ_NAMESPACE}" \ + --run-id "${SPECO_TQ_RUN_ID}" \ + --timeout "${SMOKE_TIMEOUT_SECONDS}" + +wait "${owner_pid}" +owner_pid="" +trap - EXIT INT TERM + +echo "TQ_CONNECTION_SMOKE_OK run_id=${SPECO_TQ_RUN_ID}" diff --git a/examples/tq_connection_smoke.py b/tools/tq_connection_smoke.py similarity index 77% rename from examples/tq_connection_smoke.py rename to tools/tq_connection_smoke.py index df873e76..ce198fa3 100644 --- a/examples/tq_connection_smoke.py +++ b/tools/tq_connection_smoke.py @@ -1,205 +1,187 @@ -# Copyright 2026 Bytedance Ltd. and/or its affiliates -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. """Real two-process smoke test for Ray + TransferQueue 0.1.7. - -Start a Ray head first, then run ``owner`` and ``client`` in separate shells. -The owner publishes one protocol-valid sample; the client list/get/decodes and -clears it, publishes a done marker, and exits without killing the owner. -""" - -from __future__ import annotations - -import argparse -import time - -import torch - -from verl_speco.integration.transferqueue_bridge import ( - clear_samples, - close_transfer_queue_client, - close_transfer_queue_owner, - configure_transfer_queue, - connect_ray_cluster, - connect_transfer_queue_client, - get_samples, - list_samples, - put_sample, - start_transfer_queue_owner, -) -from verl_speco.trainer.feature_store import DraftFeatureSample -from verl_speco.transport.drafter_sample_protocol import ( - ExpectedFeatureConfig, - SampleMetadata, - decode_sample, - encode_sample, - make_ready_tag, - make_sample_key, -) - - -def _config(args) -> dict: - return { - "enable": True, - "package_version": "0.1.7", - "ray": {"address": args.ray_address, "namespace": args.namespace}, - "partition_id": "speco_drafter_features", - "run_id": args.run_id, - "schema_version": 1, - "controller": {"polling_mode": True}, - "backend": { - "storage_backend": "SimpleStorage", - "SimpleStorage": { - "total_storage_size": 32, - "num_data_storage_units": 1, - }, - }, - } - - -def _record(run_id: str, sequence_no: int): - meta = SampleMetadata( - schema_version=1, - run_id=run_id, - sample_id=f"smoke-{sequence_no:04d}", - sequence_no=sequence_no, - algorithm="DSPARK", - target_model_id="smoke-target", - target_model_revision="smoke-revision", - tokenizer_fingerprint="smoke-tokenizer", - target_layer_ids=[0], - hidden_states_layout="dflash_aux", - hidden_dtype="float32", - hidden_shape=[3, 4], - feature_length=3, - full_sequence_length=3, - feature_start=0, - feature_end=3, - use_logits=False, - ) - sample = DraftFeatureSample( - algorithm="DSPARK", - input_ids=torch.tensor([1, 2, 3]) + sequence_no, - loss_mask=torch.tensor([0.0, 1.0, 1.0]), - position_ids=torch.tensor([0, 1, 2]), - hidden_states=( - torch.arange(12, dtype=torch.float32).reshape(3, 4) + sequence_no - ), - ) - return meta, sample - - -def _wait_for(predicate, timeout: float, description: str): - deadline = time.monotonic() + timeout - while time.monotonic() < deadline: - value = predicate() - if value: - return value - time.sleep(0.2) - raise TimeoutError(f"Timed out waiting for {description}") - - -def run_owner(args) -> None: - config = _config(args) - configure_transfer_queue(config) - connect_ray_cluster(args.ray_address, args.namespace) - start_transfer_queue_owner(config) - records = [_record(args.run_id, sequence_no) for sequence_no in range(2)] - keys = [make_sample_key(meta) for meta, _ in records] - done_key = f"control:v1:{args.run_id}:smoke-client-done" - try: - for (meta, sample), key in zip(records, keys, strict=True): - put_sample(key, encode_sample(sample, meta), tag=make_ready_tag(meta)) - print(f"OWNER_READY keys={keys}", flush=True) - _wait_for( - lambda: list_samples().get(done_key), - args.timeout, - "client done marker", - ) - clear_samples([done_key]) - print("OWNER_OBSERVED_CLIENT_DONE", flush=True) - finally: - close_transfer_queue_owner() - print("OWNER_CLOSED", flush=True) - - -def run_client(args) -> None: - config = _config(args) - configure_transfer_queue(config) - connect_ray_cluster(args.ray_address, args.namespace) - connect_transfer_queue_client() - metas = [_record(args.run_id, sequence_no)[0] for sequence_no in range(2)] - keys = [make_sample_key(meta) for meta in metas] - done_key = f"control:v1:{args.run_id}:smoke-client-done" - try: - _wait_for( - lambda: all(key in list_samples() for key in keys), - args.timeout, - "owner sample", - ) - tags = list_samples() - fetched = get_samples(keys) - restored = [ - decode_sample( - key, - tags[key], - fields, - ExpectedFeatureConfig( - run_id=args.run_id, - target_model_id="smoke-target", - hidden_states_layout="dflash_aux", - hidden_dtype="float32", - ), - ) - for key, fields in fetched - ] - assert restored[0].hidden_states.tolist() == torch.arange(12).reshape(3, 4).tolist() - assert restored[1].hidden_states.tolist() == ( - torch.arange(12).reshape(3, 4) + 1 - ).tolist() - clear_samples(keys) - put_sample( - done_key, - {"marker": torch.tensor([1], dtype=torch.uint8)}, - tag={ - "record_type": "control", - "status": "client_done", - "schema_version": 1, - "run_id": args.run_id, - }, - ) - print( - f"CLIENT_OK samples={len(restored)} shape={tuple(restored[0].hidden_states.shape)}", - flush=True, - ) - finally: - close_transfer_queue_client() - print("CLIENT_CLOSED_LOCAL_ONLY", flush=True) - - -def main() -> None: - parser = argparse.ArgumentParser() - parser.add_argument("role", choices=("owner", "client")) - parser.add_argument("--ray-address", required=True) - parser.add_argument("--namespace", default="speco-drafter-smoke") - parser.add_argument("--run-id", default="tq-smoke") - parser.add_argument("--timeout", type=float, default=60.0) - args = parser.parse_args() - if args.role == "owner": - run_owner(args) - else: - run_client(args) - - -if __name__ == "__main__": - main() + +Start a Ray head first, then run ``owner`` and ``client`` in separate shells. +The owner publishes one protocol-valid sample; the client list/get/decodes and +clears it, publishes a done marker, and exits without killing the owner. +""" + +from __future__ import annotations + +import argparse +import time + +import torch + +from verl_speco.integration.transferqueue_bridge import ( + close_transfer_queue_owner, + configure_transfer_queue, + connect_ray_cluster, + list_samples, + put_sample, + start_transfer_queue_owner, +) +from verl_speco.trainer.feature_store import DraftFeatureSample +from verl_speco.trainer.tq_feature_store import TQFeatureStore +from verl_speco.trainer.tq_sample_source import TQFeatureDataLoader +from verl_speco.transport.drafter_sample_protocol import ( + SampleMetadata, + encode_sample, + make_eos_record, + make_ready_tag, + make_sample_key, +) + + +def _config(args) -> dict: + return { + "enable": True, + "package_version": "0.1.7", + "ray": {"address": args.ray_address, "namespace": args.namespace}, + "partition_id": "speco_drafter_features", + "run_id": args.run_id, + "schema_version": 1, + "controller": {"polling_mode": True}, + "backend": { + "storage_backend": "SimpleStorage", + "SimpleStorage": { + "total_storage_size": 32, + "num_data_storage_units": 1, + }, + }, + } + + +def _record(run_id: str, sequence_no: int): + meta = SampleMetadata( + schema_version=1, + run_id=run_id, + sample_id=f"smoke-{sequence_no:04d}", + sequence_no=sequence_no, + algorithm="DSPARK", + target_model_id="smoke-target", + target_model_revision="smoke-revision", + tokenizer_fingerprint="smoke-tokenizer", + target_layer_ids=[0], + hidden_states_layout="dflash_aux", + hidden_dtype="float32", + hidden_shape=[3, 4], + feature_length=3, + full_sequence_length=3, + feature_start=0, + feature_end=3, + use_logits=False, + ) + sample = DraftFeatureSample( + algorithm="DSPARK", + input_ids=torch.tensor([1, 2, 3]) + sequence_no, + loss_mask=torch.tensor([0.0, 1.0, 1.0]), + position_ids=torch.tensor([0, 1, 2]), + hidden_states=( + torch.arange(12, dtype=torch.float32).reshape(3, 4) + sequence_no + ), + ) + return meta, sample + + +def _wait_for(predicate, timeout: float, description: str): + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + value = predicate() + if value: + return value + time.sleep(0.2) + raise TimeoutError(f"Timed out waiting for {description}") + + +def run_owner(args) -> None: + config = _config(args) + configure_transfer_queue(config) + connect_ray_cluster(args.ray_address, args.namespace) + start_transfer_queue_owner(config) + records = [_record(args.run_id, sequence_no) for sequence_no in range(2)] + keys = [make_sample_key(meta) for meta, _ in records] + try: + owner_ready_key = f"control:v1:{args.run_id}:owner-ready" + put_sample( + owner_ready_key, + {"marker": torch.tensor([1], dtype=torch.uint8)}, + tag={ + "record_type": "control", + "status": "owner_ready", + "schema_version": 1, + "run_id": args.run_id, + }, + ) + for (meta, sample), key in zip(records, keys, strict=True): + put_sample(key, encode_sample(sample, meta), tag=make_ready_tag(meta)) + eos_key, eos_fields, eos_tag = make_eos_record(args.run_id, len(records)) + put_sample(eos_key, eos_fields, tag=eos_tag) + print(f"OWNER_READY keys={keys}", flush=True) + _wait_for( + lambda: all(key not in list_samples() for key in keys), + args.timeout, + "client to clear the smoke sample keys", + ) + print("OWNER_OBSERVED_SAMPLES_CLEARED", flush=True) + finally: + close_transfer_queue_owner() + print("OWNER_CLOSED", flush=True) + + +def run_client(args) -> None: + config = _config(args) + store = TQFeatureStore.from_config(config) + loader = TQFeatureDataLoader(store, batch_size=2, rank=0, world_size=1) + try: + iterator = iter(loader) + batch = next(iterator) + restored = batch.local_samples + assert restored[0].hidden_states.tolist() == torch.arange(12).reshape(3, 4).tolist() + assert restored[1].hidden_states.tolist() == ( + torch.arange(12).reshape(3, 4) + 1 + ).tolist() + loader.clear_completed_batch(batch.global_keys) + try: + next(iterator) + except StopIteration: + pass + else: + raise AssertionError("TQ Consumer did not stop after draining EOS") + print( + f"CLIENT_OK samples={len(restored)} shape={tuple(restored[0].hidden_states.shape)}", + flush=True, + ) + finally: + store.close_local() + print("CLIENT_CLOSED_LOCAL_ONLY", flush=True) + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("role", choices=("owner", "client")) + parser.add_argument("--ray-address", required=True) + parser.add_argument("--namespace", default="speco-drafter-smoke") + parser.add_argument("--run-id", default="tq-smoke") + parser.add_argument("--timeout", type=float, default=60.0) + args = parser.parse_args() + if args.role == "owner": + run_owner(args) + else: + run_client(args) + + +if __name__ == "__main__": + main() diff --git a/tools/tq_delayed_test_producer.py b/tools/tq_delayed_test_producer.py new file mode 100644 index 00000000..6c89a500 --- /dev/null +++ b/tools/tq_delayed_test_producer.py @@ -0,0 +1,234 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +"""Publish protocol-valid synthetic DSpark features in delayed batches. + +This is an integration-test Producer for the standalone TQ Consumer. It does +not run vLLM: tensor shapes are derived from the target model config so the +normal DSpark preprocessing and training path can consume the samples. +""" + +from __future__ import annotations + +import argparse +import json +import time +from pathlib import Path +from typing import Any + +import torch + +from verl_speco.integration.transferqueue_bridge import ( + close_transfer_queue_client, + configure_transfer_queue, + connect_ray_cluster, + connect_transfer_queue_client, + list_samples, + put_sample, +) +from verl_speco.trainer.feature_store import DraftFeatureSample +from verl_speco.transport.drafter_sample_protocol import ( + SampleMetadata, + encode_sample, + make_eos_record, + make_ready_tag, + make_sample_key, +) + + +def _tq_config(args: argparse.Namespace) -> dict[str, Any]: + return { + "enable": True, + "package_version": "0.1.7", + "ray": {"address": args.ray_address, "namespace": args.namespace}, + "partition_id": args.partition_id, + "run_id": args.run_id, + "schema_version": 1, + "controller": {"polling_mode": True}, + "backend": { + "storage_backend": "SimpleStorage", + "SimpleStorage": { + "total_storage_size": 100000, + "num_data_storage_units": 8, + }, + }, + } + + +def _model_dimensions(model_path: str) -> tuple[int, int, int]: + config_path = Path(model_path) / "config.json" + with config_path.open("r", encoding="utf-8") as handle: + config = json.load(handle) + text_config = config.get("text_config") or config + hidden_size = int(text_config["hidden_size"]) + vocab_size = int(text_config["vocab_size"]) + num_hidden_layers = int(text_config.get("num_hidden_layers", 1)) + return hidden_size, vocab_size, num_hidden_layers + + +def _wait_for_owner(args: argparse.Namespace) -> None: + owner_ready_key = f"control:v1:{args.run_id}:owner-ready" + deadline = time.monotonic() + args.timeout + while time.monotonic() < deadline: + tag = list_samples().get(owner_ready_key) + if isinstance(tag, dict) and tag.get("status") == "owner_ready": + return + time.sleep(0.2) + raise TimeoutError(f"Timed out waiting for TQ owner key {owner_ready_key!r}") + + +def _target_layer_ids(num_target_layers: int, num_hidden_layers: int) -> list[int]: + if num_target_layers <= 0: + raise ValueError("num_target_layers must be positive") + if num_hidden_layers < 4: + raise ValueError("the DSpark target model must have at least four hidden layers") + if num_target_layers == 1: + return [num_hidden_layers // 2] + start = 1 + end = num_hidden_layers - 3 + span = end - start + return [ + int(round(start + (index * span) / (num_target_layers - 1))) + for index in range(num_target_layers) + ] + + +def _sample( + args: argparse.Namespace, + *, + sequence_no: int, + hidden_size: int, + vocab_size: int, + target_layer_ids: list[int], +) -> tuple[str, dict[str, torch.Tensor], dict[str, Any]]: + generator = torch.Generator(device="cpu") + generator.manual_seed(args.seed + sequence_no) + # The default standalone DSpark configuration enables L1 distillation. + # Its wire layout contains N auxiliary target layers followed by the + # target model's final hidden state, all concatenated on the last axis. + hidden_dim = hidden_size * (len(target_layer_ids) + 1) + input_ids = torch.randint( + low=0, + high=vocab_size, + size=(args.sequence_length,), + generator=generator, + dtype=torch.long, + ) + loss_mask = torch.ones(args.sequence_length, dtype=torch.float32) + loss_mask[: max(1, args.sequence_length // 4)] = 0 + hidden_states = torch.randn( + args.sequence_length, + hidden_dim, + generator=generator, + dtype=torch.bfloat16, + ) + sample_id = f"consumer-test-{sequence_no:08d}" + metadata = SampleMetadata( + schema_version=1, + run_id=args.run_id, + sample_id=sample_id, + sequence_no=sequence_no, + algorithm="DSPARK", + target_model_id=str(Path(args.model_path).resolve()), + target_model_revision="local-test", + tokenizer_fingerprint="synthetic-consumer-test", + target_layer_ids=target_layer_ids, + hidden_states_layout="dflash_aux_plus_last", + hidden_dtype="bfloat16", + hidden_shape=[args.sequence_length, hidden_dim], + feature_length=args.sequence_length, + full_sequence_length=args.sequence_length, + feature_start=0, + feature_end=args.sequence_length, + use_logits=False, + ) + sample = DraftFeatureSample( + algorithm="DSPARK", + input_ids=input_ids, + loss_mask=loss_mask, + position_ids=torch.arange(args.sequence_length, dtype=torch.long), + hidden_states=hidden_states, + metadata={"hidden_states_layout": "dflash_aux_plus_last"}, + ) + key = make_sample_key(metadata) + return key, encode_sample(sample, metadata), make_ready_tag(metadata) + + +def run(args: argparse.Namespace) -> None: + hidden_size, vocab_size, num_hidden_layers = _model_dimensions(args.model_path) + target_layer_ids = _target_layer_ids(args.num_target_layers, num_hidden_layers) + samples_per_batch = args.world_size * args.batch_size_per_gpu + total_samples = args.num_batches * samples_per_batch + config = _tq_config(args) + configure_transfer_queue(config) + connect_ray_cluster(args.ray_address, args.namespace) + connect_transfer_queue_client() + try: + _wait_for_owner(args) + print( + "PRODUCER_CONNECTED " + f"batches={args.num_batches} samples_per_batch={samples_per_batch} " + f"sequence_length={args.sequence_length} hidden_shape=" + f"({args.sequence_length}, {hidden_size * (args.num_target_layers + 1)})", + flush=True, + ) + time.sleep(args.initial_delay) + sequence_no = 0 + for batch_index in range(args.num_batches): + keys = [] + for _ in range(samples_per_batch): + key, fields, tag = _sample( + args, + sequence_no=sequence_no, + hidden_size=hidden_size, + vocab_size=vocab_size, + target_layer_ids=target_layer_ids, + ) + put_sample(key, fields, tag=tag) + keys.append(key) + sequence_no += 1 + print( + f"PRODUCER_BATCH_READY batch={batch_index + 1}/{args.num_batches} " + f"samples={len(keys)} sequence_no=[{sequence_no - len(keys)},{sequence_no})", + flush=True, + ) + if batch_index + 1 < args.num_batches: + time.sleep(args.batch_interval) + eos_key, eos_fields, eos_tag = make_eos_record(args.run_id, total_samples) + put_sample(eos_key, eos_fields, tag=eos_tag) + print(f"PRODUCER_EOS total_samples={total_samples} key={eos_key}", flush=True) + finally: + close_transfer_queue_client() + print("PRODUCER_CLOSED_LOCAL", flush=True) + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--model-path", required=True) + parser.add_argument("--ray-address", required=True) + parser.add_argument("--namespace", default="speco-drafter") + parser.add_argument("--partition-id", default="speco_drafter_features") + parser.add_argument("--run-id", required=True) + parser.add_argument("--world-size", type=int, default=1) + parser.add_argument("--batch-size-per-gpu", type=int, default=1) + parser.add_argument("--num-batches", type=int, default=3) + parser.add_argument("--sequence-length", type=int, default=64) + parser.add_argument("--num-target-layers", type=int, default=5) + parser.add_argument("--initial-delay", type=float, default=5.0) + parser.add_argument("--batch-interval", type=float, default=5.0) + parser.add_argument("--timeout", type=float, default=120.0) + parser.add_argument("--seed", type=int, default=2026) + args = parser.parse_args() + if args.world_size <= 0 or args.batch_size_per_gpu <= 0: + parser.error("world-size and batch-size-per-gpu must be positive") + if args.num_batches <= 0 or args.sequence_length <= 0: + parser.error("num-batches and sequence-length must be positive") + run(args) + + +if __name__ == "__main__": + main() diff --git a/verl_speco/config/speco_base.yaml b/verl_speco/config/speco_base.yaml index 13152ca9..bce79dfb 100644 --- a/verl_speco/config/speco_base.yaml +++ b/verl_speco/config/speco_base.yaml @@ -173,6 +173,7 @@ actor_rollout_ref: data_buffer_max_size: 1024 hidden_state_clip_value: 1.0e4 feature_store: + # `tq` selects the streaming standalone Consumer and does not use path. type: torch_shard path: null max_samples_per_shard: 1024 diff --git a/verl_speco/draft_train_launcher.py b/verl_speco/draft_train_launcher.py index 1f6b1a6d..087e8a5a 100644 --- a/verl_speco/draft_train_launcher.py +++ b/verl_speco/draft_train_launcher.py @@ -57,6 +57,12 @@ "speco.draft_training.standalone", "actor_rollout_ref.rollout.drafter.training.standalone", ) +_FEATURE_STORE_TYPE_KEY = "actor_rollout_ref.rollout.drafter.training.feature_store.type" +_TQ_ENABLE_KEY = "actor_rollout_ref.rollout.drafter.training.transfer_queue.enable" +_TQ_RAY_ADDRESS_KEY = ( + "actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.address" +) +_TQ_RUN_ID_KEY = "actor_rollout_ref.rollout.drafter.training.transfer_queue.run_id" _LAUNCH_OVERRIDE_KEYS = frozenset( _NPROC_KEYS @@ -166,6 +172,27 @@ def normalize_training_args( return normalized +def validate_tq_launch_config(overrides: list[str]) -> None: + """Fail early when the standalone TQ Consumer lacks connection identity.""" + + store_type = _find_override(overrides, (_FEATURE_STORE_TYPE_KEY,)) + if str(store_type or "").strip().lower() != "tq": + return + enabled = _find_override(overrides, (_TQ_ENABLE_KEY,)) + if not _parse_bool(enabled, default=False): + raise ValueError( + "feature_store.type=tq requires training.transfer_queue.enable=true" + ) + ray_address = str(_find_override(overrides, (_TQ_RAY_ADDRESS_KEY,)) or "").strip() + if not ray_address or ray_address.lower() in {"null", "none"}: + raise ValueError( + "feature_store.type=tq requires training.transfer_queue.ray.address" + ) + run_id = str(_find_override(overrides, (_TQ_RUN_ID_KEY,)) or "").strip() + if not run_id or run_id.lower() in {"null", "none"}: + raise ValueError("feature_store.type=tq requires training.transfer_queue.run_id") + + def build_torch_distributed_command( config: DraftTrainLaunchConfig, training_args: list[str], @@ -220,6 +247,7 @@ def main(argv: list[str] | None = None) -> int: ) args, training_args = parser.parse_known_args(argv) + validate_tq_launch_config(training_args) launch_config = resolve_launch_config(training_args, module=args.module) normalized_training_args = normalize_training_args(training_args, launch_config) command = build_torch_distributed_command( diff --git a/verl_speco/trainer/draft_training_loop.py b/verl_speco/trainer/draft_training_loop.py index 53c75ceb..b861f10d 100644 --- a/verl_speco/trainer/draft_training_loop.py +++ b/verl_speco/trainer/draft_training_loop.py @@ -40,6 +40,7 @@ build_feature_store_from_config, ) from verl_speco.trainer.standalone_checkpoint import rewrite_standalone_runtime_config +from verl_speco.trainer.tq_sample_source import TQFeatureDataLoader, TQLocalBatch logger = logging.getLogger(__name__) @@ -85,10 +86,12 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: str(training_cfg.get("mode", "offline") or "offline").strip().lower() ) replay_feature_store_types = {"token_replay", "jsonl_token_replay", "jsonl"} - if not feature_store_cfg.get("path"): + if feature_store_type != "tq" and not feature_store_cfg.get("path"): raise ValueError( "actor_rollout_ref.rollout.drafter.training.feature_store.path is required" ) + if feature_store_type == "tq" and training_mode != "offline": + raise ValueError("feature_store.type=tq requires standalone training.mode=offline") if feature_store_type in replay_feature_store_types and training_mode != "offline": raise ValueError( f"feature_store.type={feature_store_type} is supported only by " @@ -167,7 +170,18 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: feature_store_cfg.tokenizer_path = tokenizer_path except AttributeError: feature_store_cfg["tokenizer_path"] = tokenizer_path - store = build_feature_store_from_config(feature_store_cfg, read_only=True) + store = build_feature_store_from_config( + feature_store_cfg, + read_only=True, + transfer_queue_cfg=training_cfg.get("transfer_queue"), + ) + if feature_store_type == "tq": + current_stage = "connect_tq_feature_store" + _connect_tq_store_across_ranks( + store, + rank=rank, + device=trainer.runtime_device, + ) logger.info( "[standalone rank=%s] feature store opened elapsed=%.3fs", rank, @@ -201,17 +215,31 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: time.perf_counter() - stage_started, ) current_stage = "create_dataloader" - loader = DraftFeatureDataLoader( - store, - DraftFeatureDataLoaderConfig( + loader: Any + if feature_store_type == "tq": + tq_cfg = training_cfg.get("transfer_queue") or {} + loader = TQFeatureDataLoader( + store, batch_size=int(training_cfg.get("batch_size_per_gpu", 4)), rank=rank, world_size=world_size, - shuffle=bool(feature_store_cfg.get("shuffle", True)), - repeat=bool(feature_store_cfg.get("repeat", True)), - seed=int(training_cfg.get("seed", 0) or 0), - ), - ) + poll_interval_seconds=float( + tq_cfg.get("poll_interval_seconds", 0.5) or 0.5 + ), + drop_last=bool(tq_cfg.get("drop_last", True)), + ) + else: + loader = DraftFeatureDataLoader( + store, + DraftFeatureDataLoaderConfig( + batch_size=int(training_cfg.get("batch_size_per_gpu", 4)), + rank=rank, + world_size=world_size, + shuffle=bool(feature_store_cfg.get("shuffle", True)), + repeat=bool(feature_store_cfg.get("repeat", True)), + seed=int(training_cfg.get("seed", 0) or 0), + ), + ) logger.info( "[standalone rank=%s] dataloader ready batch_size_per_gpu=%s " "shuffle=%s repeat=%s", @@ -223,6 +251,11 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: sample_source = loader pipeline_cfg = training_cfg.get("target_feature_pipeline", {}) or {} pipeline_enabled = bool(pipeline_cfg.get("enabled", False)) + if feature_store_type == "tq" and pipeline_enabled: + raise ValueError( + "feature_store.type=tq already contains target hidden states and cannot be " + "combined with target_feature_pipeline.enabled=true" + ) if pipeline_enabled: if feature_replayer is None or not feature_replayer.backend.startswith( "vllm_" @@ -266,13 +299,21 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: sample_iterator = iter(sample_source) while max_steps <= 0 or successful_steps < max_steps: current_stage = "load_next_batch" - samples = _next_batch_across_ranks( + loaded_batch = _next_batch_across_ranks( sample_iterator, rank=rank, device=trainer.runtime_device, ) - if samples is None: + if loaded_batch is None: break + tq_local_batch = ( + loaded_batch if isinstance(loaded_batch, TQLocalBatch) else None + ) + samples = ( + tq_local_batch.local_samples + if tq_local_batch is not None + else loaded_batch + ) step_started = time.perf_counter() attempted_batches += 1 log_batch_progress = _should_log_batch_progress(attempted_batches) @@ -334,6 +375,11 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: has_batch = batch is not None current_stage = "synchronize_batch_readiness" if not _all_ranks_true(has_batch, trainer.runtime_device): + if tq_local_batch is not None: + raise RuntimeError( + "TQ Consumer could not prepare a valid batch on every rank; " + "the TQ keys were intentionally not cleared" + ) if rank == 0: logger.warning( "Skipping standalone drafter batch: at least one rank has no valid batch" @@ -362,7 +408,20 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: ) from step_error current_stage = "synchronize_training_step" if not _all_ranks_true(ok, trainer.runtime_device): + if tq_local_batch is not None: + raise RuntimeError( + "TQ Consumer training_step_from_batch failed on at least one rank; " + "the TQ keys were intentionally not cleared" + ) continue + if tq_local_batch is not None: + current_stage = "clear_tq_batch" + _clear_tq_batch_across_ranks( + cast(TQFeatureDataLoader, loader), + tq_local_batch.global_keys, + rank=rank, + device=trainer.runtime_device, + ) successful_steps += 1 optimizer_step = int(trainer.optimizer_steps_total) if optimizer_step <= initial_optimizer_step: @@ -1054,12 +1113,58 @@ def _all_ranks_true(value: bool, device: torch.device) -> bool: return bool(ready.item()) +def _clear_tq_batch_across_ranks( + loader: TQFeatureDataLoader, + global_keys: list[str] | None, + *, + rank: int, + device: torch.device, +) -> None: + """Clear once on rank 0 and report a clear failure to every training rank.""" + + local_error: BaseException | None = None + if rank == 0: + try: + loader.clear_completed_batch(global_keys) + except BaseException as exc: # noqa: BLE001 + local_error = exc + failed = torch.tensor( + 1 if local_error is not None else 0, + dtype=torch.int32, + device=device, + ) + if dist.is_initialized() and dist.get_world_size() > 1: + dist.all_reduce(failed, op=dist.ReduceOp.MAX) + if bool(failed.item()): + if local_error is not None: + raise RuntimeError("rank 0 failed to clear a completed TQ batch") from local_error + raise RuntimeError("rank 0 failed to clear a completed TQ batch") + + +def _connect_tq_store_across_ranks(store, *, rank: int, device: torch.device) -> None: + """Connect every rank before any rank enters TQ key-discovery broadcasts.""" + + local_error: BaseException | None = None + try: + store.connect() + except BaseException as exc: # noqa: BLE001 + local_error = exc + connected = _all_ranks_true(local_error is None, device) + if connected: + return + if local_error is not None: + raise RuntimeError(f"TQ Consumer failed to connect on rank={rank}") from local_error + raise RuntimeError( + f"TQ Consumer failed to connect on another rank; rank={rank} is stopping" + ) + + def _next_batch_across_ranks( source, *, rank: int, device: torch.device, -) -> list[Any] | None: +) -> Any | None: """Fetch one batch and make producer failures visible to every rank. Producer and Mooncake errors happen before the FSDP training step. Every @@ -1067,7 +1172,7 @@ def _next_batch_across_ranks( any rank is allowed to enter model collectives. This prevents healthy ranks from waiting in FSDP after another rank has already started cleanup. """ - samples: list[Any] | None = None + samples: Any | None = None local_error: BaseException | None = None exhausted = False try: diff --git a/verl_speco/trainer/feature_store.py b/verl_speco/trainer/feature_store.py index 8ebdd156..36c5ccb0 100644 --- a/verl_speco/trainer/feature_store.py +++ b/verl_speco/trainer/feature_store.py @@ -1048,12 +1048,26 @@ def build_feature_store_from_config( read_only: bool = False, metadata: dict[str, Any] | None = None, shard_prefix: str = "shard", -) -> DraftFeatureStore: + transfer_queue_cfg: Any | None = None, +) -> Any: store_type = ( str(feature_store_cfg.get("type", "torch_shard") or "torch_shard") .strip() .lower() ) + if store_type == "tq": + if not read_only: + raise ValueError("feature_store.type=tq is a read-only Consumer data source") + from verl_speco.trainer.tq_feature_store import TQFeatureStore + + tq_cfg = transfer_queue_cfg + if tq_cfg is None: + tq_cfg = feature_store_cfg.get("tq") + if tq_cfg is None: + raise ValueError( + "feature_store.type=tq requires the sibling training.transfer_queue configuration" + ) + return TQFeatureStore.from_config(tq_cfg) if store_type == "torch_shard": store_cls: type[TorchShardFeatureStore] = TorchShardFeatureStore elif store_type == "token_replay": diff --git a/verl_speco/trainer/tq_feature_store.py b/verl_speco/trainer/tq_feature_store.py new file mode 100644 index 00000000..1921fe95 --- /dev/null +++ b/verl_speco/trainer/tq_feature_store.py @@ -0,0 +1,235 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Streaming TransferQueue feature source for standalone drafter training.""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, Mapping, Sequence + +from verl_speco.integration.transferqueue_bridge import ( + clear_samples, + close_transfer_queue_client, + configure_transfer_queue, + connect_ray_cluster, + connect_transfer_queue_client, + get_samples, + list_samples, +) +from verl_speco.trainer.feature_store import DraftFeatureSample +from verl_speco.transport.drafter_sample_protocol import ( + ExpectedFeatureConfig, + decode_sample, +) + + +@dataclass(frozen=True) +class ReadyEntry: + """One discoverable sample record; payload tensors are not loaded yet.""" + + key: str + tag: dict[str, Any] + + +@dataclass(frozen=True) +class EosMetadata: + """End-of-stream control record published by the Producer.""" + + key: str + run_id: str + schema_version: int + total_samples: int + + +class TQFeatureStore: + """Thin Consumer adapter over the shared TransferQueue bridge. + + This intentionally does not implement the static ``iter_keys/read`` feature + store protocol. TQ keys are added by the Producer and removed after a + successful optimizer step, so discovery must happen for every global batch. + """ + + def __init__(self, config: Mapping[str, Any]): + self.config = _plain_dict(config) + self.run_id = str(self.config.get("run_id") or "").strip() + if not self.run_id: + raise ValueError("transfer_queue.run_id is required for a TQ Consumer") + self.schema_version = int(self.config.get("schema_version", 1)) + self.algorithm = str(self.config.get("algorithm", "DSPARK") or "DSPARK").upper() + if self.algorithm != "DSPARK": + raise ValueError( + f"Standalone TQ Consumer currently supports only DSPARK, got {self.algorithm!r}" + ) + ray_cfg = _plain_dict(self.config.get("ray") or {}) + self.ray_address = str(ray_cfg.get("address") or "").strip() + self.ray_namespace = str(ray_cfg.get("namespace") or "").strip() or None + if not self.ray_address: + raise ValueError("transfer_queue.ray.address is required for a TQ Consumer") + self._connected = False + # First version deliberately checks only run/protocol/algorithm. Tensor + # presence, lengths, shape and dtype self-consistency remain enforced by + # decode_sample; model/tokenizer/layer identity checks stay disabled. + self.expected_config = ExpectedFeatureConfig( + run_id=self.run_id, + schema_version=self.schema_version, + algorithm=self.algorithm, + ) + + @classmethod + def from_config(cls, config: Any) -> "TQFeatureStore": + return cls(_plain_dict(config)) + + def connect(self) -> None: + if self._connected: + return + if not configure_transfer_queue(self.config): + raise RuntimeError( + "TQ Consumer requires transfer_queue.enable=true and TransferQueue==0.1.7" + ) + connect_ray_cluster(self.ray_address, self.ray_namespace) + connect_transfer_queue_client() + self._connected = True + + def list_ready(self, run_id: str | None = None) -> list[ReadyEntry]: + self._require_connected() + expected_run_id = str(run_id or self.run_id) + ready: list[ReadyEntry] = [] + for key, raw_tag in list_samples().items(): + tag = dict(raw_tag) + if tag.get("record_type") != "sample" or tag.get("status") != "ready": + continue + if str(tag.get("run_id") or "") != expected_run_id: + continue + if int(tag.get("schema_version", -1)) != self.schema_version: + continue + if str(tag.get("algorithm") or "").upper() != self.algorithm: + continue + try: + int(tag["sequence_no"]) + str(tag["sample_id"]) + except (KeyError, TypeError, ValueError): + continue + ready.append(ReadyEntry(key=str(key), tag=tag)) + ready.sort(key=lambda entry: (int(entry.tag["sequence_no"]), entry.key)) + return ready + + def owner_ready(self) -> bool: + """Whether the standalone Owner published this run's readiness marker.""" + + self._require_connected() + key = f"control:v{self.schema_version}:{self.run_id}:owner-ready" + tag = list_samples().get(key) + if not isinstance(tag, Mapping): + return False + return ( + tag.get("record_type") == "control" + and tag.get("status") == "owner_ready" + and str(tag.get("run_id") or "") == self.run_id + and int(tag.get("schema_version", -1)) == self.schema_version + ) + + def get_many(self, entries: Sequence[ReadyEntry]) -> list[DraftFeatureSample]: + self._require_connected() + if not entries: + return [] + records = get_samples([entry.key for entry in entries]) + if len(records) != len(entries): + raise RuntimeError( + f"TQ returned {len(records)} records for {len(entries)} requested entries" + ) + samples: list[DraftFeatureSample] = [] + for entry, (key, fields) in zip(entries, records, strict=True): + if key != entry.key: + raise RuntimeError( + f"TQ batch result order mismatch: got key={key!r}, expected={entry.key!r}" + ) + samples.append( + decode_sample( + key=key, + tag=entry.tag, + fields=fields, + expected_config=self.expected_config, + ) + ) + return samples + + def clear_many(self, keys: Sequence[str]) -> None: + self._require_connected() + clear_samples([str(key) for key in keys]) + + def read_eos(self, run_id: str | None = None) -> EosMetadata | None: + self._require_connected() + expected_run_id = str(run_id or self.run_id) + for key, raw_tag in list_samples().items(): + tag = dict(raw_tag) + if tag.get("record_type") != "control" or tag.get("status") != "eos": + continue + if str(tag.get("run_id") or "") != expected_run_id: + continue + if int(tag.get("schema_version", -1)) != self.schema_version: + continue + return EosMetadata( + key=str(key), + run_id=expected_run_id, + schema_version=self.schema_version, + total_samples=int(tag.get("total_samples", 0)), + ) + return None + + def close_local(self) -> None: + if not self._connected: + return + close_transfer_queue_client() + self._connected = False + + def close(self) -> None: + """Compatibility with the standalone loop's existing cleanup path.""" + + self.close_local() + + def _require_connected(self) -> None: + if not self._connected: + raise RuntimeError("TQFeatureStore.connect() must be called before data access") + + +def _plain_dict(value: Any) -> dict[str, Any]: + if value is None: + return {} + try: + from omegaconf import DictConfig, OmegaConf + + if isinstance(value, DictConfig): + converted = OmegaConf.to_container(value, resolve=True) + if not isinstance(converted, dict): + raise TypeError("Expected a mapping configuration") + return dict(converted) + except ImportError: # pragma: no cover - the project depends on OmegaConf + pass + if isinstance(value, Mapping): + return { + str(key): _plain_value(item) + for key, item in value.items() + } + raise TypeError(f"Expected a mapping configuration, got {type(value)!r}") + + +def _plain_value(value: Any) -> Any: + if isinstance(value, Mapping): + return {str(key): _plain_value(item) for key, item in value.items()} + if isinstance(value, (list, tuple)): + return [_plain_value(item) for item in value] + return value + + +__all__ = ["EosMetadata", "ReadyEntry", "TQFeatureStore"] diff --git a/verl_speco/trainer/tq_sample_source.py b/verl_speco/trainer/tq_sample_source.py new file mode 100644 index 00000000..a676c7b5 --- /dev/null +++ b/verl_speco/trainer/tq_sample_source.py @@ -0,0 +1,202 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Distributed streaming sample source for a TQ-backed standalone Consumer.""" + +from __future__ import annotations + +import logging +import time +from dataclasses import dataclass +from typing import Any, Iterator, Sequence + +import torch.distributed as dist + +from verl_speco.trainer.feature_store import DraftFeatureSample +from verl_speco.trainer.tq_feature_store import ReadyEntry, TQFeatureStore + + +logger = logging.getLogger(__name__) + + +@dataclass(frozen=True) +class TQLocalBatch: + """The payload owned by one rank plus rank-0's global cleanup keys.""" + + local_keys: list[str] + local_samples: list[DraftFeatureSample] + global_keys: list[str] | None + + +def build_assignments( + entries: Sequence[ReadyEntry], *, batch_size: int, world_size: int +) -> list[list[ReadyEntry]]: + """Split one complete global batch into disjoint contiguous rank batches.""" + + if batch_size <= 0: + raise ValueError(f"batch_size must be positive, got {batch_size}") + if world_size <= 0: + raise ValueError(f"world_size must be positive, got {world_size}") + expected = batch_size * world_size + if len(entries) != expected: + raise ValueError( + f"Expected exactly {expected} ready entries for one global batch, got {len(entries)}" + ) + return [ + list(entries[rank * batch_size : (rank + 1) * batch_size]) + for rank in range(world_size) + ] + + +class TQFeatureDataLoader: + """Poll TQ on rank 0, distribute keys, and fetch payloads on owner ranks.""" + + def __init__( + self, + store: TQFeatureStore, + *, + batch_size: int, + rank: int, + world_size: int, + poll_interval_seconds: float = 0.5, + drop_last: bool = True, + ): + self.store = store + self.batch_size = int(batch_size) + self.rank = int(rank) + self.world_size = int(world_size) + self.poll_interval_seconds = max(float(poll_interval_seconds), 0.01) + self.drop_last = bool(drop_last) + if self.batch_size <= 0: + raise ValueError("TQ Consumer batch_size_per_gpu must be positive") + if self.world_size <= 0 or not (0 <= self.rank < self.world_size): + raise ValueError( + "Invalid TQ Consumer rank/world_size: " + f"rank={self.rank}, world_size={self.world_size}" + ) + if not self.drop_last: + raise ValueError( + "TQ Consumer first version requires transfer_queue.drop_last=true" + ) + if dist.is_initialized() and dist.get_world_size() != self.world_size: + raise ValueError( + "TQFeatureDataLoader world_size does not match torch.distributed world size" + ) + + def __iter__(self) -> Iterator[TQLocalBatch]: + self.store.connect() + global_batch_size = self.batch_size * self.world_size + owner_ready = False + while True: + command: dict[str, Any] | None = None + if self.rank == 0: + try: + if not owner_ready: + owner_ready = self.store.owner_ready() + if not owner_ready: + time.sleep(self.poll_interval_seconds) + continue + ready = self.store.list_ready() + if len(ready) >= global_batch_size: + selected = ready[:global_batch_size] + assignments = build_assignments( + selected, + batch_size=self.batch_size, + world_size=self.world_size, + ) + command = { + "kind": "batch", + "global_keys": [entry.key for entry in selected], + "assignments": [ + [_entry_to_wire(entry) for entry in rank_entries] + for rank_entries in assignments + ], + } + else: + eos = self.store.read_eos() + if eos is not None: + tail_keys = [entry.key for entry in ready] + if tail_keys: + logger.info( + "Dropping %s TQ tail samples after EOS because one " + "global batch requires %s", + len(tail_keys), + global_batch_size, + ) + self.store.clear_many(tail_keys) + command = {"kind": "stop"} + else: + time.sleep(self.poll_interval_seconds) + continue + except BaseException as exc: # noqa: BLE001 + command = { + "kind": "error", + "message": f"rank 0 failed while discovering TQ samples: {exc}", + } + + command = self._broadcast_command(command) + if command.get("kind") == "error": + raise RuntimeError(str(command.get("message") or "TQ discovery failed")) + if command.get("kind") == "stop": + return + if command.get("kind") != "batch": + raise RuntimeError(f"Unsupported TQ loader command: {command!r}") + assignments = command.get("assignments") + if not isinstance(assignments, list) or len(assignments) != self.world_size: + raise RuntimeError("TQ loader received malformed rank assignments") + local_entries = [_entry_from_wire(item) for item in assignments[self.rank]] + samples = self.store.get_many(local_entries) + global_keys = ( + [str(key) for key in command.get("global_keys", [])] + if self.rank == 0 + else None + ) + yield TQLocalBatch( + local_keys=[entry.key for entry in local_entries], + local_samples=samples, + global_keys=global_keys, + ) + + def clear_completed_batch(self, global_keys: Sequence[str] | None) -> None: + if self.rank != 0: + return + if not global_keys: + raise ValueError("rank 0 requires global_keys to clear a completed TQ batch") + self.store.clear_many(global_keys) + + def _broadcast_command(self, command: dict[str, Any] | None) -> dict[str, Any]: + if not dist.is_initialized() or self.world_size == 1: + if command is None: + raise RuntimeError("rank 0 did not create a TQ loader command") + return command + payload: list[Any] = [command if self.rank == 0 else None] + dist.broadcast_object_list(payload, src=0) + received = payload[0] + if not isinstance(received, dict): + raise RuntimeError("TQ loader broadcast did not contain a command mapping") + return received + + +def _entry_to_wire(entry: ReadyEntry) -> dict[str, Any]: + return {"key": entry.key, "tag": dict(entry.tag)} + + +def _entry_from_wire(value: Any) -> ReadyEntry: + if not isinstance(value, dict) or "key" not in value or "tag" not in value: + raise TypeError(f"Invalid serialized ReadyEntry: {value!r}") + if not isinstance(value["tag"], dict): + raise TypeError("Serialized ReadyEntry.tag must be a mapping") + return ReadyEntry(key=str(value["key"]), tag=dict(value["tag"])) + + +__all__ = ["TQFeatureDataLoader", "TQLocalBatch", "build_assignments"] From 2e0de5f12166f2c2ae690d002cc97beae8ce1b30 Mon Sep 17 00:00:00 2001 From: vx120 <893600387@qq.com> Date: Fri, 21 Aug 2026 09:21:15 +0800 Subject: [PATCH 34/50] Add standalone TQ producer implementation Signed-off-by: vx120 <893600387@qq.com> --- docs/standalone_tq_producer.md | 201 +++++++++ examples/run_dspark_tq_producer.sh | 40 ++ pyproject.toml | 1 + tests/special_sanity/check_example_naming.py | 7 +- verl_speco/config/speco_base.yaml | 25 + verl_speco/producer/__init__.py | 14 + verl_speco/producer/input_reader.py | 206 +++++++++ verl_speco/producer/vllm_feature_client.py | 236 ++++++++++ verl_speco/standalone_tq_producer.py | 451 +++++++++++++++++++ verl_speco/trainer/target_feature_replay.py | 287 +++++++----- 10 files changed, 1358 insertions(+), 110 deletions(-) create mode 100644 docs/standalone_tq_producer.md create mode 100644 examples/run_dspark_tq_producer.sh create mode 100644 verl_speco/producer/__init__.py create mode 100644 verl_speco/producer/input_reader.py create mode 100644 verl_speco/producer/vllm_feature_client.py create mode 100644 verl_speco/standalone_tq_producer.py diff --git a/docs/standalone_tq_producer.md b/docs/standalone_tq_producer.md new file mode 100644 index 00000000..e9c95be0 --- /dev/null +++ b/docs/standalone_tq_producer.md @@ -0,0 +1,201 @@ +# Standalone vLLM → TransferQueue Producer + +Last updated: 08/20/2026 + +本文解释这次新增的 standalone Producer:它读取已经有 `prompt` 和 +`response` 的 JSONL,向 vLLM 请求 target hidden states,并把每条样本写到 +已存在的 TransferQueue(TQ)。它不启动 TQ owner,也不启动 Consumer/训练。 + +这条路径面向第一版 DSpark standalone 训练:Producer、TQ owner 和 Consumer +是三个独立 OS 进程;Ray 只用于让它们找到同一个 TQ Controller,hidden states +不通过 Ray object store 传输。 + +## 为什么需要这个 Producer + +此前仓库已经有两块基础能力: + +- `drafter_sample_protocol.py`:规定一条 TQ sample 的 key、tag、Tensor 字段和 + EOS record; +- `transferqueue_bridge.py` 与 `tq_owner.py`:负责连接 Ray/TQ、写读清理样本和 + owner 生命周期。 + +缺少的是把预先生成的文本变成 DSpark 训练特征并发布到 TQ 的独立进程。新增的 +Producer 补上这一段,不引入第二套协议或 feature store。 + +## 数据流 + +```text +prompt/response JSONL + │ + │ 按文件顺序分配 sequence_no 和 sample_id + ▼ +Tokenizer + │ input_ids / loss_mask / feature window + ▼ +多个 vLLM endpoint(有界并发) + │ OpenAI completions 请求 → 临时 safetensors 文件 + ▼ +公共 hidden-state 转换函数 + │ DSpark DraftFeatureSample + SampleMetadata + ▼ +TransferQueue kv_put(一条输入记录对应一条 sample) + │ + ├─ put 成功:删除该请求的临时文件 + └─ 全部成功:写一个 EOS control record +``` + +Producer 在开始请求前会等到对应 `run_id` 的 `owner_ready` 控制记录。它不会创建 +Ray head、TQ Controller 或 storage backend;这些由 `verl-speco-tq-owner` 管理。 + +## 输入文件 + +输入只能是 JSONL:每个非空行是一个 JSON object,必须有字符串 `prompt` 与非空 +字符串 `response`。 + +```json +{"sample_id":"train-000017","prompt":"Question: 1 + 1 = ","response":"2"} +{"prompt":"Translate hello: ","response":"你好"} +``` + +- `sequence_no` 按非空行的文件顺序从 0 分配;并发完成顺序不会影响它。 +- `sample_id` 可选;省略时生成 `train-000000`、`train-000001` 等稳定值。 +- Producer tokenize `prompt` 和 `prompt + response`。后者必须以 prompt 的 token IDs + 为前缀;否则会报错,而不会猜测 response 的 loss-mask 边界。 +- `loss_mask` 中 prompt token 为 0,response token 为 1。 +- feature window 从 response 前一个 token 开始,长度由 + `max_feature_length` 限制;传给 vLLM 的 token IDs 截止于该 window 末端。 +- 其他 JSON 字段目前只作为 Producer 进程内来源元数据;第一版协议不会把它们写入 + TQ,所以 Consumer 不能读取这些字段。 + +## vLLM 与 hidden states + +Producer 使用 OpenAI-compatible completions API: + +```text +prompt= +max_tokens=1 +extra_body={"return_token_ids": true} +``` + +响应必须同时满足: + +1. 若返回 `choices[0].prompt_token_ids`,它必须等于请求的 token IDs; +2. `kv_transfer_params.hidden_states_path` 必须存在; +3. 该文件必须含 `token_ids` 和形状为 `[seq, layers, hidden]` 的 `hidden_states`。 + +vLLM 0.23 已内置满足这个合同的 `ExampleHiddenStatesConnector`。不需要 SpeCo +Mooncake connector。在线服务必须关闭 chunked prefill,并显式配置一个 Producer +可见的临时目录。例如: + +```bash +export MODEL_PATH=/path/to/target-model +export HIDDEN_STATES_DIR=/dev/shm/speco-hidden-states +mkdir -p "${HIDDEN_STATES_DIR}" + +vllm serve "${MODEL_PATH}" \ + --host 0.0.0.0 \ + --port 8000 \ + --speculative-config \ + '{"method":"extract_hidden_states","num_speculative_tokens":1,"draft_model_config":{"hf_config":{"eagle_aux_hidden_state_layer_ids":[1,9,17,25,33,36]}}}' \ + --kv-transfer-config \ + "{\"kv_connector\":\"ExampleHiddenStatesConnector\",\"kv_role\":\"kv_producer\",\"kv_connector_extra_config\":{\"shared_storage_path\":\"${HIDDEN_STATES_DIR}\",\"use_synchronization_lock\":true}}" \ + --no-enable-chunked-prefill +``` + +上面的 layer IDs 只是 Qwen3-4B 示例。实际值必须按 target 模型和训练配置确定; +DSpark L1 开启时,vLLM 列表是 auxiliary layer IDs 加 final layer,而 Producer 的 +`TARGET_LAYER_IDS` 只填写 auxiliary 部分。 + +官方 connector 使用持久存在的 `.lock` 文件和 `flock` 协调异步落盘。Producer +读取前等待文件锁释放;TQ `put_sample` 成功后同时删除 safetensors 和 `.lock`。 + +`feature_from_vllm_payload()` 是从旧 replay 路径提取出的公共纯函数。它校验 token +对齐、选择 feature rows、拼接 auxiliary layers;DSpark L1 开启时额外拼接 final +hidden state。旧 replay 路径仍通过薄封装调用此函数,避免两套转换规则。 + +## TQ 写入和失败语义 + +每个输入 record 只写一个协议 key: + +```text +drafter:v1::<12位sequence_no>: +``` + +写入顺序是严格的: + +```text +加载临时 safetensors +→ 校验并转换 +→ TQ kv_put +→ 删除临时文件 +``` + +因此: + +- `kv_put` 失败时临时文件保留,且 Producer 不写 EOS; +- 任一请求、转换或写入失败会停止整条 Producer,不做自动重试或 endpoint 熔断; +- 只有所有 sample 都发布完成,才写 `control:v1::eos`; +- 进程退出时只调用 `close_transfer_queue_client()`,不会调用全局 `tq.close()`, + 不会销毁共享 Controller。只有 owner 可以关闭 TQ。 + +`max_pending_samples` 是简单背压:当前 run 的 ready sample 数达到该阈值时,新的 +vLLM 请求会暂停,等待 Consumer 清理已成功训练的 key。 + +## 配置与启动 + +默认 Producer 配置位于 +`speco.standalone_tq_producer`,TQ 连接配置仍位于 +`actor_rollout_ref.rollout.drafter.training.transfer_queue`。 + +必须设置的 Producer 字段: + +| 字段 | 含义 | +| --- | --- | +| `input_path` | 上述 JSONL 文件 | +| `tokenizer_path` / `tokenizer_fingerprint` | 用于 tokenization 和 Consumer 合同校验 | +| `target_model_id` / `target_model_revision` | target checkpoint 身份 | +| `target_layer_ids` | auxiliary target layer IDs;DSpark L1 时 wire metadata 会额外写 `-1` 表示 final layer | +| `vllm_endpoints` / `vllm_model` | 一个或多个 OpenAI-compatible vLLM endpoint 与模型名 | + +必须与 owner/Consumer 一致的 TQ 字段: + +| 字段 | 固定要求 | +| --- | --- | +| `package_version` | `0.1.7` | +| `partition_id` | `speco_drafter_features` | +| `schema_version` | `1` | +| `run_id`、Ray address、Ray namespace | 三个进程必须相同 | + +使用 [run_dspark_tq_producer.sh](../examples/run_dspark_tq_producer.sh) 启动。它要求 +显式提供 `RAY_ADDRESS`、`SPECO_TQ_RUN_ID`、输入、target/tokenizer 身份、layers 和 +vLLM endpoints,并且只启动 Producer。 + +也可以通过安装后的命令入口运行: + +```bash +verl-speco-tq-producer \ + actor_rollout_ref.rollout.drafter.training.transfer_queue.enable=true \ + actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.address= \ + actor_rollout_ref.rollout.drafter.training.transfer_queue.run_id= \ + speco.standalone_tq_producer.input_path= \ + speco.standalone_tq_producer.tokenizer_path= \ + speco.standalone_tq_producer.tokenizer_fingerprint= \ + speco.standalone_tq_producer.target_model_id= \ + speco.standalone_tq_producer.target_model_revision= \ + speco.standalone_tq_producer.target_layer_ids='[2,8,14,20,26]' \ + speco.standalone_tq_producer.vllm_endpoints='[http://node0:8000/v1]' \ + speco.standalone_tq_producer.vllm_model= +``` + +完整生命周期顺序仍是:Ray/TQ backend → TQ owner → Consumer → Producer → Consumer +drain → owner shutdown。Producer 完成不代表训练完成,EOS 只表示不会再有新样本。 + +## 测试覆盖与未验证项 + +新增测试覆盖:JSONL 解析与 token 边界、多个 endpoint 的并发限制、ready 队列背压、 +成功时 sample 后 EOS 与临时文件删除、失败时无 EOS 且保留临时文件,以及旧 EAGLE3 +转换路径仍可复用公共函数。 + +这些测试使用 fake vLLM/TQ。真实 Ray + TransferQueue + vLLM 的多进程 +联调没有在当前环境执行;运行前仍需确认 vLLM 版本能返回上述 +`hidden_states_path` 以及 TQ 0.1.7 依赖环境可用。 diff --git a/examples/run_dspark_tq_producer.sh b/examples/run_dspark_tq_producer.sh new file mode 100644 index 00000000..ae8f62b5 --- /dev/null +++ b/examples/run_dspark_tq_producer.sh @@ -0,0 +1,40 @@ +#!/usr/bin/env bash +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +set -euo pipefail + +: "${RAY_ADDRESS:?Set RAY_ADDRESS to the running Ray head}" +: "${SPECO_TQ_RUN_ID:?Set SPECO_TQ_RUN_ID to the owner/consumer run id}" +: "${PRODUCER_INPUT_PATH:?Set PRODUCER_INPUT_PATH to prompt/response JSONL}" +: "${TARGET_MODEL_PATH:?Set TARGET_MODEL_PATH to the target model id/path}" +: "${TARGET_MODEL_REVISION:?Set TARGET_MODEL_REVISION to a revision or checksum}" +: "${TOKENIZER_PATH:?Set TOKENIZER_PATH to the tokenizer id/path}" +: "${TOKENIZER_FINGERPRINT:?Set TOKENIZER_FINGERPRINT to a verified fingerprint}" +: "${TARGET_LAYER_IDS:?Set TARGET_LAYER_IDS as a Hydra list, for example '[2,8,14,20,26]'}" +: "${VLLM_ENDPOINTS:?Set VLLM_ENDPOINTS as a Hydra list, for example '[http://node0:8000/v1]'}" + +python -m verl_speco.standalone_tq_producer \ + actor_rollout_ref.rollout.drafter.speculative_algorithm=DSPARK \ + actor_rollout_ref.rollout.drafter.training.transfer_queue.enable=true \ + actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.address="${RAY_ADDRESS}" \ + actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.namespace=speco-drafter \ + actor_rollout_ref.rollout.drafter.training.transfer_queue.run_id="${SPECO_TQ_RUN_ID}" \ + speco.standalone_tq_producer.input_path="${PRODUCER_INPUT_PATH}" \ + speco.standalone_tq_producer.target_model_id="${TARGET_MODEL_PATH}" \ + speco.standalone_tq_producer.target_model_revision="${TARGET_MODEL_REVISION}" \ + speco.standalone_tq_producer.tokenizer_path="${TOKENIZER_PATH}" \ + speco.standalone_tq_producer.tokenizer_fingerprint="${TOKENIZER_FINGERPRINT}" \ + speco.standalone_tq_producer.target_layer_ids="${TARGET_LAYER_IDS}" \ + speco.standalone_tq_producer.vllm_endpoints="${VLLM_ENDPOINTS}" \ + speco.standalone_tq_producer.vllm_model="${TARGET_MODEL_PATH}" diff --git a/pyproject.toml b/pyproject.toml index 445e7810..d87e9806 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -27,6 +27,7 @@ verl-speco = "verl_speco.main:main" verl-speco-draft-train = "verl_speco.draft_train_launcher:main" verl-speco-inspect-features = "verl_speco.inspect_feature_store:main" verl-speco-tq-owner = "verl_speco.tq_owner:main" +verl-speco-tq-producer = "verl_speco.standalone_tq_producer:main" [tool.setuptools.dynamic] version = { attr = "verl_speco.__version__" } diff --git a/tests/special_sanity/check_example_naming.py b/tests/special_sanity/check_example_naming.py index da49fb0c..ac375191 100644 --- a/tests/special_sanity/check_example_naming.py +++ b/tests/special_sanity/check_example_naming.py @@ -43,7 +43,12 @@ STANDALONE_SUFFIX = ("separate", "training") DEFAULT_IGNORE_DIRS: tuple[str, ...] = () -DEFAULT_IGNORE_FILES: tuple[str, ...] = () +# Dedicated TQ role launchers are lifecycle utilities, not rollout/training +# examples, so the model/drafter/rollout naming grammar does not apply. +DEFAULT_IGNORE_FILES: tuple[str, ...] = ( + "examples/run_dspark_tq_owner.sh", + "examples/run_dspark_tq_producer.sh", +) def _split_tokens(stem: str) -> list[str]: diff --git a/verl_speco/config/speco_base.yaml b/verl_speco/config/speco_base.yaml index bce79dfb..82517aed 100644 --- a/verl_speco/config/speco_base.yaml +++ b/verl_speco/config/speco_base.yaml @@ -21,6 +21,31 @@ speco: task_runner: verl_speco.integration.task_runner.SpecoTaskRunner ray_trainer: verl_speco.trainer.speco_ray_trainer.SpecoRayPPOTrainer + # One-process standalone Producer. It reads strict prompt/response JSONL, + # requests target hidden states from vLLM, and publishes one TQ key per row. + standalone_tq_producer: + input_path: null + tokenizer_path: null + tokenizer_fingerprint: null + target_model_id: null + target_model_revision: null + target_layer_ids: null + hidden_dtype: bfloat16 + trust_remote_code: false + vllm_endpoints: + - http://localhost:8000/v1 + vllm_model: null + request_timeout: 120 + max_inflight_requests: 16 + per_endpoint_concurrency: 4 + input_queue_size: 32 + publish_queue_size: 16 + max_pending_samples: 1024 + pending_poll_interval_seconds: 0.5 + owner_ready_timeout_seconds: 120 + max_sequence_length: 8192 + max_feature_length: 512 + actor_rollout_ref: rollout: drafter: diff --git a/verl_speco/producer/__init__.py b/verl_speco/producer/__init__.py new file mode 100644 index 00000000..00265cdf --- /dev/null +++ b/verl_speco/producer/__init__.py @@ -0,0 +1,14 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Standalone target-feature Producer components.""" diff --git a/verl_speco/producer/input_reader.py b/verl_speco/producer/input_reader.py new file mode 100644 index 00000000..ba77fbda --- /dev/null +++ b/verl_speco/producer/input_reader.py @@ -0,0 +1,206 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Streaming JSONL input and token preparation for the standalone Producer.""" + +from __future__ import annotations + +import json +import os +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Iterator, Mapping + +import torch + + +@dataclass(frozen=True) +class InputRecord: + sequence_no: int + sample_id: str + prompt: str + response: str + source_metadata: dict[str, Any] + + +@dataclass(frozen=True) +class TokenizedRequest: + sequence_no: int + sample_id: str + input_ids: torch.Tensor + loss_mask: torch.Tensor + position_ids: torch.Tensor + feature_positions: torch.Tensor + draft_position_ids: torch.Tensor + source_metadata: dict[str, Any] + + @property + def prompt_token_ids(self) -> list[int]: + feature_end = int(self.feature_positions[-1].item()) + 1 + return self.input_ids[:feature_end].detach().cpu().long().tolist() + + +def iter_input_records(path: str | os.PathLike[str]) -> Iterator[InputRecord]: + """Yield one strict prompt/response record per non-empty JSONL line.""" + + input_path = Path(path) + if not input_path.is_file(): + raise FileNotFoundError(f"Producer input JSONL not found: {input_path}") + sequence_no = 0 + with input_path.open("r", encoding="utf-8") as input_file: + for line_number, line in enumerate(input_file, start=1): + if not line.strip(): + continue + try: + payload = json.loads(line) + except json.JSONDecodeError as exc: + raise ValueError( + f"Invalid JSON object at {input_path}:{line_number}: {exc.msg}" + ) from exc + if not isinstance(payload, dict): + raise ValueError( + f"Producer input at {input_path}:{line_number} must be a JSON object" + ) + prompt = payload.get("prompt") + response = payload.get("response") + if not isinstance(prompt, str): + raise ValueError( + f"Producer input at {input_path}:{line_number} requires string field 'prompt'" + ) + if not isinstance(response, str) or not response: + raise ValueError( + f"Producer input at {input_path}:{line_number} requires non-empty string field 'response'" + ) + sample_id = payload.get("sample_id", f"train-{sequence_no:06d}") + if not isinstance(sample_id, str) or not sample_id: + raise ValueError( + f"Producer input at {input_path}:{line_number} has invalid sample_id" + ) + source_metadata = { + key: value + for key, value in payload.items() + if key not in {"prompt", "response", "sample_id"} + } + yield InputRecord( + sequence_no=sequence_no, + sample_id=sample_id, + prompt=prompt, + response=response, + source_metadata=source_metadata, + ) + sequence_no += 1 + + +def build_loss_mask(input_ids: torch.Tensor, prompt_length: int) -> torch.Tensor: + sequence_length = int(input_ids.numel()) + if prompt_length < 0 or prompt_length > sequence_length: + raise ValueError( + f"prompt_length must be within [0, {sequence_length}], got {prompt_length}" + ) + mask = torch.ones(sequence_length, dtype=torch.float32) + mask[:prompt_length] = 0 + return mask + + +def tokenize_record( + record: InputRecord, + tokenizer: Any, + config: Mapping[str, Any] | Any, +) -> TokenizedRequest: + """Tokenize existing prompt/response text without generating new tokens.""" + + prompt_ids = _token_ids(tokenizer(record.prompt, add_special_tokens=False)) + full_ids = _token_ids( + tokenizer(record.prompt + record.response, add_special_tokens=False) + ) + if full_ids[: len(prompt_ids)] != prompt_ids: + raise ValueError( + f"Producer sample {record.sample_id!r} has an unstable tokenizer boundary " + "between prompt and response; prompt token IDs are not a prefix of full token IDs" + ) + if len(full_ids) <= len(prompt_ids): + raise ValueError( + f"Producer sample {record.sample_id!r} produced no response tokens" + ) + input_ids = torch.tensor(full_ids, dtype=torch.int64) + if int(input_ids.numel()) <= 0: + raise ValueError( + f"Producer sample {record.sample_id!r} produced no input tokens" + ) + + max_sequence_length = int(_config_value(config, "max_sequence_length", 0) or 0) + if max_sequence_length > 0 and int(input_ids.numel()) > max_sequence_length: + raise ValueError( + f"Producer sample {record.sample_id!r} has {int(input_ids.numel())} tokens, " + f"exceeding max_sequence_length={max_sequence_length}" + ) + loss_mask = build_loss_mask(input_ids, len(prompt_ids)) + position_ids = torch.arange(int(input_ids.numel()), dtype=torch.int64) + + feature_start = max(len(prompt_ids) - 1, 0) + feature_end = int(input_ids.numel()) + max_feature_length = int(_config_value(config, "max_feature_length", 0) or 0) + if max_feature_length == 1: + raise ValueError("max_feature_length must be 0 or at least 2") + if max_feature_length > 1: + feature_end = min(feature_start + max_feature_length, feature_end) + feature_positions = torch.arange(feature_start, feature_end, dtype=torch.int64) + if int(feature_positions.numel()) <= 0: + raise ValueError( + f"Producer sample {record.sample_id!r} has an empty feature window" + ) + draft_position_ids = position_ids[feature_start:feature_end] + 1 + return TokenizedRequest( + sequence_no=record.sequence_no, + sample_id=record.sample_id, + input_ids=input_ids, + loss_mask=loss_mask, + position_ids=position_ids, + feature_positions=feature_positions, + draft_position_ids=draft_position_ids, + source_metadata=dict(record.source_metadata), + ) + + +def _token_ids(encoding: Any) -> list[int]: + value = ( + encoding.get("input_ids") + if isinstance(encoding, Mapping) + else encoding.input_ids + ) + if value is None: + raise ValueError("Tokenizer result is missing input_ids") + if torch.is_tensor(value): + value = value.detach().cpu().reshape(-1).tolist() + if not isinstance(value, (list, tuple)): + raise TypeError("Tokenizer input_ids must be a tensor, list, or tuple") + if value and isinstance(value[0], (list, tuple)): + if len(value) != 1: + raise ValueError("Tokenizer returned more than one sequence for one input") + value = value[0] + return [int(token_id) for token_id in value] + + +def _config_value(config: Any, key: str, default: Any = None) -> Any: + if isinstance(config, Mapping): + return config.get(key, default) + return getattr(config, key, default) + + +__all__ = [ + "InputRecord", + "TokenizedRequest", + "build_loss_mask", + "iter_input_records", + "tokenize_record", +] diff --git a/verl_speco/producer/vllm_feature_client.py b/verl_speco/producer/vllm_feature_client.py new file mode 100644 index 00000000..4b793a66 --- /dev/null +++ b/verl_speco/producer/vllm_feature_client.py @@ -0,0 +1,236 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Bounded asynchronous vLLM hidden-state requests for the Producer.""" + +from __future__ import annotations + +import asyncio +import errno +import inspect +import os +import time +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Mapping, Sequence + + +@dataclass(frozen=True) +class VllmEndpoint: + base_url: str + max_concurrency: int + + def __post_init__(self) -> None: + if not self.base_url: + raise ValueError("VllmEndpoint.base_url must not be empty") + if self.max_concurrency <= 0: + raise ValueError("VllmEndpoint.max_concurrency must be positive") + + +@dataclass(frozen=True) +class VllmResponse: + hidden_states_path: str + endpoint_url: str + + +@dataclass(frozen=True) +class RawVllmFeature: + payload: dict[str, Any] + temporary_path: str + endpoint_url: str + byte_size: int + + +@dataclass +class _EndpointState: + endpoint: VllmEndpoint + client: Any + semaphore: asyncio.Semaphore + inflight: int = 0 + requests: int = 0 + + +async def request_prefill( + endpoint: VllmEndpoint, + client: Any, + prompt_token_ids: list[int], + *, + model: str, + timeout: float, +) -> VllmResponse: + response = await client.completions.create( + model=model, + prompt=prompt_token_ids, + max_tokens=1, + extra_body={"return_token_ids": True}, + timeout=timeout, + ) + choices = getattr(response, "choices", None) or [] + if choices: + actual = getattr(choices[0], "prompt_token_ids", None) + if actual is not None and list(actual) != prompt_token_ids: + raise ValueError("vLLM prompt_token_ids mismatch") + params = getattr(response, "kv_transfer_params", None) + if not isinstance(params, Mapping): + raise ValueError("vLLM response missing kv_transfer_params") + path = params.get("hidden_states_path") + if not path: + raise ValueError("vLLM response missing hidden_states_path") + return VllmResponse(os.fspath(path), endpoint.base_url) + + +def load_hidden_state_result(response: VllmResponse) -> RawVllmFeature: + try: + from safetensors.torch import load_file + except ImportError as exc: + raise RuntimeError("vLLM Producer requires safetensors") from exc + path = Path(response.hidden_states_path) + _wait_for_lock(Path(f"{path}.lock")) + if not path.is_file(): + raise FileNotFoundError(f"vLLM hidden-states file not found: {path}") + return RawVllmFeature( + payload=dict(load_file(str(path), device="cpu")), + temporary_path=str(path), + endpoint_url=response.endpoint_url, + byte_size=int(path.stat().st_size), + ) + + +def delete_temporary_result(raw: RawVllmFeature) -> None: + path = Path(raw.temporary_path) + path.unlink(missing_ok=True) + Path(f"{path}.lock").unlink(missing_ok=True) + + +def choose_endpoint(states: Sequence[_EndpointState]) -> _EndpointState: + if not states: + raise RuntimeError("No vLLM endpoints are configured") + return min(states, key=lambda state: (state.inflight, state.requests)) + + +class VllmFeatureClientPool: + def __init__( + self, + endpoints: Sequence[VllmEndpoint], + *, + model: str, + max_inflight_requests: int, + request_timeout: float, + ) -> None: + if not endpoints: + raise ValueError("At least one vLLM endpoint is required") + if max_inflight_requests <= 0: + raise ValueError("max_inflight_requests must be positive") + if not model: + raise ValueError("vllm_model must not be empty") + self.endpoints = list(endpoints) + self.model = model + self.request_timeout = float(request_timeout) + self._global_semaphore = asyncio.Semaphore(max_inflight_requests) + self._states: list[_EndpointState] = [] + + async def start(self) -> None: + if self._states: + return + try: + from openai import AsyncOpenAI + except ImportError as exc: + raise RuntimeError("vLLM Producer requires the openai package") from exc + self._states = [ + _EndpointState( + endpoint=endpoint, + client=AsyncOpenAI( + base_url=endpoint.base_url, + api_key="EMPTY", + max_retries=0, + ), + semaphore=asyncio.Semaphore(endpoint.max_concurrency), + ) + for endpoint in self.endpoints + ] + + async def prefill(self, request: Any) -> RawVllmFeature: + if not self._states: + raise RuntimeError("VllmFeatureClientPool.start() must be called first") + state = choose_endpoint(self._states) + state.inflight += 1 + try: + async with self._global_semaphore, state.semaphore: + response = await request_prefill( + state.endpoint, + state.client, + list(request.prompt_token_ids), + model=self.model, + timeout=self.request_timeout, + ) + raw = await asyncio.to_thread(load_hidden_state_result, response) + state.requests += 1 + return raw + finally: + state.inflight = max(state.inflight - 1, 0) + + async def close(self) -> None: + states, self._states = self._states, [] + for state in states: + close = getattr(state.client, "close", None) + if not callable(close): + continue + result = close() + if inspect.isawaitable(result): + await result + + +def _wait_for_lock(lock_path: Path, timeout: float = 30.0) -> None: + if not lock_path.exists(): + return + try: + import fcntl + except ImportError: + # vLLM's file connector is Linux-only. Keep the old existence-based + # fallback for dependency-light tests on other platforms. + deadline = time.monotonic() + timeout + while lock_path.exists(): + if time.monotonic() >= deadline: + raise TimeoutError( + f"Timed out waiting for vLLM hidden-state lock: {lock_path}" + ) + time.sleep(0.01) + return + + deadline = time.monotonic() + timeout + with lock_path.open("rb") as lock_file: + while True: + try: + fcntl.flock(lock_file.fileno(), fcntl.LOCK_SH | fcntl.LOCK_NB) + fcntl.flock(lock_file.fileno(), fcntl.LOCK_UN) + return + except OSError as exc: + if exc.errno not in {errno.EACCES, errno.EAGAIN}: + raise + if time.monotonic() >= deadline: + raise TimeoutError( + f"Timed out waiting for vLLM hidden-state lock: {lock_path}" + ) from exc + time.sleep(0.01) + + +__all__ = [ + "RawVllmFeature", + "VllmEndpoint", + "VllmFeatureClientPool", + "VllmResponse", + "choose_endpoint", + "delete_temporary_result", + "load_hidden_state_result", + "request_prefill", +] diff --git a/verl_speco/standalone_tq_producer.py b/verl_speco/standalone_tq_producer.py new file mode 100644 index 00000000..395c8bc2 --- /dev/null +++ b/verl_speco/standalone_tq_producer.py @@ -0,0 +1,451 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Standalone vLLM target-feature Producer writing directly to TransferQueue.""" + +from __future__ import annotations + +import asyncio +import logging +from dataclasses import dataclass +from typing import Any, Mapping + +import torch + +from verl_speco.integration import transferqueue_bridge as default_transport +from verl_speco.integration.oldlogprob_layer_ids import ( + resolve_drafter_hidden_states_layout, +) +from verl_speco.producer.input_reader import ( + TokenizedRequest, + iter_input_records, + tokenize_record, +) +from verl_speco.producer.vllm_feature_client import ( + RawVllmFeature, + VllmEndpoint, + VllmFeatureClientPool, + delete_temporary_result, +) +from verl_speco.trainer.feature_store import DraftFeatureSample +from verl_speco.trainer.target_feature_replay import ( + FeatureContract, + feature_from_vllm_payload, +) +from verl_speco.transport.drafter_sample_protocol import ( + DRAFTER_TQ_PARTITION, + PROTOCOL_SCHEMA_VERSION, + SampleMetadata, + encode_sample, + make_eos_record, + make_ready_tag, + make_sample_key, +) + + +logger = logging.getLogger(__name__) +_INPUT_DONE = object() +_PUBLISH_DONE = object() + + +@dataclass +class ProducerStats: + input_count: int = 0 + published_count: int = 0 + failed_count: int = 0 + pending_bytes: int = 0 + + +@dataclass(frozen=True) +class PreparedFeature: + request: TokenizedRequest + raw: RawVllmFeature + sample: DraftFeatureSample + metadata: SampleMetadata + + +async def publish_one(result: PreparedFeature, transport: Any) -> str: + """Publish one sample and delete its temporary file only after TQ succeeds.""" + + key = make_sample_key(result.metadata) + fields = encode_sample(result.sample, result.metadata) + tag = make_ready_tag(result.metadata) + await asyncio.to_thread(transport.put_sample, key, fields, tag=tag) + delete_temporary_result(result.raw) + return key + + +def validate_producer_config(config: Any) -> None: + producer_cfg, training_cfg, tq_cfg = _config_sections(config) + required = ( + "input_path", + "tokenizer_path", + "tokenizer_fingerprint", + "target_model_id", + "target_model_revision", + "vllm_model", + ) + missing = [name for name in required if not producer_cfg.get(name)] + if missing: + raise ValueError(f"standalone_tq_producer missing required fields: {missing}") + endpoints = producer_cfg.get("vllm_endpoints") + if not isinstance(endpoints, list) or not endpoints or not all(endpoints): + raise ValueError( + "standalone_tq_producer.vllm_endpoints must be a non-empty list" + ) + target_layer_ids = producer_cfg.get("target_layer_ids") + if not isinstance(target_layer_ids, list) or not target_layer_ids: + raise ValueError( + "standalone_tq_producer.target_layer_ids must be a non-empty list" + ) + if str(training_cfg.get("speculative_algorithm", "DSPARK")).upper() != "DSPARK": + raise ValueError("Standalone TQ Producer currently supports only DSPARK") + if bool(training_cfg.get("use_logits", False)): + raise ValueError("Standalone TQ Producer does not support use_logits=true") + if int(tq_cfg.get("schema_version", 0)) != PROTOCOL_SCHEMA_VERSION: + raise ValueError( + f"transfer_queue.schema_version must be {PROTOCOL_SCHEMA_VERSION}" + ) + if tq_cfg.get("package_version") != "0.1.7": + raise ValueError("transfer_queue.package_version must be '0.1.7'") + if tq_cfg.get("partition_id") != DRAFTER_TQ_PARTITION: + raise ValueError( + f"transfer_queue.partition_id must be {DRAFTER_TQ_PARTITION!r}" + ) + if not tq_cfg.get("run_id"): + raise ValueError("transfer_queue.run_id must not be empty") + ray_cfg = tq_cfg.get("ray") or {} + if not isinstance(ray_cfg, Mapping) or not ray_cfg.get("address"): + raise ValueError( + "transfer_queue.ray.address must point to a running Ray cluster" + ) + positive_fields = ( + "max_inflight_requests", + "per_endpoint_concurrency", + "input_queue_size", + "publish_queue_size", + "max_pending_samples", + ) + invalid = [name for name in positive_fields if int(producer_cfg.get(name, 0)) <= 0] + if invalid: + raise ValueError(f"standalone_tq_producer fields must be positive: {invalid}") + + +async def run_producer( + config: Any, + *, + transport: Any = default_transport, + tokenizer: Any | None = None, + client_pool: Any | None = None, +) -> ProducerStats: + """Run the bounded input -> vLLM -> TQ pipeline and publish EOS on success.""" + + validate_producer_config(config) + producer_cfg, drafter_cfg, tq_cfg = _config_sections(config) + run_id = str(tq_cfg["run_id"]) + stats = ProducerStats() + connected = False + pool = client_pool + try: + if not transport.configure_transfer_queue(tq_cfg): + raise RuntimeError("Standalone TQ Producer requires TransferQueue==0.1.7") + ray_cfg = tq_cfg["ray"] + transport.connect_ray_cluster( + str(ray_cfg["address"]), + str(ray_cfg["namespace"]) if ray_cfg.get("namespace") else None, + ) + transport.connect_transfer_queue_client() + connected = True + await _wait_for_owner_ready( + transport, + run_id, + timeout=float(producer_cfg["owner_ready_timeout_seconds"]), + poll_interval=float(producer_cfg["pending_poll_interval_seconds"]), + ) + + if tokenizer is None: + tokenizer = await asyncio.to_thread(_load_tokenizer, producer_cfg) + if pool is None: + endpoint_concurrency = int(producer_cfg["per_endpoint_concurrency"]) + pool = VllmFeatureClientPool( + [ + VllmEndpoint(str(url).rstrip("/"), endpoint_concurrency) + for url in producer_cfg["vllm_endpoints"] + ], + model=str(producer_cfg["vllm_model"]), + max_inflight_requests=int(producer_cfg["max_inflight_requests"]), + request_timeout=float(producer_cfg["request_timeout"]), + ) + await pool.start() + + feature_contract = FeatureContract( + algorithm="DSPARK", + target_layer_ids=[int(value) for value in producer_cfg["target_layer_ids"]], + hidden_states_layout=resolve_drafter_hidden_states_layout( + "DSPARK", drafter_cfg + ), + dtype=_parse_dtype(producer_cfg["hidden_dtype"]), + target_model_id=str(producer_cfg["target_model_id"]), + target_model_revision=str(producer_cfg["target_model_revision"]), + tokenizer_fingerprint=str(producer_cfg["tokenizer_fingerprint"]), + use_logits=False, + ) + worker_count = int(producer_cfg["max_inflight_requests"]) + input_queue: asyncio.Queue[Any] = asyncio.Queue( + maxsize=int(producer_cfg["input_queue_size"]) + ) + publish_queue: asyncio.Queue[Any] = asyncio.Queue( + maxsize=int(producer_cfg["publish_queue_size"]) + ) + + async def read_inputs() -> None: + for record in iter_input_records(str(producer_cfg["input_path"])): + request = tokenize_record(record, tokenizer, producer_cfg) + await input_queue.put(request) + stats.input_count += 1 + for _ in range(worker_count): + await input_queue.put(_INPUT_DONE) + + async def request_worker() -> None: + while True: + request = await input_queue.get() + if request is _INPUT_DONE: + await publish_queue.put(_PUBLISH_DONE) + return + await _wait_for_pending_capacity( + transport, + run_id, + max_pending_samples=int(producer_cfg["max_pending_samples"]), + poll_interval=float(producer_cfg["pending_poll_interval_seconds"]), + ) + raw = await pool.prefill(request) + stats.pending_bytes += int(raw.byte_size) + sample = feature_from_vllm_payload(raw, request, feature_contract) + await publish_queue.put( + PreparedFeature( + request=request, + raw=raw, + sample=sample, + metadata=_sample_metadata( + request, sample, feature_contract, run_id, tq_cfg + ), + ) + ) + + async def publish_results() -> None: + finished_workers = 0 + while finished_workers < worker_count: + result = await publish_queue.get() + if result is _PUBLISH_DONE: + finished_workers += 1 + continue + await publish_one(result, transport) + stats.published_count += 1 + stats.pending_bytes = max( + stats.pending_bytes - int(result.raw.byte_size), 0 + ) + + tasks = [asyncio.create_task(read_inputs())] + tasks.extend(asyncio.create_task(request_worker()) for _ in range(worker_count)) + tasks.append(asyncio.create_task(publish_results())) + done, pending = await asyncio.wait(tasks, return_when=asyncio.FIRST_EXCEPTION) + failure = next( + (task.exception() for task in done if task.exception() is not None), None + ) + if failure is not None: + stats.failed_count += 1 + for task in pending: + task.cancel() + await asyncio.gather(*pending, return_exceptions=True) + raise failure + await asyncio.gather(*pending) + + eos_key, eos_fields, eos_tag = make_eos_record(run_id, stats.published_count) + await asyncio.to_thread(transport.put_sample, eos_key, eos_fields, tag=eos_tag) + logger.info( + "Standalone TQ Producer completed inputs=%s published=%s", + stats.input_count, + stats.published_count, + ) + return stats + finally: + if pool is not None: + await pool.close() + if connected: + transport.close_transfer_queue_client() + + +async def _wait_for_owner_ready( + transport: Any, + run_id: str, + *, + timeout: float, + poll_interval: float, +) -> None: + deadline = asyncio.get_running_loop().time() + timeout + while True: + records = await asyncio.to_thread(transport.list_samples) + if any( + tag.get("record_type") == "control" + and tag.get("status") == "owner_ready" + and tag.get("run_id") == run_id + for tag in records.values() + ): + return + if asyncio.get_running_loop().time() >= deadline: + raise TimeoutError( + f"Timed out waiting for TQ owner_ready for run_id={run_id!r}" + ) + await asyncio.sleep(poll_interval) + + +async def _wait_for_pending_capacity( + transport: Any, + run_id: str, + *, + max_pending_samples: int, + poll_interval: float, +) -> None: + while True: + records = await asyncio.to_thread(transport.list_samples) + ready_count = sum( + 1 + for tag in records.values() + if tag.get("record_type") == "sample" + and tag.get("status") == "ready" + and tag.get("run_id") == run_id + ) + if ready_count < max_pending_samples: + return + await asyncio.sleep(poll_interval) + + +def _sample_metadata( + request: TokenizedRequest, + sample: DraftFeatureSample, + contract: FeatureContract, + run_id: str, + tq_cfg: Mapping[str, Any], +) -> SampleMetadata: + hidden = sample.hidden_states + if not torch.is_tensor(hidden): + raise TypeError("Standalone TQ Producer requires dense hidden_states") + feature_start = int(sample.metadata["feature_start"]) + feature_end = int(sample.metadata["feature_end"]) + wire_layer_ids = list(contract.target_layer_ids) + if contract.hidden_states_layout == "dflash_aux_plus_last": + wire_layer_ids.append(-1) + return SampleMetadata( + schema_version=int(tq_cfg["schema_version"]), + run_id=run_id, + sample_id=request.sample_id, + sequence_no=request.sequence_no, + algorithm=contract.algorithm, + target_model_id=contract.target_model_id, + target_model_revision=str(contract.target_model_revision or ""), + tokenizer_fingerprint=contract.tokenizer_fingerprint, + target_layer_ids=wire_layer_ids, + hidden_states_layout=contract.hidden_states_layout, + hidden_dtype=str(hidden.dtype).removeprefix("torch."), + hidden_shape=[int(value) for value in hidden.shape], + feature_length=int(hidden.size(0)), + full_sequence_length=int(request.input_ids.numel()), + feature_start=feature_start, + feature_end=feature_end, + use_logits=contract.use_logits, + ) + + +def _config_sections( + config: Any, +) -> tuple[dict[str, Any], dict[str, Any], dict[str, Any]]: + plain = _plain_config(config) + try: + producer_cfg = plain["speco"]["standalone_tq_producer"] + drafter = plain["actor_rollout_ref"]["rollout"]["drafter"] + training_cfg = drafter["training"] + tq_cfg = training_cfg["transfer_queue"] + except (KeyError, TypeError) as exc: + raise ValueError(f"Producer configuration missing section {exc}") from exc + if not all( + isinstance(value, dict) for value in (producer_cfg, training_cfg, tq_cfg) + ): + raise TypeError("Producer configuration sections must resolve to mappings") + return ( + producer_cfg, + {**training_cfg, "speculative_algorithm": drafter.get("speculative_algorithm")}, + tq_cfg, + ) + + +def _plain_config(config: Any) -> dict[str, Any]: + value = config + try: + from omegaconf import OmegaConf + + if OmegaConf.is_config(config): + value = OmegaConf.to_container(config, resolve=True) + except ImportError: + pass + if not isinstance(value, Mapping): + raise TypeError("Producer configuration must be a mapping") + return dict(value) + + +def _load_tokenizer(config: Mapping[str, Any]) -> Any: + try: + from transformers import AutoTokenizer + except ImportError as exc: + raise RuntimeError("Standalone TQ Producer requires transformers") from exc + return AutoTokenizer.from_pretrained( + str(config["tokenizer_path"]), + trust_remote_code=bool(config.get("trust_remote_code", False)), + ) + + +def _parse_dtype(value: Any) -> torch.dtype: + name = str(value).strip().lower().removeprefix("torch.") + aliases = {"fp32": "float32", "fp16": "float16", "bf16": "bfloat16"} + dtype = getattr(torch, aliases.get(name, name), None) + if not isinstance(dtype, torch.dtype): + raise ValueError(f"Unsupported standalone_tq_producer.hidden_dtype={value!r}") + return dtype + + +def _hydra_main(config: Any) -> None: + logging.basicConfig(level=logging.INFO) + asyncio.run(run_producer(config)) + + +def main() -> None: + try: + import hydra + except ImportError as exc: + raise RuntimeError("Standalone TQ Producer requires hydra-core") from exc + hydra.main(config_path="config", config_name="speco_base", version_base=None)( + _hydra_main + )() + + +if __name__ == "__main__": + main() + + +__all__ = [ + "PreparedFeature", + "ProducerStats", + "main", + "publish_one", + "run_producer", + "validate_producer_config", +] diff --git a/verl_speco/trainer/target_feature_replay.py b/verl_speco/trainer/target_feature_replay.py index 4d590380..9e3d50e5 100644 --- a/verl_speco/trainer/target_feature_replay.py +++ b/verl_speco/trainer/target_feature_replay.py @@ -25,7 +25,7 @@ import time from dataclasses import dataclass from pathlib import Path -from typing import Any, Iterable, cast +from typing import Any, Iterable, Mapping, cast import torch from torch import nn @@ -45,6 +45,22 @@ class MooncakeReplayDescriptor: key: str +@dataclass(frozen=True) +class FeatureContract: + """Explicit inputs for converting one vLLM payload into a training sample.""" + + algorithm: str + target_layer_ids: list[int] + hidden_states_layout: str + dtype: torch.dtype + target_model_id: str + target_model_revision: str | None + tokenizer_fingerprint: str + use_logits: bool = False + target_config_fingerprint: str | None = None + source: str = "standalone_tq_producer" + + @dataclass class _VllmEndpointState: index: int @@ -355,6 +371,139 @@ def metrics(self) -> dict[str, float]: } +def feature_from_vllm_payload( + payload: Mapping[str, Any] | Any, + request: DraftReplaySample | Any, + feature_config: FeatureContract, +) -> DraftFeatureSample: + """Pure vLLM payload conversion shared by replay and standalone Producer.""" + + values = getattr(payload, "payload", payload) + if not isinstance(values, Mapping): + raise TypeError("vLLM hidden-states payload must be a mapping") + token_ids = values.get("token_ids") + hidden = values.get("hidden_states") + if not torch.is_tensor(token_ids) or not torch.is_tensor(hidden): + raise ValueError( + "vLLM hidden-states payload must contain token_ids and hidden_states" + ) + feature_positions = request.feature_positions.detach().cpu().long() + if int(feature_positions.numel()) <= 0: + raise ValueError("vLLM feature positions must not be empty") + feature_end_for_request = int(feature_positions[-1].item()) + 1 + expected_prompt_ids = ( + list(request.prompt_token_ids) + if hasattr(request, "prompt_token_ids") + else request.input_ids[:feature_end_for_request].detach().cpu().long().tolist() + ) + if token_ids.detach().cpu().long().tolist() != expected_prompt_ids: + raise ValueError("vLLM hidden-states token_ids do not match replay input") + if hidden.dim() != 3: + raise ValueError( + "vLLM hidden_states must have shape [seq, layers, hidden], " + f"got {tuple(hidden.shape)}" + ) + + algorithm = str(feature_config.algorithm).strip().upper() + if algorithm not in {"EAGLE3", "DFLASH", "DSPARK"}: + raise ValueError(f"Unsupported vLLM feature algorithm {algorithm!r}") + target_layer_ids = [int(layer_id) for layer_id in feature_config.target_layer_ids] + if not target_layer_ids: + raise ValueError("FeatureContract.target_layer_ids must not be empty") + hidden_layout = str(feature_config.hidden_states_layout) + if hidden_layout not in { + "eagle3_aux_plus_last", + "dflash_aux", + "dflash_aux_plus_last", + }: + raise ValueError(f"Unsupported vLLM hidden_states_layout {hidden_layout!r}") + + hidden_position_offset = max(len(expected_prompt_ids) - int(hidden.size(0)), 0) + include_final = hidden_layout in { + "eagle3_aux_plus_last", + "dflash_aux_plus_last", + } + required_layers = len(target_layer_ids) + (1 if include_final else 0) + if int(hidden.size(1)) < required_layers: + raise ValueError( + "vLLM hidden_states layer count is too small: " + f"got {int(hidden.size(1))}, need at least {required_layers}. " + "Start vLLM with target layer ids plus the final layer when the " + "training layout needs last hidden states." + ) + relative_positions = feature_positions - hidden_position_offset + keep_mask = (relative_positions >= 0) & (relative_positions < int(hidden.size(0))) + filtered = not bool(keep_mask.all().item()) + if filtered: + logger.warning( + "Dropping vLLM feature positions outside hidden rows dropped=%s " + "hidden_rows=%s hidden_offset=%s feature_min=%s feature_max=%s", + int((~keep_mask).sum().item()), + int(hidden.size(0)), + hidden_position_offset, + int(feature_positions.min().item()), + int(feature_positions.max().item()), + ) + feature_positions = feature_positions[keep_mask] + relative_positions = relative_positions[keep_mask] + if int(feature_positions.numel()) <= 0: + raise ValueError( + "vLLM hidden_states contain no rows for replay feature positions: " + f"hidden_rows={int(hidden.size(0))}, " + f"hidden_position_offset={hidden_position_offset}" + ) + + selected = hidden.index_select(0, relative_positions).to(dtype=feature_config.dtype) + aux_hidden = selected[:, : len(target_layer_ids), :].flatten(1) + if include_final: + final_hidden = selected[:, required_layers - 1, :] + output_hidden = torch.cat([aux_hidden, final_hidden], dim=-1) + else: + output_hidden = aux_hidden + selected_input_ids = request.input_ids.index_select(0, feature_positions).long() + selected_loss_mask = request.loss_mask.index_select(0, feature_positions).float() + draft_position_ids = request.draft_position_ids.detach().cpu().long() + if filtered: + draft_position_ids = draft_position_ids[keep_mask] + + source_metadata = getattr(request, "source_metadata", None) + if source_metadata is None: + source_metadata = getattr(request, "metadata", {}) + metadata = dict(source_metadata or {}) + feature_start = int(feature_positions[0].item()) + feature_end = int(feature_positions[-1].item()) + 1 + metadata.update( + { + "source": feature_config.source, + "target_model_path": feature_config.target_model_id, + "target_revision": feature_config.target_model_revision, + "target_config_fingerprint": feature_config.target_config_fingerprint, + "tokenizer_fingerprint": feature_config.tokenizer_fingerprint, + "target_layer_ids": target_layer_ids, + "vllm_hidden_layers": int(hidden.size(1)), + "vllm_hidden_rows": int(hidden.size(0)), + "vllm_hidden_position_offset": hidden_position_offset, + "hidden_states_layout": hidden_layout, + "feature_start": feature_start, + "feature_end": feature_end, + "hidden_position_start": feature_start, + "hidden_position_end": feature_end, + "hidden_positions": feature_positions, + "sequence_length": int(selected_input_ids.numel()), + "full_sequence_length": int(request.input_ids.numel()), + "use_logits": feature_config.use_logits, + } + ) + return DraftFeatureSample( + algorithm=algorithm, + input_ids=selected_input_ids, + loss_mask=selected_loss_mask, + hidden_states=output_hidden.cpu().contiguous(), + position_ids=draft_position_ids, + metadata=metadata, + ) + + class TargetFeatureReplayer: """Materialize target hidden states only for standalone token replay.""" @@ -1133,9 +1282,7 @@ def _release_vllm_endpoint( else: state.failures += 1 state.consecutive_failures += 1 - state.cooldown_until = ( - time.monotonic() + self.vllm_endpoint_cooldown - ) + state.cooldown_until = time.monotonic() + self.vllm_endpoint_cooldown def _request_vllm_response(self, prompt_ids: list[int]) -> Any: last_error: Exception | None = None @@ -1309,106 +1456,32 @@ def _feature_from_vllm_payload( prompt_ids: list[int], source: str, ) -> DraftFeatureSample: - token_ids = payload.get("token_ids") - hidden = payload.get("hidden_states") - if not torch.is_tensor(token_ids) or not torch.is_tensor(hidden): - raise ValueError( - "vLLM hidden-states payload must contain token_ids and hidden_states" - ) - if token_ids.detach().cpu().long().tolist() != prompt_ids: - raise ValueError("vLLM hidden-states token_ids do not match replay input") - if hidden.dim() != 3: - raise ValueError( - "vLLM hidden_states must have shape [seq, layers, hidden], " - f"got {tuple(hidden.shape)}" - ) - feature_positions = sample.feature_positions.detach().cpu().long() - hidden_position_offset = max(len(prompt_ids) - int(hidden.size(0)), 0) - expected_layers = len(self.target_layer_ids) - include_final = self.hidden_layout in { - "eagle3_aux_plus_last", - "dflash_aux_plus_last", - } - required_layers = expected_layers + (1 if include_final else 0) - if int(hidden.size(1)) < required_layers: - raise ValueError( - "vLLM hidden_states layer count is too small: " - f"got {int(hidden.size(1))}, need at least {required_layers}. " - "Start vLLM with target layer ids plus the final layer when the " - "training layout needs last hidden states." - ) - relative_feature_positions = feature_positions - hidden_position_offset - feature_keep_mask = ( - (relative_feature_positions >= 0) - & (relative_feature_positions < int(hidden.size(0))) + expected_prompt_ids = ( + sample.input_ids[: int(sample.feature_positions[-1].item()) + 1] + .detach() + .cpu() + .long() + .tolist() ) - filtered_feature_positions = not bool(feature_keep_mask.all().item()) - if filtered_feature_positions: - dropped = int((~feature_keep_mask).sum().item()) - logger.warning( - "[target replay rank=%s] dropping vLLM feature positions outside " - "hidden rows dropped=%s hidden_rows=%s hidden_offset=%s " - "feature_min=%s feature_max=%s", - self.rank, - dropped, - int(hidden.size(0)), - hidden_position_offset, - int(feature_positions.min().item()), - int(feature_positions.max().item()), - ) - feature_positions = feature_positions[feature_keep_mask] - relative_feature_positions = relative_feature_positions[feature_keep_mask] - if int(feature_positions.numel()) <= 0: - raise ValueError( - "vLLM hidden_states contain no rows for replay feature positions: " - f"hidden_rows={int(hidden.size(0))}, " - f"hidden_position_offset={hidden_position_offset}" - ) - selected = hidden.index_select(0, relative_feature_positions).to( - dtype=self.dtype - ) - aux_hidden = selected[:, :expected_layers, :].flatten(1) - if include_final: - final_hidden = selected[:, required_layers - 1, :] - hidden_states = torch.cat([aux_hidden, final_hidden], dim=-1) - else: - hidden_states = aux_hidden - selected_input_ids = sample.input_ids.index_select(0, feature_positions).long() - selected_loss_mask = sample.loss_mask.index_select(0, feature_positions).float() - draft_position_ids = sample.draft_position_ids.detach().cpu().long() - if filtered_feature_positions: - draft_position_ids = draft_position_ids[feature_keep_mask] - metadata = dict(sample.metadata) - feature_start = int(feature_positions[0].item()) - feature_end = int(feature_positions[-1].item()) + 1 - metadata.update( - { - "source": source, - "target_model_path": self.model_path, - "target_revision": self.target_revision, - "target_config_fingerprint": self.target_config_fingerprint, - "target_layer_ids": list(self.target_layer_ids), - "vllm_hidden_layers": int(hidden.size(1)), - "vllm_hidden_rows": int(hidden.size(0)), - "vllm_hidden_position_offset": hidden_position_offset, - "hidden_states_layout": self.hidden_layout, - "feature_start": feature_start, - "feature_end": feature_end, - "hidden_position_start": feature_start, - "hidden_position_end": feature_end, - "hidden_positions": feature_positions, - "sequence_length": int(selected_input_ids.numel()), - "full_sequence_length": int(sample.input_ids.numel()), - "use_logits": self.use_logits, - } - ) - return DraftFeatureSample( - algorithm=self.algorithm, - input_ids=selected_input_ids, - loss_mask=selected_loss_mask, - hidden_states=hidden_states.cpu().contiguous(), - position_ids=draft_position_ids, - metadata=metadata, + if expected_prompt_ids != prompt_ids: + raise ValueError("prompt_ids do not match replay feature window") + return feature_from_vllm_payload( + payload, + sample, + FeatureContract( + algorithm=getattr(self, "algorithm", sample.algorithm), + target_layer_ids=list(self.target_layer_ids), + hidden_states_layout=self.hidden_layout, + dtype=self.dtype, + target_model_id=self.model_path, + target_model_revision=self.target_revision, + tokenizer_fingerprint=str( + getattr(self, "tokenizer_fingerprint", "replay-unspecified") + ), + use_logits=self.use_logits, + target_config_fingerprint=self.target_config_fingerprint, + source=source, + ), ) def _build_sparse_target_logprobs( @@ -1445,9 +1518,7 @@ def metrics(self) -> dict[str, float]: } if self.backend.startswith("vllm_"): metrics["replay/vllm_requests_total"] = float(self.vllm_requests) - metrics["replay/vllm_request_time_total"] = float( - self.vllm_request_seconds - ) + metrics["replay/vllm_request_time_total"] = float(self.vllm_request_seconds) with self._endpoint_lock: metrics["replay/vllm_endpoints_total"] = float( len(self._vllm_endpoint_states) @@ -1462,9 +1533,7 @@ def metrics(self) -> dict[str, float]: ) if self.backend == "vllm_mooncake": metrics["replay/mooncake_gets_total"] = float(self.mooncake_gets) - metrics["replay/mooncake_get_time_total"] = float( - self.mooncake_get_seconds - ) + metrics["replay/mooncake_get_time_total"] = float(self.mooncake_get_seconds) total = self.cache_hits + self.cache_misses if total > 0: metrics["replay/cache_hit_ratio"] = self.cache_hits / float(total) From ab15739c0c6a86442314feb0ef854f04d535ca47 Mon Sep 17 00:00:00 2001 From: xxyyrr598 <1615681145@qq.com> Date: Fri, 21 Aug 2026 14:50:40 +0800 Subject: [PATCH 35/50] feat(tq): remove hardcoded DSPARK algorithm restriction from TQ protocol and consumer - SampleMetadata/ExpectedFeatureConfig no longer require algorithm=DSPARK - TQFeatureStore keeps Producer's algorithm tag but does not compare or gate on it - docs updated to reflect non-DSPARK algorithms can reuse the public protocol --- docs/standalone_tq_consumer_implementation.md | 20 ++++++++++++------- tests/unit/test_drafter_sample_protocol.py | 13 ++++++++++++ tools/tq_connection_smoke.py | 2 +- verl_speco/trainer/tq_feature_store.py | 10 +--------- .../transport/drafter_sample_protocol.py | 14 +++++++------ 5 files changed, 36 insertions(+), 23 deletions(-) diff --git a/docs/standalone_tq_consumer_implementation.md b/docs/standalone_tq_consumer_implementation.md index 456afd40..83c6003c 100644 --- a/docs/standalone_tq_consumer_implementation.md +++ b/docs/standalone_tq_consumer_implementation.md @@ -563,12 +563,15 @@ clear 成功后才增加 `successful_steps`,然后复用原有 metrics 和 che `TQFeatureStore` 构造的 expected config 只固定: ```python -ExpectedFeatureConfig( - run_id=<当前训练 run_id>, - schema_version=<当前 schema>, - algorithm="DSPARK", -) -``` +ExpectedFeatureConfig( + run_id=<当前训练 run_id>, + schema_version=<当前 schema>, +) +``` + +TQ 会保留 Producer 写入的 `SampleMetadata.algorithm`,但不使用它选择训练 backend, +也不额外与启动配置比较。与原有离线 feature-store 训练一致,实际 trainer/backend 只由 +`rollout.drafter.speculative_algorithm` 和既有 backend factory 决定。 因此当前不会拿 Consumer 配置额外比较: @@ -679,7 +682,10 @@ OWNER_CLOSED 当前第一版有意不实现以下复杂能力: 1. Producer 本身尚未在本次 Consumer 改动中实现。 -2. 只支持 `algorithm=DSPARK`。 +2. TQ 公共协议和 Consumer 已不再写死 `DSPARK`;当前测试 Producer、启动脚本和已验证的 + feature 语义仍是 DSPARK。其他算法若能复用当前公共 dense fields,只需由对应 Producer + 生成正确的 `DraftFeatureSample`;若字段结构不同,则在协议模块增加对应 codec,不需要改 + TQ 的 key/tag 发现、rank 分配和 clear 流程。 3. 只支持 `drop_last=true`。 4. 不支持 TQ 与 `target_feature_pipeline.enabled=true` 同时开启。 5. 不提供严格的 crash exactly-once 或 checkpoint/queue 联合恢复。 diff --git a/tests/unit/test_drafter_sample_protocol.py b/tests/unit/test_drafter_sample_protocol.py index 74defffe..393b41ab 100644 --- a/tests/unit/test_drafter_sample_protocol.py +++ b/tests/unit/test_drafter_sample_protocol.py @@ -113,6 +113,19 @@ def test_encode_rejects_shape_mismatch() -> None: encode_sample(_sample(), meta) +def test_protocol_algorithm_is_not_hardcoded_to_dspark() -> None: + meta = replace(_metadata(), algorithm="EAGLE3") + sample = replace(_sample(), algorithm="EAGLE3") + expected = replace(_expected(), algorithm="EAGLE3") + + key = make_sample_key(meta) + fields = encode_sample(sample, meta) + restored = decode_sample(key, make_ready_tag(meta), fields, expected) + + assert make_ready_tag(meta)["algorithm"] == "EAGLE3" + assert restored.algorithm == "EAGLE3" + + def test_eos_record_is_control_only() -> None: key, fields, tag = make_eos_record("run-a", 18) assert key == "control:v1:run-a:eos" diff --git a/tools/tq_connection_smoke.py b/tools/tq_connection_smoke.py index ce198fa3..d1878681 100644 --- a/tools/tq_connection_smoke.py +++ b/tools/tq_connection_smoke.py @@ -52,7 +52,7 @@ def _config(args) -> dict: "ray": {"address": args.ray_address, "namespace": args.namespace}, "partition_id": "speco_drafter_features", "run_id": args.run_id, - "schema_version": 1, + "schema_version": 1, "controller": {"polling_mode": True}, "backend": { "storage_backend": "SimpleStorage", diff --git a/verl_speco/trainer/tq_feature_store.py b/verl_speco/trainer/tq_feature_store.py index 1921fe95..e59ccc4c 100644 --- a/verl_speco/trainer/tq_feature_store.py +++ b/verl_speco/trainer/tq_feature_store.py @@ -66,24 +66,18 @@ def __init__(self, config: Mapping[str, Any]): if not self.run_id: raise ValueError("transfer_queue.run_id is required for a TQ Consumer") self.schema_version = int(self.config.get("schema_version", 1)) - self.algorithm = str(self.config.get("algorithm", "DSPARK") or "DSPARK").upper() - if self.algorithm != "DSPARK": - raise ValueError( - f"Standalone TQ Consumer currently supports only DSPARK, got {self.algorithm!r}" - ) ray_cfg = _plain_dict(self.config.get("ray") or {}) self.ray_address = str(ray_cfg.get("address") or "").strip() self.ray_namespace = str(ray_cfg.get("namespace") or "").strip() or None if not self.ray_address: raise ValueError("transfer_queue.ray.address is required for a TQ Consumer") self._connected = False - # First version deliberately checks only run/protocol/algorithm. Tensor + # First version deliberately checks only run/protocol. Tensor # presence, lengths, shape and dtype self-consistency remain enforced by # decode_sample; model/tokenizer/layer identity checks stay disabled. self.expected_config = ExpectedFeatureConfig( run_id=self.run_id, schema_version=self.schema_version, - algorithm=self.algorithm, ) @classmethod @@ -113,8 +107,6 @@ def list_ready(self, run_id: str | None = None) -> list[ReadyEntry]: continue if int(tag.get("schema_version", -1)) != self.schema_version: continue - if str(tag.get("algorithm") or "").upper() != self.algorithm: - continue try: int(tag["sequence_no"]) str(tag["sample_id"]) diff --git a/verl_speco/transport/drafter_sample_protocol.py b/verl_speco/transport/drafter_sample_protocol.py index de54b00f..174c03cf 100644 --- a/verl_speco/transport/drafter_sample_protocol.py +++ b/verl_speco/transport/drafter_sample_protocol.py @@ -77,10 +77,8 @@ def validate(self) -> None: raise ValueError("SampleMetadata.sample_id must not be empty") if self.sequence_no < 0: raise ValueError("SampleMetadata.sequence_no must be non-negative") - if self.algorithm.strip().upper() != "DSPARK": - raise ValueError( - f"Standalone TQ protocol currently requires algorithm=DSPARK, got {self.algorithm!r}" - ) + if not self.algorithm.strip(): + raise ValueError("SampleMetadata.algorithm must not be empty") if len(self.hidden_shape) != 2: raise ValueError( f"SampleMetadata.hidden_shape must be [rows, hidden_dim], got {self.hidden_shape!r}" @@ -143,7 +141,7 @@ class ExpectedFeatureConfig: run_id: str schema_version: int = PROTOCOL_SCHEMA_VERSION - algorithm: str = "DSPARK" + algorithm: str | None = None target_model_id: str | None = None target_model_revision: str | None = None tokenizer_fingerprint: str | None = None @@ -288,7 +286,11 @@ def _validate_expected(meta: SampleMetadata, expected: ExpectedFeatureConfig) -> checks = { "run_id": expected.run_id, "schema_version": expected.schema_version, - "algorithm": expected.algorithm.strip().upper(), + "algorithm": ( + expected.algorithm.strip().upper() + if expected.algorithm is not None + else None + ), "target_model_id": expected.target_model_id, "target_model_revision": expected.target_model_revision, "tokenizer_fingerprint": expected.tokenizer_fingerprint, From 1351632efb7c8ca41ecb45c63d85837fa0bbdb0d Mon Sep 17 00:00:00 2001 From: vx120 <893600387@qq.com> Date: Fri, 21 Aug 2026 16:23:44 +0800 Subject: [PATCH 36/50] feat(tq): complete standalone producer and training workflow Signed-off-by: vx120 <893600387@qq.com> --- README.md | 16 +- ...sync_vllm_mooncake_dspark_training_plan.md | 18 +- docs/standalone_tq_consumer_implementation.md | 12 +- ...standalone_tq_foundation_implementation.md | 26 +- docs/standalone_tq_producer.md | 41 +- ...standalone_vllm_tq_dspark_training_plan.md | 24 +- docs/transferqueue_integration_plan.md | 2 + examples/run_dspark_tq_producer.sh | 12 +- ...en3-8b_drafter_dspark_separate_training.sh | 21 + .../run_qwen3-8b_drafter_separate_training.sh | 143 +--- tests/examples/test_example_scripts.py | 37 +- tests/unit/test_producer_input_reader.py | 220 ++++++ .../test_standalone_tq_training_launcher.py | 290 ++++++++ tests/unit/test_tq_producer.py | 332 +++++++++ tests/unit/test_transferqueue_bridge.py | 9 +- tests/unit/test_vllm_feature_client.py | 60 ++ tools/run_dspark_tq_consumer.sh | 25 +- tools/run_dspark_tq_e2e_test.sh | 170 +++++ tools/run_dspark_tq_owner.sh | 6 +- verl_speco/config/speco_base.yaml | 5 +- .../integration/transferqueue_bridge.py | 67 +- verl_speco/producer/input_reader.py | 369 ++++++++-- verl_speco/producer/vllm_feature_client.py | 82 ++- verl_speco/standalone_tq_producer.py | 31 +- verl_speco/standalone_tq_training_launcher.py | 638 ++++++++++++++++++ verl_speco/tq_owner.py | 13 +- verl_speco/trainer/standalone_checkpoint.py | 14 + 27 files changed, 2401 insertions(+), 282 deletions(-) create mode 100644 examples/run_qwen3-8b_drafter_dspark_separate_training.sh create mode 100644 tests/unit/test_producer_input_reader.py create mode 100644 tests/unit/test_standalone_tq_training_launcher.py create mode 100644 tests/unit/test_tq_producer.py create mode 100644 tests/unit/test_vllm_feature_client.py create mode 100644 tools/run_dspark_tq_e2e_test.sh create mode 100644 verl_speco/standalone_tq_training_launcher.py diff --git a/README.md b/README.md index 8f75310c..767a7a81 100644 --- a/README.md +++ b/README.md @@ -258,9 +258,11 @@ actor_rollout_ref.rollout.drafter.speculative_algorithm=EAGLE3 ## Separate Draft Model Training -verl-SpeCo also supports a separate draft model training workflow. In this -mode, rollout workers collect drafter training features into a feature store, -and the draft model can be trained separately after feature collection. +verl-SpeCo also supports standalone DSpark draft-model training from a finite +verl-style prompt Parquet or prompt/response JSONL/Parquet file. For prompt-only +rows, a producer asks the target vLLM service to generate the response while +extracting prompt/output hidden states. It transfers each global batch through +TransferQueue, and a consumer trains the drafter independently of PPO. Quickstart: @@ -268,9 +270,11 @@ Quickstart: bash examples/run_qwen3-8b_drafter_separate_training.sh ``` -Replace the model, drafter, dataset, feature-store, and checkpoint paths in -the script before running it. The script uses `collect_only` mode for rollout -feature collection and `offline` mode for standalone drafter training. +Set the same model, dataset, drafter, checkpoint, GPU, and optimization values +used by ordinary standalone training near the top of the script. Transport +identity, Ray/TQ connection settings, and the Producer/Consumer lifecycle are +derived and managed internally. The target hidden-state vLLM service uses the +local port 8000 convention. The main mode values are: diff --git a/docs/async_vllm_mooncake_dspark_training_plan.md b/docs/async_vllm_mooncake_dspark_training_plan.md index 6cda0e40..525da0fb 100644 --- a/docs/async_vllm_mooncake_dspark_training_plan.md +++ b/docs/async_vllm_mooncake_dspark_training_plan.md @@ -1,22 +1,28 @@ # verl-SpeCo 异步 vLLM → TransferQueue(Mooncake 后端)→ DSpark 流式训练方案 +Last updated: 08/21/2026 + ## 1. 文档范围 本文只讨论当前 `verl-SpeCo-ls` 项目的独立草稿模型训练入口: ```text examples/run_qwen3-8b_drafter_separate_training.sh - → python -m verl_speco.draft_train_launcher - → torch.distributed.run - → python -m verl_speco.draft_train - → run_standalone_draft_training() + → python -m verl_speco.standalone_tq_training_launcher + ├─→ TransferQueue owner + ├─→ vLLM hidden-state producer + └─→ TransferQueue consumer + → python -m verl_speco.draft_train_launcher + → torch.distributed.run + → python -m verl_speco.draft_train + → run_standalone_draft_training() ``` 目标是训练 `mode=train/offline` 的草稿模型,不是 RL 与草稿模型一起训练,也不是先生成全量 hidden states 再长期保存到磁盘。 -输入文件已经包含提前生成好的 response。新流水线需要: +输入既可以是已有 response 的 replay 文件,也可以是 verl prompt-only Parquet。新流水线需要: -1. Producer 读取 prompt 和预生成 response,构造完整 token 序列; +1. Producer 读取 prompt;缺少 response 时由 target vLLM 生成; 2. Producer 并行请求一个或多个 vLLM endpoint 做 prefill; 3. hidden states 写入 TransferQueue(简称 TQ),TQ 的数据后端使用 Mooncake; 4. DSpark 的各训练 rank 从 TQ/Mooncake 并行读取自己的 local batch; diff --git a/docs/standalone_tq_consumer_implementation.md b/docs/standalone_tq_consumer_implementation.md index 83c6003c..ef01c311 100644 --- a/docs/standalone_tq_consumer_implementation.md +++ b/docs/standalone_tq_consumer_implementation.md @@ -1,6 +1,8 @@ -# 独立 DSpark 训练 TQ Consumer 实现说明 - -## 1. 文档范围和当前结论 +# 独立 DSpark 训练 TQ Consumer 实现说明 + +Last updated: 08/21/2026 + +## 1. 文档范围和当前结论 本文只说明当前仓库中已经实现的独立训练 Consumer。这里的 Consumer 是由 `torchrun` 启动的 DSpark 草稿模型训练任务:它持续从 TransferQueue(下文简称 TQ)发现样本,各训练 rank 分别取得自己负责的 Tensor,复用原有 DSpark 训练逻辑完成一次 optimizer step,然后由 rank 0 删除这一整个 global batch 对应的 TQ 记录。 @@ -146,7 +148,7 @@ connect_ray_cluster(ray_address, ray_namespace) connect_transfer_queue_client() ``` -`connect_ray_cluster()` 内部调用 `ray.init(address=..., namespace=...)`。`connect_transfer_queue_client()` 再调用无参数 `tq.init()`;TQ 由此在当前 Ray namespace 查找 Owner 创建的 named Controller。之后所有 KV 操作都显式携带相同的 `partition_id`。 +`connect_ray_cluster()` 内部调用 `ray.init(address=..., namespace=...)`。`connect_transfer_queue_client()` 再使用与 Owner 相同的 native 配置调用 `tq.init(config)`;TQ 会优先在当前 Ray namespace 查找 Owner 创建的 named Controller,找到时忽略本次配置并只创建本地 Client。即使 Client 意外先于 Owner 初始化,也会使用同一份 backend/controller 配置,而不会按默认配置创建服务。之后所有 KV 操作都显式携带相同的 `partition_id`。 因此,“连接同一个 TQ”实际由三层身份共同决定:同一 Ray 集群、同一 namespace 下的同一 named Controller、同一 `partition_id`。 @@ -366,7 +368,7 @@ TQ store 被限定为 `read_only=True`,意思是它是训练 Consumer source ```text configure_transfer_queue → ray.init(address, namespace) -→ tq.init() 连接 named Controller +→ tq.init(same native config) 连接 named Controller → 本 rank 设置 _connected=True ``` diff --git a/docs/standalone_tq_foundation_implementation.md b/docs/standalone_tq_foundation_implementation.md index e8f95581..fd01e7ff 100644 --- a/docs/standalone_tq_foundation_implementation.md +++ b/docs/standalone_tq_foundation_implementation.md @@ -1,6 +1,8 @@ -# Standalone TQ 公共基础层实现说明 - -## 1. 文档范围和已验证结论 +# Standalone TQ 公共基础层实现说明 + +Last updated: 08/21/2026 + +## 1. 文档范围和已验证结论 本文解释当前仓库已经实现并测试通过的 TQ 公共基础层: @@ -86,10 +88,10 @@ Controller 保存控制信息和数据位置;使用 MooncakeStore 时,大 te Owner、Producer、每个 Consumer rank 都在各自 OS 进程内拥有独立 TQ Client。 -普通 Client 通过无参: +普通 Client 也传入相同的 native 配置: ```python -tq.init() +tq.init(native_config) ``` 发现 named Controller并初始化本进程的 storage manager。Client 不是 Ray actor,Producer 和 torchrun rank 也不需要改成 Ray actor。 @@ -194,7 +196,7 @@ drop_last → _native_tq_config() → controller/backend等TQ字段 → OmegaConf DictConfig -→ tq.init() +→ tq.init(same native config) ``` ## 4. Bridge 的进程内状态 @@ -364,10 +366,10 @@ connect_ray_cluster(ray_address, namespace) connect_transfer_queue_client() ``` -`connect_transfer_queue_client()` 最终调用无参: +`connect_transfer_queue_client()` 最终调用: ```python -tq.init() +tq.init(same_native_config) ``` TQ 0.1.7 内部通过: @@ -573,7 +575,7 @@ bridge 执行: 1. 检查 TQ 已启用; 2. 丢弃 fields 中非 tensor 值; -3. 确保本进程已经 `tq.init()`; +3. 确保本进程已经使用相同 native 配置执行 `tq.init(config)`; 4. 取得配置中的 partition; 5. 调用: @@ -886,7 +888,7 @@ bridge 测试 `tests/unit/test_transferqueue_bridge.py` 覆盖: 1. Ray address/namespace 参数; 2. Owner 只向 TQ 传原生配置; -3. Client 无参 `tq.init()`; +3. Client 使用相同 native 配置调用 `tq.init(config)`; 4. put/list/get-many/clear; 5. batch 返回顺序; 6. Client close 不调用全局 close; @@ -926,7 +928,7 @@ Client 路径: ```text 连接同一个Ray -→ tq.init() +→ tq.init(same native config) → kv_list发现两个key → 一次kv_batch_get([k0,k1]) → 拆成两个fields dict @@ -966,7 +968,7 @@ Owner 普通Client → ray.init(same address, same namespace) -→ tq.init() +→ tq.init(same native config) → 找到同一个Controller DraftFeatureSample + SampleMetadata diff --git a/docs/standalone_tq_producer.md b/docs/standalone_tq_producer.md index e9c95be0..e651a04f 100644 --- a/docs/standalone_tq_producer.md +++ b/docs/standalone_tq_producer.md @@ -1,10 +1,11 @@ # Standalone vLLM → TransferQueue Producer -Last updated: 08/20/2026 +Last updated: 08/21/2026 -本文解释这次新增的 standalone Producer:它读取已经有 `prompt` 和 -`response` 的 JSONL,向 vLLM 请求 target hidden states,并把每条样本写到 -已存在的 TransferQueue(TQ)。它不启动 TQ owner,也不启动 Consumer/训练。 +本文解释 standalone Producer:它可以直接读取 verl 的 prompt-only Parquet(包括 +DAPO-Math-17k 的 chat-message `prompt`),也兼容已有 `prompt`/`response` 的 JSONL +或 Parquet。缺少 response 时由 target vLLM 生成,并在同一请求中提取 prompt 与 +output hidden states,之后把样本写到已存在的 TransferQueue(TQ)。 这条路径面向第一版 DSpark standalone 训练:Producer、TQ owner 和 Consumer 是三个独立 OS 进程;Ray 只用于让它们找到同一个 TQ Controller,hidden states @@ -25,7 +26,7 @@ Producer 补上这一段,不引入第二套协议或 feature store。 ## 数据流 ```text -prompt/response JSONL +verl prompt Parquet 或 prompt/response JSONL/Parquet │ │ 按文件顺序分配 sequence_no 和 sample_id ▼ @@ -49,8 +50,10 @@ Ray head、TQ Controller 或 storage backend;这些由 `verl-speco-tq-owner` ## 输入文件 -输入只能是 JSONL:每个非空行是一个 JSON object,必须有字符串 `prompt` 与非空 -字符串 `response`。 +输入可以是 JSONL 或 Parquet。`prompt` 可以是字符串,也可以是 verl 常用的 +`[{"role": ..., "content": ...}]` chat-message 列表。`response` 是可选字符串: +存在时直接 replay;不存在时由 target vLLM 生成。Parquet 通过 +`data.train_files` 直接传入,不需要转换。 ```json {"sample_id":"train-000017","prompt":"Question: 1 + 1 = ","response":"2"} @@ -59,6 +62,9 @@ Ray head、TQ Controller 或 storage backend;这些由 `verl-speco-tq-owner` - `sequence_no` 按非空行的文件顺序从 0 分配;并发完成顺序不会影响它。 - `sample_id` 可选;省略时生成 `train-000000`、`train-000001` 等稳定值。 +- verl 数据的 `extra_info.index` 存在时会优先作为稳定 `sample_id`。 +- chat-message prompt 通过 target tokenizer 的 `apply_chat_template()` 编码,并加上 + generation prompt;不能把 `reward_model.ground_truth` 当作模型 response。 - Producer tokenize `prompt` 和 `prompt + response`。后者必须以 prompt 的 token IDs 为前缀;否则会报错,而不会猜测 response 的 loss-mask 边界。 - `loss_mask` 中 prompt token 为 0,response token 为 1。 @@ -69,12 +75,16 @@ Ray head、TQ Controller 或 storage backend;这些由 `verl-speco-tq-owner` ## vLLM 与 hidden states -Producer 使用 OpenAI-compatible completions API: +对已有 response,Producer 使用 OpenAI-compatible completions API 做 prefill。 +对 prompt-only 数据,Producer 在一次请求中生成 response 并要求保存输出 hidden: ```text -prompt= -max_tokens=1 -extra_body={"return_token_ids": true} +prompt= +max_tokens=<内部有界长度> +extra_body={ + "return_token_ids": true, + "kv_transfer_params": {"include_output_tokens": true} +} ``` 响应必须同时满足: @@ -151,7 +161,7 @@ vLLM 请求会暂停,等待 Consumer 清理已成功训练的 key。 | 字段 | 含义 | | --- | --- | -| `input_path` | 上述 JSONL 文件 | +| `input_path` | 上述 JSONL 或 Parquet 文件 | | `tokenizer_path` / `tokenizer_fingerprint` | 用于 tokenization 和 Consumer 合同校验 | | `target_model_id` / `target_model_revision` | target checkpoint 身份 | | `target_layer_ids` | auxiliary target layer IDs;DSpark L1 时 wire metadata 会额外写 `-1` 表示 final layer | @@ -189,10 +199,15 @@ verl-speco-tq-producer \ 完整生命周期顺序仍是:Ray/TQ backend → TQ owner → Consumer → Producer → Consumer drain → owner shutdown。Producer 完成不代表训练完成,EOS 只表示不会再有新样本。 +正式独立训练入口 +`examples/run_qwen3-8b_drafter_separate_training.sh` 会通过 +`verl_speco.standalone_tq_training_launcher` 自动管理这套生命周期;上面的 Producer +脚本仅用于单独调试 Producer。 ## 测试覆盖与未验证项 -新增测试覆盖:JSONL 解析与 token 边界、多个 endpoint 的并发限制、ready 队列背压、 +新增测试覆盖:JSONL/真实 Parquet 解析、DAPO chat prompt、target response generation、 +token 边界、多个 endpoint 的并发限制、ready 队列背压、 成功时 sample 后 EOS 与临时文件删除、失败时无 EOS 且保留临时文件,以及旧 EAGLE3 转换路径仍可复用公共函数。 diff --git a/docs/standalone_vllm_tq_dspark_training_plan.md b/docs/standalone_vllm_tq_dspark_training_plan.md index 846c9053..e22c8660 100644 --- a/docs/standalone_vllm_tq_dspark_training_plan.md +++ b/docs/standalone_vllm_tq_dspark_training_plan.md @@ -1,11 +1,13 @@ # 独立 vLLM Producer + TQ + DSpark Consumer 第一版方案 +Last updated: 08/21/2026 + ## 1. 第一版要实现什么 只实现下面这条主链路: ```text -包含 prompt + 预生成 response 的输入文件 +verl prompt-only 数据或包含 prompt + response 的输入文件 → Producer 并发请求 vLLM prefill → Producer 将每条训练样本写入 TQ → Consumer 从同一个 TQ 取样本 @@ -25,7 +27,7 @@ | Producer | 1 个进程 | 读文件、并发请求 vLLM、写 TQ | | Consumer | 1 个 torchrun 任务 | 多个 rank 从 TQ 取数并训练 DSpark | -Producer 和 Consumer 不互相调用,也不通过 HTTP 或 Ray 传样本。二者先连接同一个 Ray 集群,再通过无参 `tq.init()` 找到同一个 TQ,最后使用 TQ KV API 读写样本。 +Producer 和 Consumer 不互相调用,也不通过 HTTP 或 Ray 传样本。二者先连接同一个 Ray 集群,再通过携带相同 native 配置的 `tq.init(config)` 找到同一个 TQ,最后使用 TQ KV API 读写样本。已有 Controller 时 TransferQueue 0.1.7 会忽略后续配置并只连接;若 Client 意外先初始化,同一配置可避免默认 backend 抢先生效。 ## 2. 共同的数据约定 @@ -295,7 +297,7 @@ EOS 不进入训练 batch。Consumer 看到 EOS 后继续处理剩余 ready samp ### 3.1 已验证的 TQ 0.1.7 连接机制 -`TransferQueue==0.1.7` 没有提供“把 Controller 地址直接传给第二个进程”的高层连接接口。它的无参 `tq.init()` 内部执行: +`TransferQueue==0.1.7` 没有提供“把 Controller 地址直接传给第二个进程”的高层连接接口。`tq.init(config)` 会先尝试以下已有 Controller 连接逻辑;存在时忽略传入配置,不存在时才用配置创建服务: ```python _TQ_CONTROLLER = ray.get_actor("TransferQueueController") @@ -309,8 +311,8 @@ _maybe_create_tq_client(conf) ```text TQ owner:ray.init(address) → tq.init(full_tq_config) → 创建 named Controller -Producer:ray.init(address) → tq.init() → ray.get_actor() → 创建本地 TQ client -Consumer rank 0..N:ray.init(address) → tq.init() → ray.get_actor() → 创建各自 TQ client +Producer:ray.init(address) → tq.init(same config) → ray.get_actor() → 创建本地 TQ client +Consumer rank 0..N:ray.init(address) → tq.init(same config) → ray.get_actor() → 创建各自 TQ client ``` ### 3.2 直接移植并扩展 PR #48 的 bridge @@ -370,7 +372,7 @@ def close_transfer_queue_owner() -> None: ... - 由 Producer 和每个 Consumer rank 调用; - 前置条件是当前进程已经连接 Ray; -- 调用无参 `tq.init()`,通过 `ray.get_actor("TransferQueueController")` 发现 owner; +- 调用 `tq.init(same native config)`,通过 `ray.get_actor("TransferQueueController")` 发现 owner;已有 Controller 时配置会被忽略,意外抢先时则以相同配置创建; - 只创建当前进程的 TQ client,不创建新的 Controller; - 成功后设置 `_state.initialized=True`;重复调用直接返回。 @@ -466,8 +468,8 @@ Owner 必须常驻。TQ 0.1.7 创建 Controller 时没有设置 `lifetime="detac 3. 启动 TQ owner;owner 连接 Ray并调用 tq.init(full config) 4. 等待 owner_ready 5. 启动一个或多个 vLLM servers -6. 启动 Consumer;每个 torchrun rank 连接 Ray,然后 tq.init() -7. 启动 Producer;连接 Ray,然后 tq.init() +6. 启动 Consumer;每个 torchrun rank 连接 Ray,然后 tq.init(same native config) +7. 启动 Producer;连接 Ray,然后 tq.init(same native config) 8. Producer 写 EOS,关闭本地 client并退出 9. Consumer drain、保存 final checkpoint,所有 ranks 关闭本地 client并退出 10. 给 TQ owner 发送 SIGTERM;仅 owner 执行 tq.close() @@ -488,7 +490,7 @@ Owner 必须常驻。TQ 0.1.7 创建 Controller 时没有设置 `lifetime="detac → 初始化多个 vLLM endpoint clients → 流式读取输入文件 → 为每条输入分配 sequence_no/sample_id -→ 拼接 prompt+预生成 response,得到 input_ids/loss_mask +→ 缺少 response 时由 target vLLM 生成;构造 input_ids/loss_mask → 并发请求 vLLM prefill → 读取 vLLM hidden-state 临时结果 → 转换成 DSpark DraftFeatureSample @@ -595,7 +597,7 @@ class InputRecord: sequence_no: int sample_id: str prompt: str - response: str + response: str | None source_metadata: dict[str, Any] def iter_input_records(path: str) -> Iterator[InputRecord]: ... @@ -603,7 +605,7 @@ def tokenize_record(record: InputRecord, tokenizer, config) -> TokenizedRequest: def build_loss_mask(input_ids: Tensor, prompt_length: int) -> Tensor: ... ``` -`iter_input_records()` 流式读取,不把全文件载入内存;在这里按文件顺序分配稳定的 `sequence_no`。`tokenize_record()` 拼接已经存在的 prompt/response,不调用模型生成 response;输出至少包含 `input_ids:int64[L]`、`position_ids:int64[L]`、`loss_mask:float32[L]` 和请求 vLLM 所需字段。 +`iter_input_records()` 流式读取 JSONL/Parquet,不把全文件载入内存,并按文件顺序分配稳定的 `sequence_no`。已有 response 时 `tokenize_record()` 直接拼接;prompt-only verl 数据通过 chat template 编码后由 target vLLM 生成 response,并设置 `include_output_tokens=true` 同步提取输出 hidden states。 #### `verl_speco/producer/vllm_feature_client.py` diff --git a/docs/transferqueue_integration_plan.md b/docs/transferqueue_integration_plan.md index 56b7c3f8..cac16ac1 100644 --- a/docs/transferqueue_integration_plan.md +++ b/docs/transferqueue_integration_plan.md @@ -1,5 +1,7 @@ # verl-SpeCo TransferQueue 落地方案 +Last updated: 08/21/2026 + > 目标:在**不修改上游 verl**的前提下,把 SpeCo online 训练里的逐样本特征流 > 从「`SpecoRayPPOTrainer` driver 中转 + Ray object store」改为「TransferQueue > 直传」,干掉 driver 这个数据瓶颈,并解锁流式消费与跨副本负载均衡。 diff --git a/examples/run_dspark_tq_producer.sh b/examples/run_dspark_tq_producer.sh index ae8f62b5..63c19384 100644 --- a/examples/run_dspark_tq_producer.sh +++ b/examples/run_dspark_tq_producer.sh @@ -23,12 +23,17 @@ set -euo pipefail : "${TOKENIZER_FINGERPRINT:?Set TOKENIZER_FINGERPRINT to a verified fingerprint}" : "${TARGET_LAYER_IDS:?Set TARGET_LAYER_IDS as a Hydra list, for example '[2,8,14,20,26]'}" : "${VLLM_ENDPOINTS:?Set VLLM_ENDPOINTS as a Hydra list, for example '[http://node0:8000/v1]'}" +PYTHON_BIN=${PYTHON_BIN:-python3} +TQ_NAMESPACE=${TQ_NAMESPACE:-speco-drafter} +TQ_PARTITION_ID=${TQ_PARTITION_ID:-speco_drafter_features} +VLLM_MODEL=${VLLM_MODEL:-${TARGET_MODEL_PATH}} -python -m verl_speco.standalone_tq_producer \ +exec "${PYTHON_BIN}" -m verl_speco.standalone_tq_producer \ actor_rollout_ref.rollout.drafter.speculative_algorithm=DSPARK \ actor_rollout_ref.rollout.drafter.training.transfer_queue.enable=true \ actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.address="${RAY_ADDRESS}" \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.namespace=speco-drafter \ + actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.namespace="${TQ_NAMESPACE}" \ + actor_rollout_ref.rollout.drafter.training.transfer_queue.partition_id="${TQ_PARTITION_ID}" \ actor_rollout_ref.rollout.drafter.training.transfer_queue.run_id="${SPECO_TQ_RUN_ID}" \ speco.standalone_tq_producer.input_path="${PRODUCER_INPUT_PATH}" \ speco.standalone_tq_producer.target_model_id="${TARGET_MODEL_PATH}" \ @@ -37,4 +42,5 @@ python -m verl_speco.standalone_tq_producer \ speco.standalone_tq_producer.tokenizer_fingerprint="${TOKENIZER_FINGERPRINT}" \ speco.standalone_tq_producer.target_layer_ids="${TARGET_LAYER_IDS}" \ speco.standalone_tq_producer.vllm_endpoints="${VLLM_ENDPOINTS}" \ - speco.standalone_tq_producer.vllm_model="${TARGET_MODEL_PATH}" + speco.standalone_tq_producer.vllm_model="${VLLM_MODEL}" \ + "$@" diff --git a/examples/run_qwen3-8b_drafter_dspark_separate_training.sh b/examples/run_qwen3-8b_drafter_dspark_separate_training.sh new file mode 100644 index 00000000..2cd1e8eb --- /dev/null +++ b/examples/run_qwen3-8b_drafter_dspark_separate_training.sh @@ -0,0 +1,21 @@ +#!/usr/bin/env bash +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +set -euo pipefail + +# Backward-compatible DSpark entry point. Both example names run the same +# standalone vLLM -> TransferQueue/Mooncake -> DSpark training pipeline. +script_dir=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd) + +exec bash "${script_dir}/run_qwen3-8b_drafter_separate_training.sh" "$@" diff --git a/examples/run_qwen3-8b_drafter_separate_training.sh b/examples/run_qwen3-8b_drafter_separate_training.sh index 925841aa..e03a9cd2 100644 --- a/examples/run_qwen3-8b_drafter_separate_training.sh +++ b/examples/run_qwen3-8b_drafter_separate_training.sh @@ -1,120 +1,42 @@ +#!/usr/bin/env bash +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. set -euo pipefail set -x -# GPU smoke example for separate EAGLE3 draft model training. -# -# Stage 1 collects draft-training features from a short PPO/vLLM run without -# training the drafter inside PPO. Stage 2 launches independent multi-GPU draft -# training with python -m verl_speco.draft_train_launcher, which internally -# starts torch.distributed.run. -# -# Usage: -# bash examples/run_qwen3-8b_drafter_separate_training.sh -# RUN_STAGE=collect bash examples/run_qwen3-8b_drafter_separate_training.sh -# RUN_STAGE=train bash examples/run_qwen3-8b_drafter_separate_training.sh +# One-command standalone DSpark draft-model training. The launcher internally +# starts the hidden-state target vLLM and uses the +# Producer -> TransferQueue -> Consumer path. -project_name='verl_grpo_example_eagle3_drafter' -exp_name='qwen3_8b_eagle3_separate_drafter_vllm_gpu' +project_name=verl_dspark_drafter +exp_name=qwen3_8b_dspark_separate_training -gen_tp=2 -train_sp=1 -ppo_gpus_per_node=8 draft_train_gpus_per_node=8 -ray_num_cpus=${SPECO_RAY_NUM_CPUS:-64} -ray_worker_soft_limit=${SPECO_RAY_WORKER_SOFT_LIMIT:-8} -MODEL_PATH=/path/to/model -CKPTS_DIR=/path/to/checkpoint -TRAIN_FILE=/path/to/train_file -TEST_FILE=/path/to/test_file -DRAFTER_PATH=/path/to/vllm-compatible-eagle3-drafter -FEATURE_STORE_DIR=/path/to/speco/eagle3_features -DRAFT_CKPTS_DIR=/path/to/speco/eagle3_draft_ckpts +MODEL_PATH=/path/to/Qwen3-8B +# Ordinary verl prompt Parquet is supported; target vLLM generates responses. +TRAIN_FILE=/path/to/train_file.parquet +DRAFTER_PATH=/path/to/vllm-compatible-dspark-drafter +DRAFT_CKPTS_DIR=/path/to/dspark_draft_checkpoints -RUN_STAGE=${RUN_STAGE:-both} - -if [ "${RUN_STAGE}" = "both" ] || [ "${RUN_STAGE}" = "collect" ]; then -PYTHONUNBUFFERED=1 python3 -m verl_speco.main \ - algorithm.adv_estimator=grpo \ - ray_kwargs.ray_init.num_cpus=${ray_num_cpus} \ - +ray_kwargs.ray_init._system_config.prestart_worker_first_driver=false \ - +ray_kwargs.ray_init._system_config.num_workers_soft_limit=${ray_worker_soft_limit} \ - data.train_files=${TRAIN_FILE} \ - data.val_files=${TEST_FILE} \ - data.train_batch_size=64 \ - data.max_prompt_length=512 \ - data.max_response_length=8192 \ - data.filter_overlong_prompts=True \ - data.filter_overlong_prompts_workers=256 \ - data.truncation='error' \ - actor_rollout_ref.model.path=${MODEL_PATH} \ - actor_rollout_ref.actor.optim.lr=1e-6 \ - actor_rollout_ref.model.use_remove_padding=True \ - actor_rollout_ref.actor.ppo_mini_batch_size=16 \ - actor_rollout_ref.actor.ppo_micro_batch_size_per_gpu=16 \ - actor_rollout_ref.actor.use_kl_loss=True \ - actor_rollout_ref.actor.kl_loss_coef=0.001 \ - actor_rollout_ref.actor.kl_loss_type=low_var_kl \ - actor_rollout_ref.actor.entropy_coeff=0 \ - actor_rollout_ref.actor.calculate_entropy=False \ - actor_rollout_ref.model.enable_gradient_checkpointing=True \ - actor_rollout_ref.actor.fsdp_config.param_offload=True \ - actor_rollout_ref.actor.fsdp_config.optimizer_offload=True \ - actor_rollout_ref.rollout.log_prob_micro_batch_size_per_gpu=16 \ - actor_rollout_ref.rollout.tensor_model_parallel_size=${gen_tp} \ - actor_rollout_ref.actor.ulysses_sequence_parallel_size=${train_sp} \ - actor_rollout_ref.ref.ulysses_sequence_parallel_size=${train_sp} \ - actor_rollout_ref.ref.log_prob_use_dynamic_bsz=True \ - actor_rollout_ref.actor.use_dynamic_bsz=True \ - actor_rollout_ref.rollout.log_prob_use_dynamic_bsz=True \ - actor_rollout_ref.rollout.name=vllm \ - actor_rollout_ref.rollout.gpu_memory_utilization=0.4 \ - actor_rollout_ref.rollout.n=5 \ - actor_rollout_ref.ref.log_prob_micro_batch_size_per_gpu=16 \ - actor_rollout_ref.ref.fsdp_config.param_offload=True \ - actor_rollout_ref.rollout.drafter.enable=True \ - actor_rollout_ref.rollout.drafter.enable_drafter_training=True \ - actor_rollout_ref.rollout.drafter.model_path=${DRAFTER_PATH} \ - actor_rollout_ref.rollout.drafter.speculative_algorithm="EAGLE3" \ - actor_rollout_ref.rollout.drafter.training.mode=collect_only \ - actor_rollout_ref.rollout.drafter.training.feature_store.type=torch_shard \ - actor_rollout_ref.rollout.drafter.training.feature_store.path=${FEATURE_STORE_DIR} \ - actor_rollout_ref.rollout.drafter.training.feature_store.max_samples_per_shard=256 \ - actor_rollout_ref.rollout.drafter.training.feature_store.flush_interval_steps=1 \ - actor_rollout_ref.rollout.drafter.training.collect_hidden_states_from_sgl=True \ - actor_rollout_ref.rollout.drafter.training.collect_hidden_states_from_old_logprob=True \ - actor_rollout_ref.rollout.drafter.training.old_logprob_hidden_capture_impl=forward_hook \ - actor_rollout_ref.rollout.drafter.training.use_logits=False \ - actor_rollout_ref.rollout.drafter.rollout.spec_steps=3 \ - actor_rollout_ref.rollout.drafter.rollout.spec_topk=1 \ - actor_rollout_ref.rollout.drafter.rollout.spec_verify_tokens=4 \ - actor_rollout_ref.rollout.drafter.training.step=20 \ - actor_rollout_ref.rollout.drafter.training.collect_interval_steps=1 \ - actor_rollout_ref.rollout.drafter.training.training_interval_steps=1 \ - actor_rollout_ref.rollout.drafter.training.publish_interval_steps=0 \ - actor_rollout_ref.rollout.drafter.training.publish_async=False \ - actor_rollout_ref.rollout.drafter.training.publish_dtype=bf16 \ - actor_rollout_ref.rollout.load_format="auto" \ - actor_rollout_ref.actor.strategy=fsdp2 \ - algorithm.use_kl_in_reward=False \ - trainer.val_before_train=False \ - trainer.critic_warmup=0 \ - trainer.logger='["console"]' \ - trainer.project_name=${project_name} \ - trainer.experiment_name=${exp_name}_collect \ - trainer.n_gpus_per_node=${ppo_gpus_per_node} \ - trainer.nnodes=1 \ - trainer.default_local_dir=${CKPTS_DIR} \ - trainer.save_freq=20 \ - trainer.test_freq=5 \ - trainer.total_epochs=1 $@ -fi +PYTHON_BIN=${PYTHON_BIN:-python3} -if [ "${RUN_STAGE}" = "both" ] || [ "${RUN_STAGE}" = "train" ]; then -PYTHONUNBUFFERED=1 python3 -m verl_speco.draft_train_launcher \ +PYTHONUNBUFFERED=1 "${PYTHON_BIN}" -m verl_speco.standalone_tq_training_launcher \ speco.draft_training.num_gpus_per_node=${draft_train_gpus_per_node} \ speco.draft_training.nnodes=1 \ speco.draft_training.standalone=True \ + data.train_files=${TRAIN_FILE} \ actor_rollout_ref.model.path=${MODEL_PATH} \ actor_rollout_ref.actor.strategy=fsdp2 \ actor_rollout_ref.actor.fsdp_config.param_offload=True \ @@ -124,7 +46,7 @@ PYTHONUNBUFFERED=1 python3 -m verl_speco.draft_train_launcher \ actor_rollout_ref.rollout.drafter.enable_drafter_training=True \ actor_rollout_ref.rollout.drafter.model_path=${DRAFTER_PATH} \ actor_rollout_ref.rollout.drafter.checkpoint_path=${DRAFT_CKPTS_DIR} \ - actor_rollout_ref.rollout.drafter.speculative_algorithm="EAGLE3" \ + actor_rollout_ref.rollout.drafter.speculative_algorithm=DSPARK \ actor_rollout_ref.rollout.drafter.training.mode=offline \ actor_rollout_ref.rollout.drafter.training.max_steps=10 \ actor_rollout_ref.rollout.drafter.training.save_interval_steps=5 \ @@ -133,9 +55,6 @@ PYTHONUNBUFFERED=1 python3 -m verl_speco.draft_train_launcher \ actor_rollout_ref.rollout.drafter.training.lr_warmup_steps=0 \ actor_rollout_ref.rollout.drafter.training.warmup_style=constant \ actor_rollout_ref.rollout.drafter.training.use_logits=False \ - actor_rollout_ref.rollout.drafter.training.feature_store.type=torch_shard \ - actor_rollout_ref.rollout.drafter.training.feature_store.path=${FEATURE_STORE_DIR} \ - actor_rollout_ref.rollout.drafter.training.feature_store.shuffle=True \ - actor_rollout_ref.rollout.drafter.training.feature_store.repeat=True \ - actor_rollout_ref.rollout.drafter.training.feature_store.strict_schema=True $@ -fi + trainer.project_name=${project_name} \ + trainer.experiment_name=${exp_name} \ + "$@" diff --git a/tests/examples/test_example_scripts.py b/tests/examples/test_example_scripts.py index dbb70a06..ade09ffc 100644 --- a/tests/examples/test_example_scripts.py +++ b/tests/examples/test_example_scripts.py @@ -22,6 +22,12 @@ ROOT = Path(__file__).resolve().parents[2] EXAMPLES = sorted((ROOT / "examples").glob("*.sh")) +PPO_EXAMPLES = [ + script + for script in EXAMPLES + if not script.name.endswith("_separate_training.sh") + and script.name != "run_dspark_tq_producer.sh" +] def _require_working_bash() -> str: @@ -40,7 +46,7 @@ def test_example_shell_syntax_is_valid(script: Path) -> None: subprocess.run([bash, "-n", str(script)], check=True) -@pytest.mark.parametrize("script", EXAMPLES, ids=lambda path: path.name) +@pytest.mark.parametrize("script", PPO_EXAMPLES, ids=lambda path: path.name) def test_example_keeps_speco_entrypoint_and_required_drafter_switches( script: Path, ) -> None: @@ -62,6 +68,35 @@ def test_example_keeps_speco_entrypoint_and_required_drafter_switches( assert "actor_rollout_ref.rollout.drafter.training.publish_async=" in source +def test_standalone_tq_training_example_uses_unified_launcher() -> None: + source = ( + ROOT / "examples" / "run_qwen3-8b_drafter_separate_training.sh" + ).read_text(encoding="utf-8") + + assert "-m verl_speco.standalone_tq_training_launcher" in source + assert "data.train_files=${TRAIN_FILE}" in source + assert "actor_rollout_ref.rollout.drafter.enable=True" in source + assert "actor_rollout_ref.rollout.drafter.enable_drafter_training=True" in source + assert "actor_rollout_ref.rollout.drafter.model_path=${DRAFTER_PATH}" in source + assert "actor_rollout_ref.rollout.drafter.speculative_algorithm=DSPARK" in source + + +def test_standalone_tq_compatibility_example_delegates_to_formal_entry() -> None: + source = ( + ROOT / "examples" / "run_qwen3-8b_drafter_dspark_separate_training.sh" + ).read_text(encoding="utf-8") + + assert 'run_qwen3-8b_drafter_separate_training.sh" "$@"' in source + + +def test_standalone_tq_producer_example_uses_producer_entrypoint() -> None: + source = (ROOT / "examples" / "run_dspark_tq_producer.sh").read_text( + encoding="utf-8" + ) + + assert "-m verl_speco.standalone_tq_producer" in source + + def test_vllm_eagle3_example_keeps_runtime_agnostic_training_switches() -> None: source = (ROOT / "examples" / "run_qwen3-8b_drafter_eagle3_vllm.sh").read_text( encoding="utf-8" diff --git a/tests/unit/test_producer_input_reader.py b/tests/unit/test_producer_input_reader.py new file mode 100644 index 00000000..a79ff659 --- /dev/null +++ b/tests/unit/test_producer_input_reader.py @@ -0,0 +1,220 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from __future__ import annotations + +import json +from pathlib import Path +from types import SimpleNamespace + +import pytest + +from verl_speco.producer import input_reader + + +def test_iter_input_records_reads_jsonl(tmp_path: Path) -> None: + input_path = tmp_path / "train.jsonl" + input_path.write_text( + json.dumps({"prompt": "Q: ", "response": "A", "source": "test"}) + "\n", + encoding="utf-8", + ) + + records = list(input_reader.iter_input_records(input_path)) + + assert len(records) == 1 + assert records[0].sample_id == "train-000000" + assert records[0].prompt == "Q: " + assert records[0].response == "A" + assert records[0].source_metadata == {"source": "test"} + + +def test_iter_input_records_reads_parquet_rows(monkeypatch, tmp_path: Path) -> None: + input_path = tmp_path / "train.parquet" + input_path.write_bytes(b"PAR1-test-fixture") + rows = [ + {"sample_id": "first", "prompt": "Q1: ", "response": "A1"}, + {"prompt": "Q2: ", "response": "A2", "split": "train"}, + ] + + class FakeParquetFile: + def __init__(self, path: Path) -> None: + assert path == input_path + + def iter_batches(self): + yield SimpleNamespace(to_pylist=lambda: rows) + + monkeypatch.setattr( + input_reader.importlib, + "import_module", + lambda name: SimpleNamespace(ParquetFile=FakeParquetFile), + ) + + records = list(input_reader.iter_input_records(input_path)) + + assert [record.sample_id for record in records] == ["first", "train-000001"] + assert records[1].source_metadata == {"split": "train"} + + +def test_iter_input_records_reads_real_parquet_when_available(tmp_path: Path) -> None: + pyarrow = pytest.importorskip("pyarrow") + parquet = pytest.importorskip("pyarrow.parquet") + input_path = tmp_path / "train.parquet" + parquet.write_table( + pyarrow.Table.from_pylist( + [{"prompt": "real Q: ", "response": "real A", "split": "train"}] + ), + input_path, + ) + + records = list(input_reader.iter_input_records(input_path)) + + assert len(records) == 1 + assert records[0].prompt == "real Q: " + assert records[0].response == "real A" + assert records[0].source_metadata == {"split": "train"} + + +def test_iter_input_records_reads_real_dapo_style_parquet_when_available( + tmp_path: Path, +) -> None: + pyarrow = pytest.importorskip("pyarrow") + parquet = pytest.importorskip("pyarrow.parquet") + input_path = tmp_path / "dapo.parquet" + parquet.write_table( + pyarrow.Table.from_pylist( + [ + { + "data_source": "math_dapo", + "prompt": [{"role": "user", "content": "Solve Q"}], + "reward_model": { + "ground_truth": "42", + "style": "rule-lighteval/MATH_v2", + }, + "extra_info": {"index": "dapo-real-row"}, + } + ] + ), + input_path, + ) + + record = next(input_reader.iter_input_records(input_path)) + + assert record.prompt == ({"role": "user", "content": "Solve Q"},) + assert record.response is None + assert record.sample_id == "dapo-real-row" + + +def test_dapo_parquet_prompt_is_prepared_for_target_generation( + monkeypatch, tmp_path: Path +) -> None: + input_path = tmp_path / "train.parquet" + input_path.write_bytes(b"PAR1-test-fixture") + + class FakeParquetFile: + def __init__(self, path: Path) -> None: + assert path == input_path + + def iter_batches(self): + yield SimpleNamespace( + to_pylist=lambda: [ + { + "prompt": [{"role": "user", "content": "Solve Q"}], + "reward_model": {"ground_truth": "42"}, + "extra_info": {"index": "dapo-row-id"}, + } + ] + ) + + monkeypatch.setattr( + input_reader.importlib, + "import_module", + lambda name: SimpleNamespace(ParquetFile=FakeParquetFile), + ) + + record = next(input_reader.iter_input_records(input_path)) + + class ChatTokenizer: + def apply_chat_template(self, messages, *, tokenize, add_generation_prompt): + assert messages == [{"role": "user", "content": "Solve Q"}] + assert tokenize is True + assert add_generation_prompt is True + return [10, 11, 12] + + generation = input_reader.prepare_generation_request( + record, + ChatTokenizer(), + {"max_sequence_length": 16, "generation_max_tokens": 4}, + ) + finalized = input_reader.finalize_generated_request( + generation, + [10, 11, 12, 20, 21], + {"max_sequence_length": 16, "max_feature_length": 8}, + ) + + assert record.sample_id == "dapo-row-id" + assert record.response is None + assert generation.prompt_token_ids == (10, 11, 12) + assert generation.max_tokens == 4 + assert finalized.input_ids.tolist() == [10, 11, 12, 20, 21] + assert finalized.loss_mask.tolist() == [0, 0, 0, 1, 1] + assert finalized.prompt_token_ids == [10, 11, 12, 20, 21] + + +def test_finalize_generated_request_aligns_connector_excluding_final_token() -> None: + request = input_reader.GenerationRequest( + sequence_no=0, + sample_id="generated-row", + prompt_token_ids=(10, 11, 12), + max_tokens=4, + source_metadata={}, + ) + + finalized = input_reader.finalize_generated_request( + request, + [10, 11, 12, 20], + {"max_sequence_length": 16, "max_feature_length": 8}, + expected_response_token_ids=[20, 21], + ) + + assert finalized.input_ids.tolist() == [10, 11, 12, 20, 21] + assert finalized.loss_mask.tolist() == [0, 0, 0, 1, 1] + assert finalized.prompt_token_ids == [10, 11, 12, 20] + assert finalized.feature_positions.tolist() == [2, 3] + assert finalized.draft_position_ids.tolist() == [3, 4] + + +def test_finalize_generated_request_rejects_misaligned_connector_tokens() -> None: + request = input_reader.GenerationRequest( + sequence_no=0, + sample_id="generated-row", + prompt_token_ids=(10, 11), + max_tokens=4, + source_metadata={}, + ) + + with pytest.raises(ValueError, match="excluding its final token"): + input_reader.finalize_generated_request( + request, + [10, 11, 99], + {"max_sequence_length": 16, "max_feature_length": 8}, + expected_response_token_ids=[20, 21], + ) + + +def test_non_utf8_non_parquet_input_has_actionable_error(tmp_path: Path) -> None: + input_path = tmp_path / "train.data" + input_path.write_bytes(b"plain-prefix\xc0binary") + + with pytest.raises(ValueError, match="not UTF-8 JSONL or a Parquet file"): + list(input_reader.iter_input_records(input_path)) diff --git a/tests/unit/test_standalone_tq_training_launcher.py b/tests/unit/test_standalone_tq_training_launcher.py new file mode 100644 index 00000000..38e53960 --- /dev/null +++ b/tests/unit/test_standalone_tq_training_launcher.py @@ -0,0 +1,290 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from __future__ import annotations + +from pathlib import Path +import json +import threading + +from omegaconf import OmegaConf +import pytest + +from verl_speco.standalone_tq_training_launcher import ( + _preflight_input_file, + _target_final_layer_id, + build_pipeline_commands, + resolve_pipeline_config, + run_pipeline, + start_ray_session, +) +import verl_speco.tq_owner as tq_owner + + +def _training_args() -> list[str]: + return [ + "data.train_files=/data/train.jsonl", + "actor_rollout_ref.model.path=/models/Qwen3-8B", + "actor_rollout_ref.rollout.drafter.model_path=/models/dspark", + "actor_rollout_ref.rollout.drafter.speculative_algorithm=DSPARK", + "actor_rollout_ref.rollout.drafter.training.max_steps=10", + ] + + +def test_pipeline_config_derives_transport_identity_from_training_args() -> None: + config = resolve_pipeline_config(_training_args(), environ={}) + + assert config.input_path == "/data/train.jsonl" + assert config.model_path == "/models/Qwen3-8B" + assert config.tokenizer_path == "/models/Qwen3-8B" + assert config.algorithm == "DSPARK" + assert config.target_layer_ids == (1, 9, 17, 25, 33) + assert config.vllm_endpoint == "http://127.0.0.1:8000/v1" + assert config.run_id.startswith("dspark-") + + +def test_pipeline_config_accepts_one_hydra_list_train_file() -> None: + args = _training_args() + args[0] = "data.train_files=['/data/train.jsonl']" + + config = resolve_pipeline_config(args, environ={}) + + assert config.input_path == "/data/train.jsonl" + + +def test_target_final_layer_id_uses_local_model_config(tmp_path) -> None: + (tmp_path / "config.json").write_text( + json.dumps({"text_config": {"num_hidden_layers": 48}}), + encoding="utf-8", + ) + + assert _target_final_layer_id(str(tmp_path), (2, 10, 20)) == 48 + + +def test_pipeline_config_rejects_multiple_train_files() -> None: + args = _training_args() + args[0] = "data.train_files=[a.jsonl,b.jsonl]" + + with pytest.raises(ValueError, match="exactly one train file"): + resolve_pipeline_config(args, environ={}) + + +def test_preflight_accepts_verl_prompt_parquet(monkeypatch, tmp_path) -> None: + input_path = tmp_path / "train.parquet" + input_path.write_bytes(b"PAR1-test-fixture") + + class FakeParquetFile: + def __init__(self, path: Path) -> None: + assert path == input_path + + def iter_batches(self): + class Batch: + @staticmethod + def to_pylist(): + return [ + { + "prompt": [{"role": "user", "content": "Solve Q"}], + "reward_model": {"ground_truth": "42"}, + } + ] + + yield Batch() + + from verl_speco.producer import input_reader + + monkeypatch.setattr( + input_reader.importlib, + "import_module", + lambda name: type("ParquetModule", (), {"ParquetFile": FakeParquetFile}), + ) + + _preflight_input_file(str(input_path)) + + +def test_pipeline_commands_hide_and_replace_tq_overrides() -> None: + args = [ + *_training_args(), + "actor_rollout_ref.rollout.drafter.training.feature_store.type=torch_shard", + "actor_rollout_ref.rollout.drafter.training.transfer_queue.run_id=user-value", + ] + config = resolve_pipeline_config(args, environ={}) + + commands = build_pipeline_commands( + config, + args, + ray_address="10.0.0.1:6379", + python_executable="python", + ) + + assert commands.owner[:3] == ["python", "-m", "verl_speco.tq_owner"] + assert commands.vllm is not None + assert commands.vllm[:3] == ["vllm", "serve", "/models/Qwen3-8B"] + assert "ExampleHiddenStatesConnector" in " ".join(commands.vllm) + assert commands.vllm_endpoint == "http://127.0.0.1:8000/v1" + assert commands.producer[:3] == [ + "python", + "-m", + "verl_speco.standalone_tq_producer", + ] + assert commands.consumer[:3] == [ + "python", + "-m", + "verl_speco.draft_train_launcher", + ] + assert any(item.endswith("feature_store.type=tq") for item in commands.consumer) + assert not any("run_id=user-value" in item for item in commands.consumer) + assert any(f"run_id={config.run_id}" in item for item in commands.consumer) + assert any( + "vllm_endpoints=[http://127.0.0.1:8000/v1]" in item + for item in commands.producer + ) + + +class _FakeRuntimeContext: + gcs_address = "127.0.0.1:61234" + + +class _FakeRay: + def __init__(self) -> None: + self.init_kwargs = None + self.shutdown_called = False + + def init(self, **kwargs) -> None: + self.init_kwargs = kwargs + + def get_runtime_context(self) -> _FakeRuntimeContext: + return _FakeRuntimeContext() + + def shutdown(self) -> None: + self.shutdown_called = True + + +def test_ray_session_starts_local_control_plane_without_exposed_address() -> None: + ray = _FakeRay() + + session = start_ray_session( + environ={"RAY_ADDRESS": "172.51.9.253:35195"}, ray_module=ray + ) + + assert session.address == "127.0.0.1:61234" + assert ray.init_kwargs == { + "address": "local", + "namespace": "speco-drafter", + "include_dashboard": False, + } + session.close() + assert ray.shutdown_called + + +class _FakeProcess: + def __init__(self, role: str) -> None: + self.role = role + self.returncode = 0 if role == "consumer" else None + self.terminated = False + + def poll(self): + return self.returncode + + def terminate(self) -> None: + self.terminated = True + self.returncode = -15 + + def wait(self, timeout=None): + return self.returncode + + def kill(self) -> None: + self.returncode = -9 + + +def test_pipeline_starts_owner_then_consumer_then_producer() -> None: + started: list[str] = [] + processes: list[_FakeProcess] = [] + child_environments: list[dict[str, str]] = [] + vllm_started = False + + def fake_popen(command, *, env): + nonlocal vllm_started + if command[0] == "vllm": + role = "vllm" + vllm_started = True + else: + module = command[2] + role = { + "verl_speco.tq_owner": "owner", + "verl_speco.draft_train_launcher": "consumer", + "verl_speco.standalone_tq_producer": "producer", + }[module] + process = _FakeProcess(role) + started.append(role) + processes.append(process) + child_environments.append(dict(env)) + if role == "owner": + Path(env["SPECO_TQ_OWNER_READY_FILE"]).touch() + return process + + config = resolve_pipeline_config(_training_args(), environ={}) + commands = build_pipeline_commands( + config, + _training_args(), + ray_address="127.0.0.1:61234", + ) + + assert ( + run_pipeline( + commands, + ray_address="127.0.0.1:61234", + environ={"RAY_ADDRESS": "172.51.9.253:35195"}, + popen=fake_popen, + endpoint_ready=lambda _: vllm_started, + ) + == 0 + ) + assert started == ["vllm", "owner", "consumer", "producer"] + assert all(env["RAY_ADDRESS"] == "127.0.0.1:61234" for env in child_environments) + assert all(process.poll() is not None for process in processes) + + +def test_owner_writes_internal_ready_file(monkeypatch, tmp_path) -> None: + ready_file = tmp_path / "owner.ready" + monkeypatch.setenv("SPECO_TQ_OWNER_READY_FILE", str(ready_file)) + monkeypatch.setattr(tq_owner, "configure_transfer_queue", lambda config: True) + monkeypatch.setattr(tq_owner, "connect_ray_cluster", lambda *args: None) + monkeypatch.setattr(tq_owner, "start_transfer_queue_owner", lambda config: None) + monkeypatch.setattr(tq_owner, "close_transfer_queue_owner", lambda: None) + monkeypatch.setattr(tq_owner, "publish_owner_ready", lambda *args: "ready-key") + stop_event = threading.Event() + stop_event.set() + config = OmegaConf.create( + { + "actor_rollout_ref": { + "rollout": { + "drafter": { + "training": { + "transfer_queue": { + "enable": True, + "run_id": "test-run", + "schema_version": 1, + "ray": { + "address": "127.0.0.1:6379", + "namespace": "speco-drafter", + }, + } + } + } + } + } + } + ) + + assert tq_owner.run_owner(config, stop_event=stop_event) == 0 + assert ready_file.is_file() diff --git a/tests/unit/test_tq_producer.py b/tests/unit/test_tq_producer.py new file mode 100644 index 00000000..6ac2fd43 --- /dev/null +++ b/tests/unit/test_tq_producer.py @@ -0,0 +1,332 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from __future__ import annotations + +import asyncio +import json +from pathlib import Path +from typing import Any + +import pytest +import torch + +from verl_speco.producer.vllm_feature_client import RawVllmFeature +from verl_speco.standalone_tq_producer import run_producer, validate_producer_config + + +def _config(input_path: Path) -> dict[str, Any]: + return { + "speco": { + "standalone_tq_producer": { + "input_path": str(input_path), + "tokenizer_path": "/target", + "tokenizer_fingerprint": "sha256:tokenizer", + "target_model_id": "/target", + "target_model_revision": "rev-a", + "target_layer_ids": [2, 8], + "hidden_dtype": "float32", + "trust_remote_code": False, + "vllm_endpoints": ["http://vllm:8000/v1"], + "vllm_model": "/target", + "request_timeout": 10, + "max_inflight_requests": 2, + "per_endpoint_concurrency": 1, + "input_queue_size": 2, + "publish_queue_size": 2, + "max_pending_samples": 8, + "pending_poll_interval_seconds": 0.01, + "owner_ready_timeout_seconds": 1, + "max_sequence_length": 16, + "max_feature_length": 8, + "generation_max_tokens": 4, + } + }, + "actor_rollout_ref": { + "rollout": { + "drafter": { + "speculative_algorithm": "DSPARK", + "training": { + "use_logits": False, + "dspark_l1_loss_alpha": 0.9, + "transfer_queue": { + "enable": True, + "package_version": "0.1.7", + "ray": { + "address": "ray-head:6379", + "namespace": "speco-drafter", + }, + "partition_id": "speco_drafter_features", + "run_id": "run-a", + "schema_version": 1, + }, + }, + } + } + }, + } + + +class _Tokenizer: + def __call__(self, text: str, *, add_special_tokens: bool) -> dict[str, list[int]]: + assert add_special_tokens is False + values = { + "Q1: ": [1, 2], + "Q1: A1": [1, 2, 3, 4], + "Q2: ": [5, 6], + "Q2: A2": [5, 6, 7, 8], + } + return {"input_ids": values[text]} + + +class _ChatTokenizer: + def apply_chat_template(self, messages, *, tokenize, add_generation_prompt): + assert messages == [{"role": "user", "content": "Q3"}] + assert tokenize is True + assert add_generation_prompt is True + return [9, 10] + + +class _Transport: + def __init__(self, *, fail_sample_put: bool = False): + self.fail_sample_put = fail_sample_put + self.records: dict[str, dict[str, Any]] = { + "control:v1:run-a:owner-ready": { + "record_type": "control", + "status": "owner_ready", + "schema_version": 1, + "run_id": "run-a", + } + } + self.payloads: dict[str, dict[str, torch.Tensor]] = {} + self.closed = False + + def configure_transfer_queue(self, config: dict[str, Any]) -> bool: + return bool(config["enable"]) + + def connect_ray_cluster(self, address: str, namespace: str | None) -> None: + assert (address, namespace) == ("ray-head:6379", "speco-drafter") + + def connect_transfer_queue_client(self) -> None: + pass + + def list_samples(self) -> dict[str, dict[str, Any]]: + return dict(self.records) + + def put_sample( + self, + key: str, + fields: dict[str, torch.Tensor], + *, + tag: dict[str, Any], + ) -> None: + if self.fail_sample_put and tag.get("record_type") == "sample": + raise RuntimeError("put failed") + self.records[key] = dict(tag) + self.payloads[key] = fields + + def close_transfer_queue_client(self) -> None: + self.closed = True + + +class _Pool: + def __init__(self, root: Path, *, close_error: BaseException | None = None): + self.root = root + self.close_error = close_error + self.paths: list[Path] = [] + self.started = False + self.closed = False + self.generate_calls = 0 + + async def start(self) -> None: + self.started = True + + async def prefill(self, request: Any) -> RawVllmFeature: + path = self.root / f"{request.sample_id}.safetensors" + path.write_bytes(b"temporary") + self.paths.append(path) + token_ids = torch.tensor(request.prompt_token_ids, dtype=torch.int64) + hidden = torch.arange(token_ids.numel() * 3 * 2, dtype=torch.float32).reshape( + token_ids.numel(), 3, 2 + ) + return RawVllmFeature( + payload={"token_ids": token_ids, "hidden_states": hidden}, + temporary_path=str(path), + endpoint_url="http://vllm:8000/v1", + byte_size=path.stat().st_size, + ) + + async def generate(self, request: Any) -> RawVllmFeature: + self.generate_calls += 1 + path = self.root / f"{request.sample_id}.safetensors" + path.write_bytes(b"temporary") + self.paths.append(path) + # ExampleHiddenStatesConnector excludes the final generated token because + # it was never consumed by a model forward pass. + token_ids = torch.tensor([*request.prompt_token_ids, 11], dtype=torch.int64) + hidden = torch.arange(token_ids.numel() * 3 * 2, dtype=torch.float32).reshape( + token_ids.numel(), 3, 2 + ) + return RawVllmFeature( + payload={"token_ids": token_ids, "hidden_states": hidden}, + temporary_path=str(path), + endpoint_url="http://vllm:8000/v1", + byte_size=path.stat().st_size, + generated_token_ids=(11, 12), + ) + + async def close(self) -> None: + self.closed = True + if self.close_error is not None: + raise self.close_error + + +def _write_input(path: Path) -> None: + records = [ + {"sample_id": "sample-1", "prompt": "Q1: ", "response": "A1"}, + {"sample_id": "sample-2", "prompt": "Q2: ", "response": "A2"}, + ] + path.write_text( + "".join(json.dumps(record) + "\n" for record in records), + encoding="utf-8", + ) + + +def test_run_producer_publishes_samples_then_eos(tmp_path: Path) -> None: + input_path = tmp_path / "input.jsonl" + _write_input(input_path) + transport = _Transport() + pool = _Pool(tmp_path) + + stats = asyncio.run( + run_producer( + _config(input_path), + transport=transport, + tokenizer=_Tokenizer(), + client_pool=pool, + ) + ) + + sample_keys = [ + key + for key, tag in transport.records.items() + if tag.get("record_type") == "sample" + ] + eos_tags = [tag for tag in transport.records.values() if tag.get("status") == "eos"] + assert stats.input_count == stats.published_count == 2 + assert stats.failed_count == stats.pending_bytes == 0 + assert len(sample_keys) == 2 + assert eos_tags == [ + { + "record_type": "control", + "status": "eos", + "schema_version": 1, + "run_id": "run-a", + "total_samples": 2, + } + ] + assert all(not path.exists() for path in pool.paths) + assert pool.started and pool.closed and transport.closed + first_fields = transport.payloads[sorted(sample_keys)[0]] + assert tuple(first_fields["hidden_states"].shape) == (3, 6) + + +def test_run_producer_generates_response_for_verl_chat_prompt(tmp_path: Path) -> None: + input_path = tmp_path / "dapo.jsonl" + input_path.write_text( + json.dumps( + { + "prompt": [{"role": "user", "content": "Q3"}], + "reward_model": {"ground_truth": "42"}, + "extra_info": {"index": "dapo-row"}, + } + ) + + "\n", + encoding="utf-8", + ) + transport = _Transport() + pool = _Pool(tmp_path) + + stats = asyncio.run( + run_producer( + _config(input_path), + transport=transport, + tokenizer=_ChatTokenizer(), + client_pool=pool, + ) + ) + + sample_keys = [ + key + for key, tag in transport.records.items() + if tag.get("record_type") == "sample" + ] + assert stats.input_count == stats.published_count == 1 + assert pool.generate_calls == 1 + assert len(sample_keys) == 1 + fields = transport.payloads[sample_keys[0]] + assert fields["input_ids"].tolist() == [10, 11] + assert fields["loss_mask"].tolist() == [0.0, 1.0] + + +def test_run_producer_put_failure_keeps_temporary_file_and_omits_eos( + tmp_path: Path, +) -> None: + input_path = tmp_path / "input.jsonl" + _write_input(input_path) + transport = _Transport(fail_sample_put=True) + pool = _Pool(tmp_path) + + with pytest.raises(RuntimeError, match="put failed"): + asyncio.run( + run_producer( + _config(input_path), + transport=transport, + tokenizer=_Tokenizer(), + client_pool=pool, + ) + ) + + assert any(path.exists() for path in pool.paths) + assert not any(tag.get("status") == "eos" for tag in transport.records.values()) + assert pool.closed and transport.closed + + +def test_validate_producer_rejects_consumer_partition_mismatch(tmp_path: Path) -> None: + config = _config(tmp_path / "input.jsonl") + config["actor_rollout_ref"]["rollout"]["drafter"]["training"]["transfer_queue"][ + "partition_id" + ] = "other" + + with pytest.raises(ValueError, match="partition_id"): + validate_producer_config(config) + + +def test_pool_close_failure_does_not_skip_transport_close(tmp_path: Path) -> None: + input_path = tmp_path / "input.jsonl" + _write_input(input_path) + transport = _Transport() + pool = _Pool(tmp_path, close_error=RuntimeError("pool close failed")) + + with pytest.raises(RuntimeError, match="pool close failed"): + asyncio.run( + run_producer( + _config(input_path), + transport=transport, + tokenizer=_Tokenizer(), + client_pool=pool, + ) + ) + + assert pool.closed and transport.closed diff --git a/tests/unit/test_transferqueue_bridge.py b/tests/unit/test_transferqueue_bridge.py index 66049d68..f792819c 100644 --- a/tests/unit/test_transferqueue_bridge.py +++ b/tests/unit/test_transferqueue_bridge.py @@ -14,7 +14,6 @@ from __future__ import annotations import sys -from types import SimpleNamespace import pytest import torch @@ -144,7 +143,13 @@ def test_client_put_list_get_many_clear_and_local_close(fake_runtime) -> None: assert bridge.configure_transfer_queue(_config()) bridge.connect_ray_cluster("ray-head:6379", "speco-drafter") bridge.connect_transfer_queue_client() - assert fake_tq.init_calls == [None] + native = bridge._to_plain_dict(fake_tq.init_calls[0]) + assert set(native) == {"controller", "backend"} + assert native["backend"]["storage_backend"] == "SimpleStorage" + assert native["backend"]["SimpleStorage"] == { + "total_storage_size": 16, + "num_data_storage_units": 1, + } bridge.put_sample( "k0", diff --git a/tests/unit/test_vllm_feature_client.py b/tests/unit/test_vllm_feature_client.py new file mode 100644 index 00000000..b09bb601 --- /dev/null +++ b/tests/unit/test_vllm_feature_client.py @@ -0,0 +1,60 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from __future__ import annotations + +import asyncio +from types import SimpleNamespace + +from verl_speco.producer.vllm_feature_client import ( + VllmEndpoint, + request_generate, +) + + +def test_request_generate_includes_output_tokens_in_hidden_states() -> None: + calls = [] + + class Completions: + async def create(self, **kwargs): + calls.append(kwargs) + return SimpleNamespace( + choices=[ + SimpleNamespace( + prompt_token_ids=[1, 2], + token_ids=[3, 4], + ) + ], + kv_transfer_params={"hidden_states_path": "/tmp/result.safetensors"}, + ) + + client = SimpleNamespace(completions=Completions()) + + response = asyncio.run( + request_generate( + VllmEndpoint("http://vllm:8000/v1", 1), + client, + [1, 2], + model="target", + max_tokens=128, + timeout=30, + ) + ) + + assert response.generated_token_ids == (3, 4) + assert calls[0]["max_tokens"] == 128 + assert calls[0]["extra_body"] == { + "return_token_ids": True, + "kv_transfer_params": {"include_output_tokens": True}, + } diff --git a/tools/run_dspark_tq_consumer.sh b/tools/run_dspark_tq_consumer.sh index 24e1d1d8..faeccdaa 100644 --- a/tools/run_dspark_tq_consumer.sh +++ b/tools/run_dspark_tq_consumer.sh @@ -1,4 +1,18 @@ -set -euo pipefail +#!/usr/bin/env bash +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +set -euo pipefail set -x # Standalone DSpark Consumer. Start Ray and verl_speco.tq_owner first, then run @@ -14,10 +28,11 @@ DSPARK_NUM_TARGET_LAYERS=${DSPARK_NUM_TARGET_LAYERS:-5} RAY_ADDRESS=${RAY_ADDRESS:-127.0.0.1:6379} TQ_NAMESPACE=${TQ_NAMESPACE:-speco-drafter} TQ_PARTITION_ID=${TQ_PARTITION_ID:-speco_drafter_features} -SPECO_TQ_RUN_ID=${SPECO_TQ_RUN_ID:-dspark-standalone-run} - -CUDA_VISIBLE_DEVICES=${TRAIN_DEVICES} PYTHONUNBUFFERED=1 \ -python3 -m verl_speco.draft_train_launcher \ +SPECO_TQ_RUN_ID=${SPECO_TQ_RUN_ID:-dspark-standalone-run} +PYTHON_BIN=${PYTHON_BIN:-python3} + +CUDA_VISIBLE_DEVICES=${TRAIN_DEVICES} PYTHONUNBUFFERED=1 \ +exec "${PYTHON_BIN}" -m verl_speco.draft_train_launcher \ speco.draft_training.nproc_per_node=${TRAIN_GPUS} \ speco.draft_training.nnodes=1 \ actor_rollout_ref.model.path=${MODEL_PATH} \ diff --git a/tools/run_dspark_tq_e2e_test.sh b/tools/run_dspark_tq_e2e_test.sh new file mode 100644 index 00000000..eef3c026 --- /dev/null +++ b/tools/run_dspark_tq_e2e_test.sh @@ -0,0 +1,170 @@ +#!/usr/bin/env bash +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# Run the real standalone TQ Producer and Consumer in one end-to-end test. +set -euo pipefail + +script_dir=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd) +repo_root=$(cd -- "${script_dir}/.." && pwd) +cd "${repo_root}" + +: "${RAY_ADDRESS:?Set RAY_ADDRESS to the running Ray head}" +: "${MODEL_PATH:?Set MODEL_PATH to the target model directory}" +: "${DRAFTER_PATH:?Set DRAFTER_PATH to the DSpark drafter directory}" +: "${PRODUCER_INPUT_PATH:?Set PRODUCER_INPUT_PATH to prompt/response JSONL}" +: "${TARGET_MODEL_REVISION:?Set TARGET_MODEL_REVISION to a revision or checksum}" +: "${TOKENIZER_FINGERPRINT:?Set TOKENIZER_FINGERPRINT to a verified fingerprint}" +: "${TARGET_LAYER_IDS:?Set TARGET_LAYER_IDS as a Hydra list, for example '[2,8,14,20,26]'}" +: "${VLLM_ENDPOINTS:?Set VLLM_ENDPOINTS as a Hydra list, for example '[http://node0:8000/v1]'}" + +PYTHON_BIN=${PYTHON_BIN:-python3} +TOKENIZER_PATH=${TOKENIZER_PATH:-${MODEL_PATH}} +TARGET_MODEL_PATH=${TARGET_MODEL_PATH:-${MODEL_PATH}} +TQ_NAMESPACE=${TQ_NAMESPACE:-speco-drafter} +TQ_PARTITION_ID=${TQ_PARTITION_ID:-speco_drafter_features} +SPECO_TQ_RUN_ID=${SPECO_TQ_RUN_ID:-dspark-e2e-$(date +%Y%m%d-%H%M%S)-$$} +TRAIN_DEVICES=${TRAIN_DEVICES:-0} +TRAIN_GPUS=${TRAIN_GPUS:-1} +BATCH_SIZE_PER_GPU=${BATCH_SIZE_PER_GPU:-1} +DSPARK_NUM_TARGET_LAYERS=${DSPARK_NUM_TARGET_LAYERS:-5} +DRAFT_CKPTS_DIR=${DRAFT_CKPTS_DIR:-/tmp/speco-dspark-tq-e2e-${SPECO_TQ_RUN_ID}} +E2E_TIMEOUT_SECONDS=${E2E_TIMEOUT_SECONDS:-1800} + +if [[ "${TQ_PARTITION_ID}" != "speco_drafter_features" ]]; then + echo "TQ_PARTITION_ID must be speco_drafter_features for protocol v1" >&2 + exit 2 +fi +if [[ ! -f "${PRODUCER_INPUT_PATH}" ]]; then + echo "Producer input does not exist: ${PRODUCER_INPUT_PATH}" >&2 + exit 2 +fi +input_samples=$(awk 'NF { count++ } END { print count + 0 }' "${PRODUCER_INPUT_PATH}") +global_batch_size=$((TRAIN_GPUS * BATCH_SIZE_PER_GPU)) +if (( input_samples < global_batch_size )); then + echo "Producer input has ${input_samples} non-empty records, but one Consumer global batch needs ${global_batch_size}" >&2 + exit 2 +fi + +work_dir=$(mktemp -d "${TMPDIR:-/tmp}/speco-tq-e2e.XXXXXX") +owner_pid="" +consumer_pid="" +producer_pid="" + +cleanup() { + local pid + for pid in "${producer_pid}" "${consumer_pid}" "${owner_pid}"; do + if [[ -n "${pid}" ]] && kill -0 "${pid}" 2>/dev/null; then + kill "${pid}" 2>/dev/null || true + wait "${pid}" 2>/dev/null || true + fi + done + rm -rf -- "${work_dir}" +} +trap cleanup EXIT INT TERM + +echo "[1/4] Starting TQ owner run_id=${SPECO_TQ_RUN_ID}" +RAY_ADDRESS="${RAY_ADDRESS}" \ +TQ_NAMESPACE="${TQ_NAMESPACE}" \ +TQ_PARTITION_ID="${TQ_PARTITION_ID}" \ +SPECO_TQ_RUN_ID="${SPECO_TQ_RUN_ID}" \ + bash tools/run_dspark_tq_owner.sh & +owner_pid=$! + +echo "[2/4] Starting real DSpark Consumer" +( + set +e + MODEL_PATH="${MODEL_PATH}" \ + DRAFTER_PATH="${DRAFTER_PATH}" \ + DRAFT_CKPTS_DIR="${DRAFT_CKPTS_DIR}" \ + TRAIN_DEVICES="${TRAIN_DEVICES}" \ + TRAIN_GPUS="${TRAIN_GPUS}" \ + BATCH_SIZE_PER_GPU="${BATCH_SIZE_PER_GPU}" \ + MAX_STEPS=0 \ + DSPARK_NUM_TARGET_LAYERS="${DSPARK_NUM_TARGET_LAYERS}" \ + RAY_ADDRESS="${RAY_ADDRESS}" \ + TQ_NAMESPACE="${TQ_NAMESPACE}" \ + TQ_PARTITION_ID="${TQ_PARTITION_ID}" \ + SPECO_TQ_RUN_ID="${SPECO_TQ_RUN_ID}" \ + bash tools/run_dspark_tq_consumer.sh "$@" + echo "$?" > "${work_dir}/consumer.status" +) & +consumer_pid=$! + +echo "[3/4] Starting real vLLM-backed Producer" +( + set +e + PYTHON_BIN="${PYTHON_BIN}" \ + RAY_ADDRESS="${RAY_ADDRESS}" \ + TQ_NAMESPACE="${TQ_NAMESPACE}" \ + TQ_PARTITION_ID="${TQ_PARTITION_ID}" \ + SPECO_TQ_RUN_ID="${SPECO_TQ_RUN_ID}" \ + PRODUCER_INPUT_PATH="${PRODUCER_INPUT_PATH}" \ + TARGET_MODEL_PATH="${TARGET_MODEL_PATH}" \ + TARGET_MODEL_REVISION="${TARGET_MODEL_REVISION}" \ + TOKENIZER_PATH="${TOKENIZER_PATH}" \ + TOKENIZER_FINGERPRINT="${TOKENIZER_FINGERPRINT}" \ + TARGET_LAYER_IDS="${TARGET_LAYER_IDS}" \ + VLLM_ENDPOINTS="${VLLM_ENDPOINTS}" \ + bash examples/run_dspark_tq_producer.sh + echo "$?" > "${work_dir}/producer.status" +) & +producer_pid=$! + +started_at=${SECONDS} +while [[ ! -f "${work_dir}/producer.status" || ! -f "${work_dir}/consumer.status" ]]; do + if ! kill -0 "${owner_pid}" 2>/dev/null; then + echo "TQ owner exited before the end-to-end test completed" >&2 + exit 1 + fi + if (( SECONDS - started_at >= E2E_TIMEOUT_SECONDS )); then + echo "E2E timed out after ${E2E_TIMEOUT_SECONDS} seconds" >&2 + exit 124 + fi + if [[ -f "${work_dir}/producer.status" ]]; then + producer_status=$(<"${work_dir}/producer.status") + if [[ "${producer_status}" -ne 0 ]]; then + echo "Producer failed with exit code ${producer_status}" >&2 + exit "${producer_status}" + fi + fi + if [[ -f "${work_dir}/consumer.status" ]]; then + consumer_status=$(<"${work_dir}/consumer.status") + if [[ "${consumer_status}" -ne 0 ]]; then + echo "Consumer failed with exit code ${consumer_status}" >&2 + exit "${consumer_status}" + fi + fi + sleep 1 +done + +producer_status=$(<"${work_dir}/producer.status") +consumer_status=$(<"${work_dir}/consumer.status") +if [[ "${producer_status}" -ne 0 || "${consumer_status}" -ne 0 ]]; then + echo "E2E failed: producer=${producer_status} consumer=${consumer_status}" >&2 + exit 1 +fi + +wait "${producer_pid}" +producer_pid="" +wait "${consumer_pid}" +consumer_pid="" + +echo "[4/4] Producer published EOS and Consumer drained the run; stopping owner" +kill "${owner_pid}" 2>/dev/null || true +wait "${owner_pid}" 2>/dev/null || true +owner_pid="" +trap - EXIT INT TERM +rm -rf -- "${work_dir}" + +echo "DSPARK_TQ_E2E_TEST_OK run_id=${SPECO_TQ_RUN_ID} checkpoints=${DRAFT_CKPTS_DIR}" diff --git a/tools/run_dspark_tq_owner.sh b/tools/run_dspark_tq_owner.sh index d1c0f245..c152407b 100644 --- a/tools/run_dspark_tq_owner.sh +++ b/tools/run_dspark_tq_owner.sh @@ -18,10 +18,12 @@ set -euo pipefail : "${SPECO_TQ_RUN_ID:?Set SPECO_TQ_RUN_ID to a unique pipeline run id}" TQ_NAMESPACE=${TQ_NAMESPACE:-speco-drafter} TQ_PARTITION_ID=${TQ_PARTITION_ID:-speco_drafter_features} +PYTHON_BIN=${PYTHON_BIN:-python3} -python -m verl_speco.tq_owner \ +exec "${PYTHON_BIN}" -m verl_speco.tq_owner \ actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.address="${RAY_ADDRESS}" \ actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.namespace="${TQ_NAMESPACE}" \ actor_rollout_ref.rollout.drafter.training.transfer_queue.partition_id="${TQ_PARTITION_ID}" \ actor_rollout_ref.rollout.drafter.training.transfer_queue.run_id="${SPECO_TQ_RUN_ID}" \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.backend.storage_backend=SimpleStorage + actor_rollout_ref.rollout.drafter.training.transfer_queue.backend.storage_backend=SimpleStorage \ + "$@" diff --git a/verl_speco/config/speco_base.yaml b/verl_speco/config/speco_base.yaml index 82517aed..2c0e5aa9 100644 --- a/verl_speco/config/speco_base.yaml +++ b/verl_speco/config/speco_base.yaml @@ -21,8 +21,8 @@ speco: task_runner: verl_speco.integration.task_runner.SpecoTaskRunner ray_trainer: verl_speco.trainer.speco_ray_trainer.SpecoRayPPOTrainer - # One-process standalone Producer. It reads strict prompt/response JSONL, - # requests target hidden states from vLLM, and publishes one TQ key per row. + # One-process standalone Producer. Prompt-only verl rows are generated by the + # target vLLM; rows with a response are replayed directly for hidden states. standalone_tq_producer: input_path: null tokenizer_path: null @@ -45,6 +45,7 @@ speco: owner_ready_timeout_seconds: 120 max_sequence_length: 8192 max_feature_length: 512 + generation_max_tokens: 511 actor_rollout_ref: rollout: diff --git a/verl_speco/integration/transferqueue_bridge.py b/verl_speco/integration/transferqueue_bridge.py index 79407aa5..497e76ed 100644 --- a/verl_speco/integration/transferqueue_bridge.py +++ b/verl_speco/integration/transferqueue_bridge.py @@ -1,3 +1,17 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + """TransferQueue bridge for SPECO drafter feature transport. This module lets SPECO route large per-sample drafter-training tensors (hidden @@ -51,7 +65,6 @@ _TQ_IMPORTABLE = True except ImportError: - _TQ_IMPORTABLE = False class KVBatchMeta: # type: ignore[no-redef] @@ -78,12 +91,12 @@ def _raise(*args: Any, **kwargs: Any) -> Any: # --------------------------------------------------------------------------- _state_lock = threading.Lock() -_state = { - "enabled": False, # config says enable=true - "configured": False, # configure_transfer_queue has run - "initialized": False, # tq.init() has run in this process - "config": None, # the transfer_queue sub-config (plain dict) - "owner": False, # this process created the task-level TQ system +_state: dict[str, Any] = { + "enabled": False, # config says enable=true + "configured": False, # configure_transfer_queue has run + "initialized": False, # tq.init() has run in this process + "config": None, # the transfer_queue sub-config (plain dict) + "owner": False, # this process created the task-level TQ system "ray_initialized_here": False, "ray_address": None, "ray_namespace": None, @@ -214,7 +227,9 @@ def connect_transfer_queue_client() -> None: if not _TQ_IMPORTABLE: raise RuntimeError("TransferQueue==0.1.7 is required to connect a TQ client") if not bool(_state["enabled"]): - raise RuntimeError("configure_transfer_queue() must enable TQ before client connect") + raise RuntimeError( + "configure_transfer_queue() must enable TQ before client connect" + ) _ensure_initialized() @@ -223,8 +238,9 @@ def init_transfer_queue(config: Any) -> bool: Mirrors verl ``main_ppo_sync`` calling ``tq.init(config.transfer_queue)`` in the TaskRunner before workers spawn. Other Ray processes lazily call - ``tq.init()`` and connect to the named TransferQueue controller. Returns - whether TQ is usable; no-op (returns False) when disabled or not installed. + ``tq.init(config)`` with the same native configuration and connect to the + named TransferQueue controller. Returns whether TQ is usable; no-op + (returns False) when disabled or not installed. """ tq_cfg = _extract_tq_config(_drafter_training_cfg(config)) @@ -236,7 +252,10 @@ def init_transfer_queue(config: Any) -> bool: _state["enabled"] = True _state["initialized"] = True _state["owner"] = True - logger.info("[SpeCo TQ] TransferQueue bootstrapped in task runner (partition=%s)", _SPECO_TQ_PARTITION) + logger.info( + "[SpeCo TQ] TransferQueue bootstrapped in task runner (partition=%s)", + _SPECO_TQ_PARTITION, + ) return True @@ -248,11 +267,13 @@ def _drafter_training_cfg(config: Any) -> Any: def _ensure_initialized() -> None: - """Lazily ``tq.init()`` once per worker process (mirrors verl TQ_INITIALIZED). + """Lazily ``tq.init(config)`` once per worker process. - A no-argument initialization discovers the named TransferQueue controller - on the connected Ray cluster. It deliberately does not create a separate - per-worker configuration. + TransferQueue 0.1.7 first tries to discover the named controller and ignores + the supplied configuration when one already exists. Supplying the same + native configuration in every process is therefore safe for ordinary + clients and also prevents an unexpectedly early client from creating a + default-configured controller. """ if _state["initialized"]: @@ -260,7 +281,12 @@ def _ensure_initialized() -> None: with _state_lock: if _state["initialized"]: return - tq.init() + configured = _state.get("config") + if not isinstance(configured, Mapping): + raise RuntimeError( + "configure_transfer_queue() must provide TQ configuration before init" + ) + tq.init(_as_tq_config(_native_tq_config(configured))) _state["initialized"] = True @@ -302,6 +328,7 @@ def _partition_id() -> str: # Key / put / get / close # --------------------------------------------------------------------------- + def make_sample_key(global_step: Any, replica_rank: Any, request_id: Any) -> str: """Build a deterministic, cluster-unique key for one drafter sample. @@ -377,7 +404,9 @@ def list_samples() -> dict[str, dict[str, Any]]: # 0.1.7 returns key -> tag when partition_id is supplied. Accept the # partition -> (key -> tag) wrapper as well to keep the bridge version-safe. nested = result.get(_partition_id()) - if isinstance(nested, Mapping) and all(isinstance(v, Mapping) for v in nested.values()): + if isinstance(nested, Mapping) and all( + isinstance(v, Mapping) for v in nested.values() + ): result = nested records: dict[str, dict[str, Any]] = {} for key, tag in result.items(): @@ -386,7 +415,9 @@ def list_samples() -> dict[str, dict[str, Any]]: elif isinstance(tag, Mapping): records[str(key)] = dict(tag) else: - raise TypeError(f"TQ tag for key {key!r} must be a mapping, got {type(tag)!r}") + raise TypeError( + f"TQ tag for key {key!r} must be a mapping, got {type(tag)!r}" + ) return records diff --git a/verl_speco/producer/input_reader.py b/verl_speco/producer/input_reader.py index ba77fbda..6cc2e433 100644 --- a/verl_speco/producer/input_reader.py +++ b/verl_speco/producer/input_reader.py @@ -11,10 +11,11 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. -"""Streaming JSONL input and token preparation for the standalone Producer.""" +"""Streaming verl/JSONL input and token preparation for the Producer.""" from __future__ import annotations +import importlib import json import os from dataclasses import dataclass @@ -28,8 +29,8 @@ class InputRecord: sequence_no: int sample_id: str - prompt: str - response: str + prompt: str | tuple[dict[str, str], ...] + response: str | None source_metadata: dict[str, Any] @@ -43,62 +44,151 @@ class TokenizedRequest: feature_positions: torch.Tensor draft_position_ids: torch.Tensor source_metadata: dict[str, Any] + vllm_prompt_token_ids: tuple[int, ...] @property def prompt_token_ids(self) -> list[int]: - feature_end = int(self.feature_positions[-1].item()) + 1 - return self.input_ids[:feature_end].detach().cpu().long().tolist() + return list(self.vllm_prompt_token_ids) + + +@dataclass(frozen=True) +class GenerationRequest: + """One prompt-only row that needs target-model response generation.""" + + sequence_no: int + sample_id: str + prompt_token_ids: tuple[int, ...] + max_tokens: int + source_metadata: dict[str, Any] def iter_input_records(path: str | os.PathLike[str]) -> Iterator[InputRecord]: - """Yield one strict prompt/response record per non-empty JSONL line.""" + """Yield strict prompt/response records from one JSONL or Parquet file.""" input_path = Path(path) if not input_path.is_file(): - raise FileNotFoundError(f"Producer input JSONL not found: {input_path}") + raise FileNotFoundError(f"Producer input file not found: {input_path}") + sequence_no = 0 - with input_path.open("r", encoding="utf-8") as input_file: - for line_number, line in enumerate(input_file, start=1): - if not line.strip(): - continue - try: - payload = json.loads(line) - except json.JSONDecodeError as exc: - raise ValueError( - f"Invalid JSON object at {input_path}:{line_number}: {exc.msg}" - ) from exc - if not isinstance(payload, dict): - raise ValueError( - f"Producer input at {input_path}:{line_number} must be a JSON object" - ) - prompt = payload.get("prompt") - response = payload.get("response") - if not isinstance(prompt, str): + for location, payload in _iter_payloads(input_path): + if not isinstance(payload, dict): + raise ValueError( + f"Producer input at {location} must be a JSON-style object" + ) + prompt = _normalize_prompt(payload.get("prompt"), location) + response = payload.get("response") + if response is not None and (not isinstance(response, str) or not response): + raise ValueError( + f"Producer input at {location} field 'response' must be a non-empty " + "string when present" + ) + sample_id = payload.get("sample_id") or _extra_info_index(payload) + if sample_id is None: + sample_id = f"train-{sequence_no:06d}" + if not isinstance(sample_id, str) or not sample_id: + raise ValueError(f"Producer input at {location} has invalid sample_id") + source_metadata = { + key: value + for key, value in payload.items() + if key not in {"prompt", "response", "sample_id"} + } + yield InputRecord( + sequence_no=sequence_no, + sample_id=sample_id, + prompt=prompt, + response=response, + source_metadata=source_metadata, + ) + sequence_no += 1 + + +def _normalize_prompt(value: Any, location: str) -> str | tuple[dict[str, str], ...]: + if isinstance(value, str): + return value + if isinstance(value, (list, tuple)) and value: + messages: list[dict[str, str]] = [] + for index, message in enumerate(value): + if not isinstance(message, Mapping): raise ValueError( - f"Producer input at {input_path}:{line_number} requires string field 'prompt'" + f"Producer input at {location} prompt message {index} must be " + "an object" ) - if not isinstance(response, str) or not response: + role = message.get("role") + content = message.get("content") + if not isinstance(role, str) or not role: raise ValueError( - f"Producer input at {input_path}:{line_number} requires non-empty string field 'response'" + f"Producer input at {location} prompt message {index} requires " + "string field 'role'" ) - sample_id = payload.get("sample_id", f"train-{sequence_no:06d}") - if not isinstance(sample_id, str) or not sample_id: + if not isinstance(content, str): raise ValueError( - f"Producer input at {input_path}:{line_number} has invalid sample_id" + f"Producer input at {location} prompt message {index} requires " + "string field 'content'" ) - source_metadata = { - key: value - for key, value in payload.items() - if key not in {"prompt", "response", "sample_id"} - } - yield InputRecord( - sequence_no=sequence_no, - sample_id=sample_id, - prompt=prompt, - response=response, - source_metadata=source_metadata, - ) - sequence_no += 1 + messages.append({"role": role, "content": content}) + return tuple(messages) + raise ValueError( + f"Producer input at {location} requires 'prompt' as a string or " + "chat-message list" + ) + + +def _extra_info_index(payload: Mapping[str, Any]) -> str | None: + extra_info = payload.get("extra_info") + if not isinstance(extra_info, Mapping): + return None + value = extra_info.get("index") + return value if isinstance(value, str) and value else None + + +def _iter_payloads(input_path: Path) -> Iterator[tuple[str, Any]]: + if _is_parquet(input_path): + yield from _iter_parquet_payloads(input_path) + return + yield from _iter_jsonl_payloads(input_path) + + +def _is_parquet(input_path: Path) -> bool: + if input_path.suffix.lower() in {".parquet", ".pq"}: + return True + with input_path.open("rb") as input_file: + return input_file.read(4) == b"PAR1" + + +def _iter_jsonl_payloads(input_path: Path) -> Iterator[tuple[str, Any]]: + try: + with input_path.open("r", encoding="utf-8") as input_file: + for line_number, line in enumerate(input_file, start=1): + if not line.strip(): + continue + try: + payload = json.loads(line) + except json.JSONDecodeError as exc: + raise ValueError( + f"Invalid JSON object at {input_path}:{line_number}: {exc.msg}" + ) from exc + yield f"{input_path}:{line_number}", payload + except UnicodeDecodeError as exc: + raise ValueError( + f"Producer input {input_path} is not UTF-8 JSONL or a Parquet file" + ) from exc + + +def _iter_parquet_payloads(input_path: Path) -> Iterator[tuple[str, Any]]: + try: + parquet = importlib.import_module("pyarrow.parquet") + except ImportError as exc: + raise RuntimeError( + "Reading a Parquet training file requires pyarrow; install the normal " + "verl data dependencies in the training environment" + ) from exc + + parquet_file = parquet.ParquetFile(input_path) + row_number = 0 + for batch in parquet_file.iter_batches(): + for payload in batch.to_pylist(): + row_number += 1 + yield f"{input_path}:row {row_number}", payload def build_loss_mask(input_ids: torch.Tensor, prompt_length: int) -> torch.Tensor: @@ -117,38 +207,176 @@ def tokenize_record( tokenizer: Any, config: Mapping[str, Any] | Any, ) -> TokenizedRequest: - """Tokenize existing prompt/response text without generating new tokens.""" + """Tokenize one row that already contains a response.""" - prompt_ids = _token_ids(tokenizer(record.prompt, add_special_tokens=False)) - full_ids = _token_ids( - tokenizer(record.prompt + record.response, add_special_tokens=False) - ) + if record.response is None: + raise ValueError( + f"Producer sample {record.sample_id!r} has no response; prepare it for " + "target-model generation instead" + ) + prompt_ids = _prompt_ids(record.prompt, tokenizer) + if isinstance(record.prompt, str): + full_ids = _token_ids( + tokenizer(record.prompt + record.response, add_special_tokens=False) + ) + else: + full_ids = _token_ids( + tokenizer.apply_chat_template( + [*record.prompt, {"role": "assistant", "content": record.response}], + tokenize=True, + add_generation_prompt=False, + ) + ) if full_ids[: len(prompt_ids)] != prompt_ids: raise ValueError( f"Producer sample {record.sample_id!r} has an unstable tokenizer boundary " - "between prompt and response; prompt token IDs are not a prefix of full token IDs" + "between prompt and response; prompt token IDs are not a prefix of full " + "token IDs" ) if len(full_ids) <= len(prompt_ids): raise ValueError( f"Producer sample {record.sample_id!r} produced no response tokens" ) - input_ids = torch.tensor(full_ids, dtype=torch.int64) - if int(input_ids.numel()) <= 0: + return _build_tokenized_request( + sequence_no=record.sequence_no, + sample_id=record.sample_id, + prompt_length=len(prompt_ids), + full_ids=full_ids, + source_metadata=record.source_metadata, + config=config, + ) + + +def prepare_generation_request( + record: InputRecord, + tokenizer: Any, + config: Mapping[str, Any] | Any, +) -> GenerationRequest: + """Tokenize a prompt-only row and bound target-model generation.""" + + if record.response is not None: + raise ValueError( + f"Producer sample {record.sample_id!r} already contains a response" + ) + prompt_ids = _prompt_ids(record.prompt, tokenizer) + max_sequence_length = int(_config_value(config, "max_sequence_length", 0) or 0) + max_tokens = int(_config_value(config, "generation_max_tokens", 0) or 0) + if max_tokens <= 0: + raise ValueError("generation_max_tokens must be positive") + if max_sequence_length > 0: + max_tokens = min(max_tokens, max_sequence_length - len(prompt_ids)) + if max_tokens <= 0: + raise ValueError( + f"Producer sample {record.sample_id!r} prompt has {len(prompt_ids)} tokens " + "and leaves no generation capacity within " + f"max_sequence_length={max_sequence_length}" + ) + return GenerationRequest( + sequence_no=record.sequence_no, + sample_id=record.sample_id, + prompt_token_ids=tuple(prompt_ids), + max_tokens=max_tokens, + source_metadata=dict(record.source_metadata), + ) + + +def finalize_generated_request( + request: GenerationRequest, + hidden_state_token_ids: Any, + config: Mapping[str, Any] | Any, + *, + expected_response_token_ids: Any | None = None, +) -> TokenizedRequest: + """Build a training request from vLLM generation and hidden-state tokens.""" + + hidden_ids = _token_ids(hidden_state_token_ids) + prompt_ids = list(request.prompt_token_ids) + if hidden_ids[: len(prompt_ids)] != prompt_ids: + raise ValueError( + f"vLLM hidden-state token sequence for sample {request.sample_id!r} " + "does not " + "start with the rendered prompt token IDs" + ) + if expected_response_token_ids is not None: + response_ids = _token_ids(expected_response_token_ids) + full_ids = [*prompt_ids, *response_ids] + # ExampleHiddenStatesConnector deliberately excludes the final sampled + # token: that token was emitted by the model but was never fed through a + # subsequent forward pass, so no hidden state exists for it. + expected_hidden_ids = full_ids[:-1] + if hidden_ids != expected_hidden_ids: + raise ValueError( + f"vLLM hidden-state token IDs for sample {request.sample_id!r} do not " + "match the prompt plus generated completion excluding its final token " + f"(hidden={len(hidden_ids)}, expected={len(expected_hidden_ids)}, " + f"completion={len(response_ids)})" + ) + else: + # Prefilled records already contain their complete response, and their + # caller supplies the full token sequence directly. + full_ids = hidden_ids + if len(full_ids) <= len(prompt_ids): raise ValueError( - f"Producer sample {record.sample_id!r} produced no input tokens" + f"vLLM generated no response tokens for sample {request.sample_id!r}; " + "the hidden-state server must enable include_output_tokens" ) + return _build_tokenized_request( + sequence_no=request.sequence_no, + sample_id=request.sample_id, + prompt_length=len(prompt_ids), + full_ids=full_ids, + source_metadata=request.source_metadata, + config=config, + vllm_prompt_token_ids=hidden_ids, + feature_end_limit=len(hidden_ids), + ) + + +def _prompt_ids(prompt: str | tuple[dict[str, str], ...], tokenizer: Any) -> list[int]: + if isinstance(prompt, str): + return _token_ids(tokenizer(prompt, add_special_tokens=False)) + apply_chat_template = getattr(tokenizer, "apply_chat_template", None) + if not callable(apply_chat_template): + raise RuntimeError( + "Chat-message prompts require a tokenizer with apply_chat_template()" + ) + return _token_ids( + apply_chat_template( + list(prompt), + tokenize=True, + add_generation_prompt=True, + ) + ) + + +def _build_tokenized_request( + *, + sequence_no: int, + sample_id: str, + prompt_length: int, + full_ids: list[int], + source_metadata: Mapping[str, Any], + config: Mapping[str, Any] | Any, + vllm_prompt_token_ids: list[int] | None = None, + feature_end_limit: int | None = None, +) -> TokenizedRequest: + input_ids = torch.tensor(full_ids, dtype=torch.int64) + if int(input_ids.numel()) <= 0: + raise ValueError(f"Producer sample {sample_id!r} produced no input tokens") max_sequence_length = int(_config_value(config, "max_sequence_length", 0) or 0) if max_sequence_length > 0 and int(input_ids.numel()) > max_sequence_length: raise ValueError( - f"Producer sample {record.sample_id!r} has {int(input_ids.numel())} tokens, " + f"Producer sample {sample_id!r} has {int(input_ids.numel())} tokens, " f"exceeding max_sequence_length={max_sequence_length}" ) - loss_mask = build_loss_mask(input_ids, len(prompt_ids)) + loss_mask = build_loss_mask(input_ids, prompt_length) position_ids = torch.arange(int(input_ids.numel()), dtype=torch.int64) - feature_start = max(len(prompt_ids) - 1, 0) + feature_start = max(prompt_length - 1, 0) feature_end = int(input_ids.numel()) + if feature_end_limit is not None: + feature_end = min(feature_end, int(feature_end_limit)) max_feature_length = int(_config_value(config, "max_feature_length", 0) or 0) if max_feature_length == 1: raise ValueError("max_feature_length must be 0 or at least 2") @@ -156,28 +384,34 @@ def tokenize_record( feature_end = min(feature_start + max_feature_length, feature_end) feature_positions = torch.arange(feature_start, feature_end, dtype=torch.int64) if int(feature_positions.numel()) <= 0: - raise ValueError( - f"Producer sample {record.sample_id!r} has an empty feature window" - ) + raise ValueError(f"Producer sample {sample_id!r} has an empty feature window") draft_position_ids = position_ids[feature_start:feature_end] + 1 return TokenizedRequest( - sequence_no=record.sequence_no, - sample_id=record.sample_id, + sequence_no=sequence_no, + sample_id=sample_id, input_ids=input_ids, loss_mask=loss_mask, position_ids=position_ids, feature_positions=feature_positions, draft_position_ids=draft_position_ids, - source_metadata=dict(record.source_metadata), + source_metadata=dict(source_metadata), + vllm_prompt_token_ids=tuple( + vllm_prompt_token_ids + if vllm_prompt_token_ids is not None + else full_ids[:feature_end] + ), ) def _token_ids(encoding: Any) -> list[int]: - value = ( - encoding.get("input_ids") - if isinstance(encoding, Mapping) - else encoding.input_ids - ) + if isinstance(encoding, (list, tuple)) or torch.is_tensor(encoding): + value = encoding + else: + value = ( + encoding.get("input_ids") + if isinstance(encoding, Mapping) + else encoding.input_ids + ) if value is None: raise ValueError("Tokenizer result is missing input_ids") if torch.is_tensor(value): @@ -198,9 +432,12 @@ def _config_value(config: Any, key: str, default: Any = None) -> Any: __all__ = [ + "GenerationRequest", "InputRecord", "TokenizedRequest", "build_loss_mask", + "finalize_generated_request", "iter_input_records", + "prepare_generation_request", "tokenize_record", ] diff --git a/verl_speco/producer/vllm_feature_client.py b/verl_speco/producer/vllm_feature_client.py index 4b793a66..b57961dc 100644 --- a/verl_speco/producer/vllm_feature_client.py +++ b/verl_speco/producer/vllm_feature_client.py @@ -17,6 +17,7 @@ import asyncio import errno +import importlib import inspect import os import time @@ -41,6 +42,7 @@ def __post_init__(self) -> None: class VllmResponse: hidden_states_path: str endpoint_url: str + generated_token_ids: tuple[int, ...] = () @dataclass(frozen=True) @@ -49,6 +51,7 @@ class RawVllmFeature: temporary_path: str endpoint_url: str byte_size: int + generated_token_ids: tuple[int, ...] = () @dataclass @@ -89,6 +92,51 @@ async def request_prefill( return VllmResponse(os.fspath(path), endpoint.base_url) +async def request_generate( + endpoint: VllmEndpoint, + client: Any, + prompt_token_ids: list[int], + *, + model: str, + max_tokens: int, + timeout: float, +) -> VllmResponse: + """Generate a response and request hidden states for prompt and output tokens.""" + + response = await client.completions.create( + model=model, + prompt=prompt_token_ids, + max_tokens=max_tokens, + extra_body={ + "return_token_ids": True, + "kv_transfer_params": {"include_output_tokens": True}, + }, + timeout=timeout, + ) + choices = getattr(response, "choices", None) or [] + if not choices: + raise ValueError("vLLM generation response has no choices") + actual_prompt = getattr(choices[0], "prompt_token_ids", None) + if actual_prompt is not None and list(actual_prompt) != prompt_token_ids: + raise ValueError("vLLM generation prompt_token_ids mismatch") + generated = getattr(choices[0], "token_ids", None) + if not isinstance(generated, (list, tuple)) or not generated: + raise ValueError( + "vLLM generation response missing token_ids; enable return_token_ids support" + ) + params = getattr(response, "kv_transfer_params", None) + if not isinstance(params, Mapping): + raise ValueError("vLLM generation response missing kv_transfer_params") + path = params.get("hidden_states_path") + if not path: + raise ValueError("vLLM generation response missing hidden_states_path") + return VllmResponse( + os.fspath(path), + endpoint.base_url, + tuple(int(token_id) for token_id in generated), + ) + + def load_hidden_state_result(response: VllmResponse) -> RawVllmFeature: try: from safetensors.torch import load_file @@ -103,6 +151,7 @@ def load_hidden_state_result(response: VllmResponse) -> RawVllmFeature: temporary_path=str(path), endpoint_url=response.endpoint_url, byte_size=int(path.stat().st_size), + generated_token_ids=response.generated_token_ids, ) @@ -160,19 +209,35 @@ async def start(self) -> None: ] async def prefill(self, request: Any) -> RawVllmFeature: + return await self._request(request, generate=False) + + async def generate(self, request: Any) -> RawVllmFeature: + return await self._request(request, generate=True) + + async def _request(self, request: Any, *, generate: bool) -> RawVllmFeature: if not self._states: raise RuntimeError("VllmFeatureClientPool.start() must be called first") state = choose_endpoint(self._states) state.inflight += 1 try: async with self._global_semaphore, state.semaphore: - response = await request_prefill( - state.endpoint, - state.client, - list(request.prompt_token_ids), - model=self.model, - timeout=self.request_timeout, - ) + if generate: + response = await request_generate( + state.endpoint, + state.client, + list(request.prompt_token_ids), + model=self.model, + max_tokens=int(request.max_tokens), + timeout=self.request_timeout, + ) + else: + response = await request_prefill( + state.endpoint, + state.client, + list(request.prompt_token_ids), + model=self.model, + timeout=self.request_timeout, + ) raw = await asyncio.to_thread(load_hidden_state_result, response) state.requests += 1 return raw @@ -194,7 +259,7 @@ def _wait_for_lock(lock_path: Path, timeout: float = 30.0) -> None: if not lock_path.exists(): return try: - import fcntl + fcntl: Any = importlib.import_module("fcntl") except ImportError: # vLLM's file connector is Linux-only. Keep the old existence-based # fallback for dependency-light tests on other platforms. @@ -233,4 +298,5 @@ def _wait_for_lock(lock_path: Path, timeout: float = 30.0) -> None: "delete_temporary_result", "load_hidden_state_result", "request_prefill", + "request_generate", ] diff --git a/verl_speco/standalone_tq_producer.py b/verl_speco/standalone_tq_producer.py index 395c8bc2..71e5e9d6 100644 --- a/verl_speco/standalone_tq_producer.py +++ b/verl_speco/standalone_tq_producer.py @@ -27,8 +27,11 @@ resolve_drafter_hidden_states_layout, ) from verl_speco.producer.input_reader import ( + GenerationRequest, TokenizedRequest, + finalize_generated_request, iter_input_records, + prepare_generation_request, tokenize_record, ) from verl_speco.producer.vllm_feature_client import ( @@ -135,6 +138,7 @@ def validate_producer_config(config: Any) -> None: "input_queue_size", "publish_queue_size", "max_pending_samples", + "generation_max_tokens", ) invalid = [name for name in positive_fields if int(producer_cfg.get(name, 0)) <= 0] if invalid: @@ -210,7 +214,11 @@ async def run_producer( async def read_inputs() -> None: for record in iter_input_records(str(producer_cfg["input_path"])): - request = tokenize_record(record, tokenizer, producer_cfg) + request = ( + prepare_generation_request(record, tokenizer, producer_cfg) + if record.response is None + else tokenize_record(record, tokenizer, producer_cfg) + ) await input_queue.put(request) stats.input_count += 1 for _ in range(worker_count): @@ -228,7 +236,16 @@ async def request_worker() -> None: max_pending_samples=int(producer_cfg["max_pending_samples"]), poll_interval=float(producer_cfg["pending_poll_interval_seconds"]), ) - raw = await pool.prefill(request) + if isinstance(request, GenerationRequest): + raw = await pool.generate(request) + request = finalize_generated_request( + request, + raw.payload.get("token_ids"), + producer_cfg, + expected_response_token_ids=raw.generated_token_ids, + ) + else: + raw = await pool.prefill(request) stats.pending_bytes += int(raw.byte_size) sample = feature_from_vllm_payload(raw, request, feature_contract) await publish_queue.put( @@ -279,10 +296,12 @@ async def publish_results() -> None: ) return stats finally: - if pool is not None: - await pool.close() - if connected: - transport.close_transfer_queue_client() + try: + if pool is not None: + await pool.close() + finally: + if connected: + transport.close_transfer_queue_client() async def _wait_for_owner_ready( diff --git a/verl_speco/standalone_tq_training_launcher.py b/verl_speco/standalone_tq_training_launcher.py new file mode 100644 index 00000000..f0a5e403 --- /dev/null +++ b/verl_speco/standalone_tq_training_launcher.py @@ -0,0 +1,638 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Single-entry launcher for Producer -> TransferQueue -> draft training. + +The example script keeps the ordinary standalone-training interface. This +module owns the internal Ray/TQ identity and the owner, Producer and Consumer +process lifecycle so transport-specific overrides do not leak into examples. +""" + +from __future__ import annotations + +import argparse +import hashlib +import importlib +import json +import logging +import os +from collections.abc import Callable, Mapping, Sequence +from dataclasses import dataclass +from pathlib import Path +import subprocess +import sys +import tempfile +import time +from typing import Any +from urllib.error import HTTPError, URLError +from urllib.parse import urlparse +from urllib.request import urlopen +import uuid + + +logger = logging.getLogger(__name__) + +_MODEL_PATH_KEY = "actor_rollout_ref.model.path" +_DRAFTER_PATH_KEY = "actor_rollout_ref.rollout.drafter.model_path" +_ALGORITHM_KEY = "actor_rollout_ref.rollout.drafter.speculative_algorithm" +_TRAIN_FILES_KEY = "data.train_files" +_TOKENIZER_PATH_KEY = ( + "actor_rollout_ref.rollout.drafter.training.feature_store.tokenizer_path" +) +_DSPARK_LAYER_IDS_KEY = ( + "actor_rollout_ref.rollout.drafter.training.dspark_target_layer_ids" +) + +_TQ_PREFIX = "actor_rollout_ref.rollout.drafter.training.transfer_queue" +_FEATURE_STORE_PREFIX = "actor_rollout_ref.rollout.drafter.training.feature_store" +_INTERNAL_OVERRIDE_KEYS = frozenset( + { + f"{_FEATURE_STORE_PREFIX}.type", + f"{_FEATURE_STORE_PREFIX}.path", + f"{_FEATURE_STORE_PREFIX}.shuffle", + f"{_FEATURE_STORE_PREFIX}.repeat", + f"{_TQ_PREFIX}.enable", + f"{_TQ_PREFIX}.ray.address", + f"{_TQ_PREFIX}.ray.namespace", + f"{_TQ_PREFIX}.partition_id", + f"{_TQ_PREFIX}.run_id", + f"{_TQ_PREFIX}.drop_last", + f"{_TQ_PREFIX}.backend.storage_backend", + f"{_TQ_PREFIX}.backend.SimpleStorage.total_storage_size", + f"{_TQ_PREFIX}.backend.SimpleStorage.num_data_storage_units", + } +) + +_DEFAULT_TARGET_LAYER_IDS = (1, 9, 17, 25, 33) +_DEFAULT_VLLM_ENDPOINT = "http://127.0.0.1:8000/v1" +_DEFAULT_VLLM_GPU_MEMORY_UTILIZATION = "0.4" +_VLLM_HIDDEN_STATES_DIR = "__SPECO_HIDDEN_STATES_DIR__" +_TQ_NAMESPACE = "speco-drafter" +_TQ_PARTITION = "speco_drafter_features" + + +@dataclass(frozen=True) +class PipelineConfig: + input_path: str + model_path: str + tokenizer_path: str + algorithm: str + target_layer_ids: tuple[int, ...] + vllm_endpoint: str + run_id: str + + +@dataclass(frozen=True) +class PipelineCommands: + vllm: list[str] | None + vllm_endpoint: str + owner: list[str] + producer: list[str] + consumer: list[str] + + +@dataclass(frozen=True) +class RaySession: + module: Any + address: str + + def close(self) -> None: + self.module.shutdown() + + +def _split_override(item: str) -> tuple[str, str] | None: + if "=" not in item or item.startswith("-"): + return None + key, value = item.split("=", 1) + return key, value + + +def _find_override(overrides: Sequence[str], key: str) -> str | None: + for item in reversed(overrides): + parsed = _split_override(item) + if parsed is not None and parsed[0] == key: + return parsed[1] + return None + + +def _strip_quotes(value: str) -> str: + value = value.strip() + if len(value) >= 2 and value[0] == value[-1] and value[0] in {"'", '"'}: + return value[1:-1] + return value + + +def _single_train_file(value: str | None) -> str: + if value is None: + raise ValueError(f"Standalone TQ training requires {_TRAIN_FILES_KEY}") + text = _strip_quotes(value) + if text.startswith("[") and text.endswith("]"): + items = [_strip_quotes(item) for item in text[1:-1].split(",") if item.strip()] + if len(items) != 1: + raise ValueError("Standalone TQ Producer requires exactly one train file") + text = items[0] + if not text: + raise ValueError("Standalone TQ Producer train file must not be empty") + return text + + +def _parse_layer_ids(value: str | None) -> tuple[int, ...]: + if value is None or _strip_quotes(value).lower() in {"", "null", "none"}: + return _DEFAULT_TARGET_LAYER_IDS + text = _strip_quotes(value) + if not (text.startswith("[") and text.endswith("]")): + raise ValueError(f"{_DSPARK_LAYER_IDS_KEY} must be a Hydra integer list") + try: + result = tuple(int(item.strip()) for item in text[1:-1].split(",")) + except ValueError as exc: + raise ValueError(f"{_DSPARK_LAYER_IDS_KEY} must contain only integers") from exc + if not result or any(value < 0 for value in result): + raise ValueError(f"{_DSPARK_LAYER_IDS_KEY} must contain non-negative IDs") + return result + + +def _required_override(overrides: Sequence[str], key: str) -> str: + value = _find_override(overrides, key) + normalized = _strip_quotes(value or "") + if not normalized or normalized.startswith("/path/to/"): + raise ValueError(f"Standalone TQ training requires a real {key}") + return normalized + + +def _stable_path_identity(kind: str, path: str) -> str: + digest = hashlib.sha256(path.encode("utf-8")).hexdigest() + return f"{kind}-path-sha256-{digest}" + + +def _target_final_layer_id(model_path: str, target_layer_ids: Sequence[int]) -> int: + """Resolve the final transformer-layer output ID from a local HF config.""" + + config_path = Path(model_path) / "config.json" + if config_path.is_file(): + try: + model_config = json.loads(config_path.read_text(encoding="utf-8")) + except (OSError, UnicodeDecodeError, json.JSONDecodeError) as exc: + raise ValueError(f"Cannot read target model config: {config_path}") from exc + for candidate in ( + model_config.get("num_hidden_layers"), + (model_config.get("text_config") or {}).get("num_hidden_layers"), + ): + if candidate is not None and int(candidate) > 0: + return int(candidate) + raise ValueError(f"Target model config has no num_hidden_layers: {config_path}") + # Keep dry-run and model-registry IDs usable. The formal Qwen3-4B/8B + # defaults select layer 33 and use transformer output 36 as the final state. + return max(int(layer_id) for layer_id in target_layer_ids) + 3 + + +def resolve_pipeline_config( + training_args: Sequence[str], + *, + environ: Mapping[str, str] | None = None, +) -> PipelineConfig: + """Derive all Producer/TQ settings from ordinary training arguments.""" + + env = os.environ if environ is None else environ + model_path = _required_override(training_args, _MODEL_PATH_KEY) + _required_override(training_args, _DRAFTER_PATH_KEY) + input_path = _single_train_file(_find_override(training_args, _TRAIN_FILES_KEY)) + tokenizer_path = _strip_quotes( + _find_override(training_args, _TOKENIZER_PATH_KEY) or model_path + ) + algorithm = _strip_quotes( + _find_override(training_args, _ALGORITHM_KEY) or "DSPARK" + ).upper() + if algorithm != "DSPARK": + raise ValueError("Standalone Producer/TQ/Consumer launcher requires DSPARK") + target_layer_ids = _parse_layer_ids( + _find_override(training_args, _DSPARK_LAYER_IDS_KEY) + ) + endpoint = str(env.get("SPECO_VLLM_ENDPOINT", _DEFAULT_VLLM_ENDPOINT)).strip() + if not endpoint: + raise ValueError("SPECO_VLLM_ENDPOINT must not be empty") + return PipelineConfig( + input_path=input_path, + model_path=model_path, + tokenizer_path=tokenizer_path, + algorithm=algorithm, + target_layer_ids=target_layer_ids, + vllm_endpoint=endpoint.rstrip("/"), + run_id=f"dspark-{uuid.uuid4().hex}", + ) + + +def start_ray_session( + *, + environ: Mapping[str, str] | None = None, + ray_module: Any | None = None, +) -> RaySession: + """Create a task-local Ray control plane for the hidden TQ pipeline.""" + + env = os.environ if environ is None else environ + ray_runtime = ray_module + if ray_runtime is None: + try: + ray_runtime = importlib.import_module("ray") + except ImportError as exc: + raise RuntimeError( + "Standalone TQ training requires Ray and TransferQueue==0.1.7" + ) from exc + # ``ray.init()`` consults RAY_ADDRESS when no explicit address is supplied. + # This launcher owns the complete Producer/TQ/Consumer lifetime, so an + # inherited address (often left by another job) must never select its + # control plane. ``local`` explicitly starts this task's Ray runtime. + init_kwargs: dict[str, Any] = { + "address": "local", + "namespace": _TQ_NAMESPACE, + "include_dashboard": False, + } + num_cpus = str(env.get("SPECO_RAY_NUM_CPUS", "")).strip() + if num_cpus: + init_kwargs["num_cpus"] = int(num_cpus) + ray_runtime.init(**init_kwargs) + address = str(ray_runtime.get_runtime_context().gcs_address).strip() + if not address: + ray_runtime.shutdown() + raise RuntimeError("Ray did not report a GCS address for TQ clients") + return RaySession(module=ray_runtime, address=address) + + +def _hydra_list(values: Sequence[Any]) -> str: + return "[" + ",".join(str(value) for value in values) + "]" + + +def _replace_internal_overrides( + training_args: Sequence[str], internal: Sequence[str] +) -> list[str]: + cleaned: list[str] = [] + for item in training_args: + parsed = _split_override(item) + if parsed is not None and parsed[0] in _INTERNAL_OVERRIDE_KEYS: + continue + cleaned.append(item) + return [*cleaned, *internal] + + +def build_pipeline_commands( + config: PipelineConfig, + training_args: Sequence[str], + *, + ray_address: str, + python_executable: str = sys.executable, +) -> PipelineCommands: + """Build the internal commands without exposing transport options.""" + + tq_overrides = [ + f"{_TQ_PREFIX}.enable=true", + f"{_TQ_PREFIX}.ray.address={ray_address}", + f"{_TQ_PREFIX}.ray.namespace={_TQ_NAMESPACE}", + f"{_TQ_PREFIX}.partition_id={_TQ_PARTITION}", + f"{_TQ_PREFIX}.run_id={config.run_id}", + f"{_TQ_PREFIX}.drop_last=true", + f"{_TQ_PREFIX}.backend.storage_backend=SimpleStorage", + f"{_TQ_PREFIX}.backend.SimpleStorage.total_storage_size=17179869184", + f"{_TQ_PREFIX}.backend.SimpleStorage.num_data_storage_units=8", + ] + parsed_endpoint = urlparse(config.vllm_endpoint) + if parsed_endpoint.scheme not in {"http", "https"} or not parsed_endpoint.hostname: + raise ValueError(f"Invalid vLLM endpoint: {config.vllm_endpoint!r}") + vllm_port = parsed_endpoint.port or ( + 443 if parsed_endpoint.scheme == "https" else 80 + ) + # extract_hidden_states uses the model's layer-output convention. Qwen3-4B/8B + # have 36 transformer layers; the default DSpark auxiliary selection ends at + # 33 and requests the final layer output as 36. + final_layer_id = _target_final_layer_id(config.model_path, config.target_layer_ids) + speculative_config = { + "method": "extract_hidden_states", + "num_speculative_tokens": 1, + "draft_model_config": { + "hf_config": { + "eagle_aux_hidden_state_layer_ids": [ + *config.target_layer_ids, + final_layer_id, + ] + } + }, + } + kv_transfer_config = { + "kv_connector": "ExampleHiddenStatesConnector", + "kv_role": "kv_producer", + "kv_connector_extra_config": { + "shared_storage_path": _VLLM_HIDDEN_STATES_DIR, + "use_synchronization_lock": True, + }, + } + vllm = None + if parsed_endpoint.hostname in {"127.0.0.1", "localhost", "0.0.0.0"}: + vllm = [ + "vllm", + "serve", + config.model_path, + "--host", + "127.0.0.1", + "--port", + str(vllm_port), + "--gpu-memory-utilization", + _DEFAULT_VLLM_GPU_MEMORY_UTILIZATION, + "--speculative-config", + json.dumps(speculative_config, separators=(",", ":")), + "--kv-transfer-config", + json.dumps(kv_transfer_config, separators=(",", ":")), + "--no-enable-chunked-prefill", + ] + owner = [ + python_executable, + "-m", + "verl_speco.tq_owner", + *tq_overrides, + ] + producer = [ + python_executable, + "-m", + "verl_speco.standalone_tq_producer", + f"{_ALGORITHM_KEY}={config.algorithm}", + *tq_overrides, + f"speco.standalone_tq_producer.input_path={config.input_path}", + f"speco.standalone_tq_producer.tokenizer_path={config.tokenizer_path}", + "speco.standalone_tq_producer.tokenizer_fingerprint=" + + _stable_path_identity("tokenizer", config.tokenizer_path), + f"speco.standalone_tq_producer.target_model_id={config.model_path}", + "speco.standalone_tq_producer.target_model_revision=" + + _stable_path_identity("target", config.model_path), + "speco.standalone_tq_producer.target_layer_ids=" + + _hydra_list(config.target_layer_ids), + "speco.standalone_tq_producer.vllm_endpoints=" + + _hydra_list((config.vllm_endpoint,)), + f"speco.standalone_tq_producer.vllm_model={config.model_path}", + ] + consumer_internal = [ + f"{_FEATURE_STORE_PREFIX}.type=tq", + f"{_FEATURE_STORE_PREFIX}.path=null", + f"{_FEATURE_STORE_PREFIX}.shuffle=false", + f"{_FEATURE_STORE_PREFIX}.repeat=false", + *tq_overrides, + f"{_DSPARK_LAYER_IDS_KEY}={_hydra_list(config.target_layer_ids)}", + ] + consumer = [ + python_executable, + "-m", + "verl_speco.draft_train_launcher", + *_replace_internal_overrides(training_args, consumer_internal), + ] + return PipelineCommands( + vllm=vllm, + vllm_endpoint=config.vllm_endpoint, + owner=owner, + producer=producer, + consumer=consumer, + ) + + +def _wait_for_owner_ready( + owner: subprocess.Popen[Any], + ready_file: Path, + *, + timeout_seconds: float, + monotonic: Callable[[], float] = time.monotonic, + sleep: Callable[[float], None] = time.sleep, +) -> None: + deadline = monotonic() + timeout_seconds + while not ready_file.is_file(): + returncode = owner.poll() + if returncode is not None: + raise RuntimeError(f"TransferQueue owner exited early ({returncode})") + if monotonic() >= deadline: + raise TimeoutError("Timed out waiting for TransferQueue owner readiness") + sleep(0.1) + + +def _vllm_is_ready(endpoint: str, *, timeout_seconds: float = 1.0) -> bool: + try: + with urlopen( + f"{endpoint.rstrip('/')}/models", timeout=timeout_seconds + ) as response: + return 200 <= int(response.status) < 300 + except (HTTPError, URLError, OSError, TimeoutError): + return False + + +def _wait_for_vllm_ready( + process: subprocess.Popen[Any], + endpoint: str, + *, + timeout_seconds: float, + endpoint_ready: Callable[[str], bool] = _vllm_is_ready, + monotonic: Callable[[], float] = time.monotonic, + sleep: Callable[[float], None] = time.sleep, +) -> None: + deadline = monotonic() + timeout_seconds + while not endpoint_ready(endpoint): + returncode = process.poll() + if returncode is not None: + raise RuntimeError( + f"hidden-state vLLM exited before becoming ready ({returncode})" + ) + if monotonic() >= deadline: + raise TimeoutError(f"Timed out waiting for hidden-state vLLM at {endpoint}") + sleep(1.0) + + +def _stop_process(process: subprocess.Popen[Any] | None) -> None: + if process is None or process.poll() is not None: + return + process.terminate() + try: + process.wait(timeout=10) + except subprocess.TimeoutExpired: + process.kill() + process.wait() + + +def run_pipeline( + commands: PipelineCommands, + *, + ray_address: str, + environ: Mapping[str, str] | None = None, + popen: Callable[..., subprocess.Popen[Any]] = subprocess.Popen, + poll_interval_seconds: float = 0.2, + owner_ready_timeout_seconds: float = 120, + vllm_ready_timeout_seconds: float = 900, + endpoint_ready: Callable[[str], bool] = _vllm_is_ready, +) -> int: + """Run the internal processes and return the training Consumer status.""" + + base_env = dict(os.environ if environ is None else environ) + # Ray itself also reads RAY_ADDRESS. Pin every child (including the + # torchrun ranks created by the Consumer launcher) to the control plane + # created above instead of allowing a stale inherited value to win. + base_env["RAY_ADDRESS"] = ray_address + owner: subprocess.Popen[Any] | None = None + producer: subprocess.Popen[Any] | None = None + consumer: subprocess.Popen[Any] | None = None + vllm: subprocess.Popen[Any] | None = None + hidden_states_temp: tempfile.TemporaryDirectory[str] | None = None + try: + with tempfile.TemporaryDirectory(prefix="speco-tq-launch-") as temp_dir: + ready_file = Path(temp_dir) / "owner.ready" + hidden_states_temp = tempfile.TemporaryDirectory( + prefix="speco-vllm-hidden-states-" + ) + hidden_states_dir = Path(hidden_states_temp.name) + config_endpoint = commands.vllm_endpoint + if not endpoint_ready(config_endpoint): + if commands.vllm is None: + raise RuntimeError( + "The configured remote hidden-state vLLM is unavailable at " + f"{config_endpoint}" + ) + vllm_command = [ + part.replace(_VLLM_HIDDEN_STATES_DIR, str(hidden_states_dir)) + for part in commands.vllm + ] + logger.info("Starting hidden-state vLLM at %s", config_endpoint) + vllm = popen(vllm_command, env=base_env) + owner_env = {**base_env, "SPECO_TQ_OWNER_READY_FILE": str(ready_file)} + logger.info("Starting TransferQueue owner") + owner = popen(commands.owner, env=owner_env) + _wait_for_owner_ready( + owner, + ready_file, + timeout_seconds=owner_ready_timeout_seconds, + ) + if vllm is not None: + _wait_for_vllm_ready( + vllm, + config_endpoint, + timeout_seconds=vllm_ready_timeout_seconds, + endpoint_ready=endpoint_ready, + ) + elif not endpoint_ready(config_endpoint): + raise RuntimeError( + f"hidden-state vLLM became unavailable at {config_endpoint}" + ) + logger.info("Starting standalone DSpark Consumer") + consumer = popen(commands.consumer, env=base_env) + logger.info("Starting standalone vLLM Producer") + producer = popen(commands.producer, env=base_env) + + while True: + owner_status = owner.poll() + vllm_status = None if vllm is None else vllm.poll() + producer_status = producer.poll() + consumer_status = consumer.poll() + if owner_status is not None: + raise RuntimeError( + f"TransferQueue owner exited during training ({owner_status})" + ) + if producer_status is not None and producer_status != 0: + return int(producer_status) + if vllm_status is not None: + raise RuntimeError( + f"hidden-state vLLM exited during training ({vllm_status})" + ) + if consumer_status is not None: + return int(consumer_status) + time.sleep(poll_interval_seconds) + except KeyboardInterrupt: + logger.warning("Standalone DSpark training interrupted") + return 130 + finally: + _stop_process(producer) + _stop_process(consumer) + _stop_process(owner) + _stop_process(vllm) + if hidden_states_temp is not None: + hidden_states_temp.cleanup() + + +def _format_command(command: Sequence[str]) -> str: + import shlex + + return " ".join(shlex.quote(part) for part in command) + + +def _preflight_input_file(input_path: str) -> None: + path = Path(input_path) + if not path.is_file(): + raise ValueError(f"Training file does not exist: {input_path}") + from verl_speco.producer.input_reader import iter_input_records + + try: + next(iter_input_records(path)) + except StopIteration as exc: + raise ValueError(f"Training file contains no samples: {input_path}") from exc + + +def main(argv: Sequence[str] | None = None) -> int: + parser = argparse.ArgumentParser( + description="Launch standalone DSpark training through Producer/TQ/Consumer." + ) + parser.add_argument("--dry-run", action="store_true") + parser.add_argument("--python-executable", default=sys.executable) + args, training_args = parser.parse_known_args(argv) + logging.basicConfig(level=logging.INFO) + + try: + config = resolve_pipeline_config(training_args) + if args.dry_run: + commands = build_pipeline_commands( + config, + training_args, + ray_address="127.0.0.1:6379", + python_executable=args.python_executable, + ) + printable_commands = [ + ("vllm", commands.vllm), + ("owner", commands.owner), + ("producer", commands.producer), + ("consumer", commands.consumer), + ] + for role, command in printable_commands: + if command is None: + print(f"{role}: external service at {commands.vllm_endpoint}") + continue + print(f"{role}: {_format_command(command)}") + return 0 + _preflight_input_file(config.input_path) + ray_session = start_ray_session() + try: + commands = build_pipeline_commands( + config, + training_args, + ray_address=ray_session.address, + python_executable=args.python_executable, + ) + logger.info("Using task-local Ray control plane at %s", ray_session.address) + return run_pipeline(commands, ray_address=ray_session.address) + finally: + ray_session.close() + except (OSError, RuntimeError, TimeoutError, ValueError) as exc: + logger.error("Standalone TQ training failed: %s", exc) + return 2 + + +if __name__ == "__main__": + raise SystemExit(main()) + + +__all__ = [ + "PipelineCommands", + "PipelineConfig", + "RaySession", + "build_pipeline_commands", + "main", + "resolve_pipeline_config", + "run_pipeline", + "start_ray_session", +] diff --git a/verl_speco/tq_owner.py b/verl_speco/tq_owner.py index 0d8d1400..cae5552b 100644 --- a/verl_speco/tq_owner.py +++ b/verl_speco/tq_owner.py @@ -16,6 +16,8 @@ from __future__ import annotations import logging +import os +from pathlib import Path import signal import threading from typing import Any @@ -76,13 +78,13 @@ def run_owner(config: Any, *, stop_event: threading.Event | None = None) -> int: # and enable only this process's copied configuration. tq_cfg["enable"] = True if not configure_transfer_queue(tq_cfg): - raise RuntimeError( - "Standalone TQ owner requires TransferQueue==0.1.7" - ) + raise RuntimeError("Standalone TQ owner requires TransferQueue==0.1.7") ray_cfg = tq_cfg.get("ray", {}) ray_address = ray_cfg.get("address") if not ray_address: - raise ValueError("transfer_queue.ray.address must point to a running Ray cluster") + raise ValueError( + "transfer_queue.ray.address must point to a running Ray cluster" + ) namespace = ray_cfg.get("namespace") event = stop_event or threading.Event() if stop_event is None: @@ -98,6 +100,9 @@ def run_owner(config: Any, *, stop_event: threading.Event | None = None) -> int: int(tq_cfg.get("schema_version", 1)), ) logger.info("TQ owner ready key=%s", ready_key) + ready_file = os.environ.get("SPECO_TQ_OWNER_READY_FILE") + if ready_file: + Path(ready_file).touch() wait_until_stopped(event) return 0 finally: diff --git a/verl_speco/trainer/standalone_checkpoint.py b/verl_speco/trainer/standalone_checkpoint.py index 4c91850e..6eae949e 100644 --- a/verl_speco/trainer/standalone_checkpoint.py +++ b/verl_speco/trainer/standalone_checkpoint.py @@ -1,3 +1,17 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + """Standalone drafter checkpoint runtime config helpers.""" from __future__ import annotations From 4887e8bcae1e8d71604d09f5c03c4fb18cfae7dc Mon Sep 17 00:00:00 2001 From: xxyyrr598 <1615681145@qq.com> Date: Tue, 25 Aug 2026 09:49:38 +0800 Subject: [PATCH 37/50] feat(tq): support drafter hidden state vllm training and enhance standalone TQ workflow --- .../run_qwen3-8b_drafter_hidden_state_vllm.sh | 88 +++++++++ .../run_qwen3-8b_drafter_separate_training.sh | 168 ++++++++++++++++-- tests/examples/test_example_scripts.py | 19 ++ tests/unit/test_producer_input_reader.py | 93 ++++++++++ .../test_standalone_tq_training_launcher.py | 79 +++++++- tests/unit/test_tq_producer.py | 32 ++++ tests/unit/test_vllm_feature_client.py | 7 +- verl_speco/config/speco_base.yaml | 3 + verl_speco/producer/input_reader.py | 147 +++++++++++++-- verl_speco/producer/vllm_feature_client.py | 7 +- verl_speco/standalone_tq_producer.py | 120 +++++++++++-- verl_speco/standalone_tq_training_launcher.py | 163 ++++++++++++++--- 12 files changed, 844 insertions(+), 82 deletions(-) create mode 100644 examples/run_qwen3-8b_drafter_hidden_state_vllm.sh diff --git a/examples/run_qwen3-8b_drafter_hidden_state_vllm.sh b/examples/run_qwen3-8b_drafter_hidden_state_vllm.sh new file mode 100644 index 00000000..c318d306 --- /dev/null +++ b/examples/run_qwen3-8b_drafter_hidden_state_vllm.sh @@ -0,0 +1,88 @@ +#!/usr/bin/env bash +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +set -euo pipefail +set -x + +# Start only the target-model vLLM used by the standalone TQ Producer. +# Run this script in its own terminal before starting the training script. +# +# Ascend example: +# DEVICE_ENV=ASCEND_RT_VISIBLE_DEVICES VLLM_DEVICES_0=0,1 VLLM_DEVICES_1=2,3 VLLM_TP=2 \ +# bash examples/run_qwen3-8b_drafter_hidden_state_vllm.sh +# CUDA example: +# DEVICE_ENV=CUDA_VISIBLE_DEVICES VLLM_DEVICES_0=0,1 VLLM_DEVICES_1=2,3 VLLM_TP=2 \ +# bash examples/run_qwen3-8b_drafter_hidden_state_vllm.sh + +MODEL_PATH=${MODEL_PATH:-/path/to/Qwen3-8B} +DEVICE_ENV=${DEVICE_ENV:-ASCEND_RT_VISIBLE_DEVICES} +VLLM_DEVICES_0=${VLLM_DEVICES_0:-0} +VLLM_DEVICES_1=${VLLM_DEVICES_1:-1} +VLLM_TP=${VLLM_TP:-1} +VLLM_HOST=${VLLM_HOST:-127.0.0.1} +VLLM_PORT_0=${VLLM_PORT_0:-8000} +VLLM_PORT_1=${VLLM_PORT_1:-8001} +VLLM_GPU_MEMORY_UTILIZATION=${VLLM_GPU_MEMORY_UTILIZATION:-0.8} +VLLM_MAX_NUM_SEQS=${VLLM_MAX_NUM_SEQS:-256} +# Auxiliary training layers followed by the target model's final hidden-state +# layer. Keep the auxiliary prefix aligned with DSPARK_TARGET_LAYER_IDS in the +# standalone training script. DSpark L1 loss consumes the final entry. +VLLM_HIDDEN_STATE_LAYER_IDS=${VLLM_HIDDEN_STATE_LAYER_IDS:-'[1,9,17,25,33,36]'} +HIDDEN_STATES_DIR=${HIDDEN_STATES_DIR:-/tmp/speco-vllm-hidden-states} + +mkdir -p "${HIDDEN_STATES_DIR}/service-0" "${HIDDEN_STATES_DIR}/service-1" + +SPECULATIVE_CONFIG=$(printf '{"method":"extract_hidden_states","num_speculative_tokens":1,"draft_model_config":{"hf_config":{"eagle_aux_hidden_state_layer_ids":%s}}}' "${VLLM_HIDDEN_STATE_LAYER_IDS}") + +start_vllm() { + local devices=$1 + local port=$2 + local hidden_states_dir=$3 + shift 3 + local kv_transfer_config + kv_transfer_config=$(printf '{"kv_connector":"ExampleHiddenStatesConnector","kv_role":"kv_producer","kv_connector_extra_config":{"shared_storage_path":"%s","use_synchronization_lock":true}}' "${hidden_states_dir}") + env "${DEVICE_ENV}=${devices}" vllm serve "${MODEL_PATH}" \ + --host "${VLLM_HOST}" \ + --port "${port}" \ + --tensor-parallel-size "${VLLM_TP}" \ + --gpu-memory-utilization "${VLLM_GPU_MEMORY_UTILIZATION}" \ + --max-num-seqs "${VLLM_MAX_NUM_SEQS}" \ + --speculative-config "${SPECULATIVE_CONFIG}" \ + --kv-transfer-config "${kv_transfer_config}" \ + --no-enable-chunked-prefill \ + "$@" & + STARTED_PID=$! +} + +PID_0="" +PID_1="" +cleanup() { + [[ -n "${PID_0}" ]] && kill "${PID_0}" 2>/dev/null || true + [[ -n "${PID_1}" ]] && kill "${PID_1}" 2>/dev/null || true + [[ -n "${PID_0}" ]] && wait "${PID_0}" 2>/dev/null || true + [[ -n "${PID_1}" ]] && wait "${PID_1}" 2>/dev/null || true +} +trap cleanup EXIT INT TERM + +start_vllm "${VLLM_DEVICES_0}" "${VLLM_PORT_0}" "${HIDDEN_STATES_DIR}/service-0" "$@" +PID_0=${STARTED_PID} +start_vllm "${VLLM_DEVICES_1}" "${VLLM_PORT_1}" "${HIDDEN_STATES_DIR}/service-1" "$@" +PID_1=${STARTED_PID} + +echo "VLLM_SERVICES_STARTED pid_0=${PID_0} endpoint_0=http://${VLLM_HOST}:${VLLM_PORT_0}/v1 pid_1=${PID_1} endpoint_1=http://${VLLM_HOST}:${VLLM_PORT_1}/v1" +set +e +wait -n "${PID_0}" "${PID_1}" +status=$? +set -e +exit "${status}" diff --git a/examples/run_qwen3-8b_drafter_separate_training.sh b/examples/run_qwen3-8b_drafter_separate_training.sh index e03a9cd2..5a49d844 100644 --- a/examples/run_qwen3-8b_drafter_separate_training.sh +++ b/examples/run_qwen3-8b_drafter_separate_training.sh @@ -15,22 +15,123 @@ set -euo pipefail set -x -# One-command standalone DSpark draft-model training. The launcher internally -# starts the hidden-state target vLLM and uses the -# Producer -> TransferQueue -> Consumer path. +# Standalone DSpark draft-model training using an already-running hidden-state +# vLLM. Start run_qwen3-8b_drafter_hidden_state_vllm.sh in another terminal +# first. This process owns Ray/TQ, Producer and Consumer, but it must not own +# the target vLLM so inference and training can use different accelerators. -project_name=verl_dspark_drafter -exp_name=qwen3_8b_dspark_separate_training +project_name=${PROJECT_NAME:-verl_dspark_drafter} +exp_name=${EXP_NAME:-qwen3_8b_dspark_separate_training} -draft_train_gpus_per_node=8 +draft_train_gpus_per_node=${TRAIN_GPUS:-2} -MODEL_PATH=/path/to/Qwen3-8B +MODEL_PATH=${MODEL_PATH:-/path/to/Qwen3-8B} # Ordinary verl prompt Parquet is supported; target vLLM generates responses. -TRAIN_FILE=/path/to/train_file.parquet -DRAFTER_PATH=/path/to/vllm-compatible-dspark-drafter -DRAFT_CKPTS_DIR=/path/to/dspark_draft_checkpoints +TRAIN_FILE=${TRAIN_FILE:-/path/to/train_file.parquet} +# Optional. Leave empty to initialize DSpark from the target-model/config +# fallback; set it only when loading or resuming an existing drafter. +DRAFTER_PATH=${DRAFTER_PATH:-} +DRAFT_CKPTS_DIR=${DRAFT_CKPTS_DIR:-/path/to/dspark_draft_checkpoints} PYTHON_BIN=${PYTHON_BIN:-python3} +DEVICE_ENV=${DEVICE_ENV:-ASCEND_RT_VISIBLE_DEVICES} +TRAIN_DEVICES=${TRAIN_DEVICES:-2,3} +SPECO_VLLM_ENDPOINTS=${SPECO_VLLM_ENDPOINTS:-'[http://127.0.0.1:8000/v1,http://127.0.0.1:8001/v1]'} +VLLM_READY_TIMEOUT_SECONDS=${VLLM_READY_TIMEOUT_SECONDS:-120} + +# Producer -> vLLM concurrency and bounded queues. MAX_INFLIGHT_REQUESTS is the +# process-wide request limit; PER_ENDPOINT_CONCURRENCY applies independently to +# every URL in SPECO_VLLM_ENDPOINTS. +VLLM_REQUEST_TIMEOUT=${VLLM_REQUEST_TIMEOUT:-120} +VLLM_MAX_INFLIGHT_REQUESTS=${VLLM_MAX_INFLIGHT_REQUESTS:-16} +VLLM_PER_ENDPOINT_CONCURRENCY=${VLLM_PER_ENDPOINT_CONCURRENCY:-4} +PRODUCER_INPUT_QUEUE_SIZE=${PRODUCER_INPUT_QUEUE_SIZE:-32} +PRODUCER_PUBLISH_QUEUE_SIZE=${PRODUCER_PUBLISH_QUEUE_SIZE:-16} +PRODUCER_MAX_PENDING_SAMPLES=${PRODUCER_MAX_PENDING_SAMPLES:-1024} +PRODUCER_PENDING_POLL_INTERVAL=${PRODUCER_PENDING_POLL_INTERVAL:-0.5} +PRODUCER_MAX_SEQUENCE_LENGTH=${PRODUCER_MAX_SEQUENCE_LENGTH:-8192} +PRODUCER_MAX_FEATURE_LENGTH=${PRODUCER_MAX_FEATURE_LENGTH:-512} +PRODUCER_GENERATION_MAX_TOKENS=${PRODUCER_GENERATION_MAX_TOKENS:-511} + +# Standalone trainer. +MAX_STEPS=${MAX_STEPS:-10} +SAVE_INTERVAL_STEPS=${SAVE_INTERVAL_STEPS:-5} +SAVE_FINAL_CHECKPOINT=${SAVE_FINAL_CHECKPOINT:-true} +BATCH_SIZE_PER_GPU=${BATCH_SIZE_PER_GPU:-2} +LEARNING_RATE=${LEARNING_RATE:-1e-6} +LR_WARMUP_STEPS=${LR_WARMUP_STEPS:-0} +LR_SCHEDULER_TYPE=${LR_SCHEDULER_TYPE:-constant} +LR_DECAY_STEPS=${LR_DECAY_STEPS:-100} +MIN_LR_RATIO=${MIN_LR_RATIO:-0.1} +PARAM_OFFLOAD=${PARAM_OFFLOAD:-true} +OPTIMIZER_OFFLOAD=${OPTIMIZER_OFFLOAD:-true} + +# DSpark architecture, sampling and losses. TARGET_LAYER_IDS must match the +# auxiliary layers exposed by both hidden-state vLLM services. +DSPARK_BLOCK_SIZE=${DSPARK_BLOCK_SIZE:-7} +DSPARK_NUM_ANCHORS=${DSPARK_NUM_ANCHORS:-32} +DSPARK_MAX_WINDOW=${DSPARK_MAX_WINDOW:-512} +DSPARK_LOSS_MODE=${DSPARK_LOSS_MODE:-full_vocab} +DSPARK_SAMPLED_CE_NEGATIVES=${DSPARK_SAMPLED_CE_NEGATIVES:-0} +DSPARK_LOSS_DECAY_GAMMA=${DSPARK_LOSS_DECAY_GAMMA:-7} +DSPARK_NUM_TARGET_LAYERS=${DSPARK_NUM_TARGET_LAYERS:-5} +DSPARK_NUM_HIDDEN_LAYERS=${DSPARK_NUM_HIDDEN_LAYERS:-5} +DSPARK_TARGET_LAYER_IDS=${DSPARK_TARGET_LAYER_IDS:-'[1,9,17,25,33]'} +DSPARK_MARKOV_RANK=${DSPARK_MARKOV_RANK:-256} +DSPARK_MARKOV_HEAD_TYPE=${DSPARK_MARKOV_HEAD_TYPE:-vanilla} +DSPARK_CE_LOSS_ALPHA=${DSPARK_CE_LOSS_ALPHA:-0.1} +DSPARK_L1_LOSS_ALPHA=${DSPARK_L1_LOSS_ALPHA:-0.45} +DSPARK_L1_CHUNK_SIZE=${DSPARK_L1_CHUNK_SIZE:-0} +# The current DSpark trainer rejects nonzero confidence loss because target +# acceptance labels are not part of the standalone feature protocol yet. +DSPARK_CONFIDENCE_LOSS_ALPHA=${DSPARK_CONFIDENCE_LOSS_ALPHA:-0.0} +DSPARK_DEBUG_LOG=${DSPARK_DEBUG_LOG:-false} +DSPARK_DEBUG_LOG_FIRST_N=${DSPARK_DEBUG_LOG_FIRST_N:-2} +DSPARK_DEBUG_LOG_INTERVAL=${DSPARK_DEBUG_LOG_INTERVAL:-100} + +export "${DEVICE_ENV}=${TRAIN_DEVICES}" +export SPECO_VLLM_ENDPOINTS + +# Fail before entering the unified launcher when the separately managed vLLM +# is absent. Otherwise a localhost endpoint would make the launcher start its +# fallback vLLM inside the training process and on the training devices. +"${PYTHON_BIN}" - "${SPECO_VLLM_ENDPOINTS}" "${VLLM_READY_TIMEOUT_SECONDS}" <<'PY' +import sys +import time +from urllib.error import URLError +from urllib.request import urlopen + +raw_endpoints = sys.argv[1].strip() +if not (raw_endpoints.startswith("[") and raw_endpoints.endswith("]")): + raise SystemExit("SPECO_VLLM_ENDPOINTS must use [url0,url1] syntax") +endpoints = [ + item.strip().strip("'\"").rstrip("/") + for item in raw_endpoints[1:-1].split(",") + if item.strip() +] +if not endpoints: + raise SystemExit("SPECO_VLLM_ENDPOINTS must contain at least one URL") +timeout_seconds = float(sys.argv[2]) +deadline = time.monotonic() + timeout_seconds +pending = set(endpoints) +while pending: + for endpoint in list(pending): + try: + with urlopen(f"{endpoint}/models", timeout=2) as response: + if 200 <= response.status < 300: + print(f"EXTERNAL_VLLM_READY endpoint={endpoint}", flush=True) + pending.remove(endpoint) + except (OSError, URLError): + pass + if pending and time.monotonic() >= deadline: + raise SystemExit( + "external hidden-state vLLM is not ready at: " + + ", ".join(sorted(pending)) + + "; start examples/run_qwen3-8b_drafter_hidden_state_vllm.sh first" + ) + if pending: + time.sleep(1) +PY PYTHONUNBUFFERED=1 "${PYTHON_BIN}" -m verl_speco.standalone_tq_training_launcher \ speco.draft_training.num_gpus_per_node=${draft_train_gpus_per_node} \ @@ -39,8 +140,8 @@ PYTHONUNBUFFERED=1 "${PYTHON_BIN}" -m verl_speco.standalone_tq_training_launcher data.train_files=${TRAIN_FILE} \ actor_rollout_ref.model.path=${MODEL_PATH} \ actor_rollout_ref.actor.strategy=fsdp2 \ - actor_rollout_ref.actor.fsdp_config.param_offload=True \ - actor_rollout_ref.actor.fsdp_config.optimizer_offload=True \ + actor_rollout_ref.actor.fsdp_config.param_offload=${PARAM_OFFLOAD} \ + actor_rollout_ref.actor.fsdp_config.optimizer_offload=${OPTIMIZER_OFFLOAD} \ actor_rollout_ref.rollout.tensor_model_parallel_size=1 \ actor_rollout_ref.rollout.drafter.enable=True \ actor_rollout_ref.rollout.drafter.enable_drafter_training=True \ @@ -48,13 +149,44 @@ PYTHONUNBUFFERED=1 "${PYTHON_BIN}" -m verl_speco.standalone_tq_training_launcher actor_rollout_ref.rollout.drafter.checkpoint_path=${DRAFT_CKPTS_DIR} \ actor_rollout_ref.rollout.drafter.speculative_algorithm=DSPARK \ actor_rollout_ref.rollout.drafter.training.mode=offline \ - actor_rollout_ref.rollout.drafter.training.max_steps=10 \ - actor_rollout_ref.rollout.drafter.training.save_interval_steps=5 \ - actor_rollout_ref.rollout.drafter.training.batch_size_per_gpu=2 \ - actor_rollout_ref.rollout.drafter.training.lr=1e-6 \ - actor_rollout_ref.rollout.drafter.training.lr_warmup_steps=0 \ - actor_rollout_ref.rollout.drafter.training.warmup_style=constant \ + actor_rollout_ref.rollout.drafter.training.max_steps=${MAX_STEPS} \ + actor_rollout_ref.rollout.drafter.training.save_interval_steps=${SAVE_INTERVAL_STEPS} \ + actor_rollout_ref.rollout.drafter.training.save_final_checkpoint=${SAVE_FINAL_CHECKPOINT} \ + actor_rollout_ref.rollout.drafter.training.batch_size_per_gpu=${BATCH_SIZE_PER_GPU} \ + actor_rollout_ref.rollout.drafter.training.lr=${LEARNING_RATE} \ + actor_rollout_ref.rollout.drafter.training.lr_warmup_steps=${LR_WARMUP_STEPS} \ + actor_rollout_ref.rollout.drafter.training.lr_scheduler_type=${LR_SCHEDULER_TYPE} \ + actor_rollout_ref.rollout.drafter.training.lr_decay_steps=${LR_DECAY_STEPS} \ + actor_rollout_ref.rollout.drafter.training.min_lr_ratio=${MIN_LR_RATIO} \ actor_rollout_ref.rollout.drafter.training.use_logits=False \ + actor_rollout_ref.rollout.drafter.training.dspark_block_size=${DSPARK_BLOCK_SIZE} \ + actor_rollout_ref.rollout.drafter.training.dspark_num_anchors=${DSPARK_NUM_ANCHORS} \ + actor_rollout_ref.rollout.drafter.training.dspark_max_window=${DSPARK_MAX_WINDOW} \ + actor_rollout_ref.rollout.drafter.training.dspark_loss_mode=${DSPARK_LOSS_MODE} \ + actor_rollout_ref.rollout.drafter.training.dspark_sampled_ce_negatives=${DSPARK_SAMPLED_CE_NEGATIVES} \ + actor_rollout_ref.rollout.drafter.training.dspark_loss_decay_gamma=${DSPARK_LOSS_DECAY_GAMMA} \ + actor_rollout_ref.rollout.drafter.training.dspark_num_target_layers=${DSPARK_NUM_TARGET_LAYERS} \ + actor_rollout_ref.rollout.drafter.training.dspark_num_hidden_layers=${DSPARK_NUM_HIDDEN_LAYERS} \ + actor_rollout_ref.rollout.drafter.training.dspark_target_layer_ids=${DSPARK_TARGET_LAYER_IDS} \ + actor_rollout_ref.rollout.drafter.training.dspark_markov_rank=${DSPARK_MARKOV_RANK} \ + actor_rollout_ref.rollout.drafter.training.dspark_markov_head_type=${DSPARK_MARKOV_HEAD_TYPE} \ + actor_rollout_ref.rollout.drafter.training.dspark_ce_loss_alpha=${DSPARK_CE_LOSS_ALPHA} \ + actor_rollout_ref.rollout.drafter.training.dspark_l1_loss_alpha=${DSPARK_L1_LOSS_ALPHA} \ + actor_rollout_ref.rollout.drafter.training.dspark_l1_chunk_size=${DSPARK_L1_CHUNK_SIZE} \ + actor_rollout_ref.rollout.drafter.training.dspark_confidence_loss_alpha=${DSPARK_CONFIDENCE_LOSS_ALPHA} \ + actor_rollout_ref.rollout.drafter.training.dspark_debug_log=${DSPARK_DEBUG_LOG} \ + actor_rollout_ref.rollout.drafter.training.dspark_debug_log_first_n=${DSPARK_DEBUG_LOG_FIRST_N} \ + actor_rollout_ref.rollout.drafter.training.dspark_debug_log_interval=${DSPARK_DEBUG_LOG_INTERVAL} \ + speco.standalone_tq_producer.request_timeout=${VLLM_REQUEST_TIMEOUT} \ + speco.standalone_tq_producer.max_inflight_requests=${VLLM_MAX_INFLIGHT_REQUESTS} \ + speco.standalone_tq_producer.per_endpoint_concurrency=${VLLM_PER_ENDPOINT_CONCURRENCY} \ + speco.standalone_tq_producer.input_queue_size=${PRODUCER_INPUT_QUEUE_SIZE} \ + speco.standalone_tq_producer.publish_queue_size=${PRODUCER_PUBLISH_QUEUE_SIZE} \ + speco.standalone_tq_producer.max_pending_samples=${PRODUCER_MAX_PENDING_SAMPLES} \ + speco.standalone_tq_producer.pending_poll_interval_seconds=${PRODUCER_PENDING_POLL_INTERVAL} \ + speco.standalone_tq_producer.max_sequence_length=${PRODUCER_MAX_SEQUENCE_LENGTH} \ + speco.standalone_tq_producer.max_feature_length=${PRODUCER_MAX_FEATURE_LENGTH} \ + speco.standalone_tq_producer.generation_max_tokens=${PRODUCER_GENERATION_MAX_TOKENS} \ trainer.project_name=${project_name} \ trainer.experiment_name=${exp_name} \ "$@" diff --git a/tests/examples/test_example_scripts.py b/tests/examples/test_example_scripts.py index ade09ffc..0699f5d0 100644 --- a/tests/examples/test_example_scripts.py +++ b/tests/examples/test_example_scripts.py @@ -27,6 +27,7 @@ for script in EXAMPLES if not script.name.endswith("_separate_training.sh") and script.name != "run_dspark_tq_producer.sh" + and script.name != "run_qwen3-8b_drafter_hidden_state_vllm.sh" ] @@ -79,6 +80,24 @@ def test_standalone_tq_training_example_uses_unified_launcher() -> None: assert "actor_rollout_ref.rollout.drafter.enable_drafter_training=True" in source assert "actor_rollout_ref.rollout.drafter.model_path=${DRAFTER_PATH}" in source assert "actor_rollout_ref.rollout.drafter.speculative_algorithm=DSPARK" in source + assert "speco.standalone_tq_producer.max_inflight_requests=" in source + assert "speco.standalone_tq_producer.per_endpoint_concurrency=" in source + assert "actor_rollout_ref.rollout.drafter.training.dspark_ce_loss_alpha=" in source + assert "actor_rollout_ref.rollout.drafter.training.dspark_l1_loss_alpha=" in source + + +def test_standalone_tq_hidden_state_vllm_uses_separate_devices() -> None: + source = ( + ROOT / "examples" / "run_qwen3-8b_drafter_hidden_state_vllm.sh" + ).read_text(encoding="utf-8") + + assert 'VLLM_DEVICES_0=${VLLM_DEVICES_0:-0}' in source + assert 'VLLM_DEVICES_1=${VLLM_DEVICES_1:-1}' in source + assert 'env "${DEVICE_ENV}=${devices}" vllm serve "${MODEL_PATH}"' in source + assert '--tensor-parallel-size "${VLLM_TP}"' in source + assert '--max-num-seqs "${VLLM_MAX_NUM_SEQS}"' in source + assert '"${VLLM_HIDDEN_STATE_LAYER_IDS}"' in source + assert '"kv_connector":"ExampleHiddenStatesConnector"' in source def test_standalone_tq_compatibility_example_delegates_to_formal_entry() -> None: diff --git a/tests/unit/test_producer_input_reader.py b/tests/unit/test_producer_input_reader.py index a79ff659..2cb976c2 100644 --- a/tests/unit/test_producer_input_reader.py +++ b/tests/unit/test_producer_input_reader.py @@ -194,6 +194,99 @@ def test_finalize_generated_request_aligns_connector_excluding_final_token() -> assert finalized.draft_position_ids.tolist() == [3, 4] +def test_prefilled_response_limits_vllm_prefix_after_selecting_training_window() -> None: + prepared = input_reader._build_tokenized_request( + sequence_no=0, + sample_id="long-response", + prompt_length=3, + full_ids=list(range(20)), + source_metadata={}, + config={"max_sequence_length": 8, "max_feature_length": 4}, + ) + + assert prepared.input_ids.numel() == 20 + assert prepared.feature_positions.tolist() == [2, 3, 4, 5] + assert prepared.prompt_token_ids == [0, 1, 2, 3, 4, 5] + + +def test_prefilled_response_rejects_prompt_prefix_beyond_vllm_limit() -> None: + with pytest.raises( + ValueError, + match=r"vLLM prefill of 13 tokens.*max_sequence_length=8", + ): + input_reader._build_tokenized_request( + sequence_no=0, + sample_id="long-prompt", + prompt_length=10, + full_ids=list(range(20)), + source_metadata={}, + config={"max_sequence_length": 8, "max_feature_length": 4}, + ) + + +def test_iter_jsonl_conversation_splits_final_assistant_response(tmp_path: Path) -> None: + input_path = tmp_path / "conversation.jsonl" + input_path.write_text( + json.dumps( + { + "conversation": [ + {"role": "human", "content": "Question"}, + {"role": "assistant", "content": "Answer"}, + ] + } + ) + + "\n", + encoding="utf-8", + ) + + record = next(input_reader.iter_input_records(input_path)) + + assert record.prompt == ({"role": "user", "content": "Question"},) + assert record.response == "Answer" + + +def test_iter_jsonl_legacy_conversations_normalizes_from_value(tmp_path: Path) -> None: + input_path = tmp_path / "conversations.jsonl" + input_path.write_text( + json.dumps( + { + "conversations": [ + {"from": "human", "value": "Question"}, + {"from": "gpt", "value": "Answer"}, + ] + } + ) + + "\n", + encoding="utf-8", + ) + + record = next(input_reader.iter_input_records(input_path)) + + assert record.prompt == ({"role": "user", "content": "Question"},) + assert record.response == "Answer" + + +def test_prepare_generated_prefill_request_uses_full_sequence_without_final_token() -> None: + request = input_reader.GenerationRequest( + sequence_no=0, + sample_id="generated-row", + prompt_token_ids=(10, 11, 12), + max_tokens=4, + source_metadata={}, + ) + + prepared = input_reader.prepare_generated_prefill_request( + request, + [20, 21], + {"max_sequence_length": 16, "max_feature_length": 8}, + ) + + assert prepared.input_ids.tolist() == [10, 11, 12, 20, 21] + assert prepared.loss_mask.tolist() == [0, 0, 0, 1, 1] + assert prepared.prompt_token_ids == [10, 11, 12, 20] + assert prepared.feature_positions.tolist() == [2, 3] + + def test_finalize_generated_request_rejects_misaligned_connector_tokens() -> None: request = input_reader.GenerationRequest( sequence_no=0, diff --git a/tests/unit/test_standalone_tq_training_launcher.py b/tests/unit/test_standalone_tq_training_launcher.py index 38e53960..5e6baf4c 100644 --- a/tests/unit/test_standalone_tq_training_launcher.py +++ b/tests/unit/test_standalone_tq_training_launcher.py @@ -49,7 +49,7 @@ def test_pipeline_config_derives_transport_identity_from_training_args() -> None assert config.tokenizer_path == "/models/Qwen3-8B" assert config.algorithm == "DSPARK" assert config.target_layer_ids == (1, 9, 17, 25, 33) - assert config.vllm_endpoint == "http://127.0.0.1:8000/v1" + assert config.vllm_endpoints == ("http://127.0.0.1:8000/v1",) assert config.run_id.startswith("dspark-") @@ -62,6 +62,34 @@ def test_pipeline_config_accepts_one_hydra_list_train_file() -> None: assert config.input_path == "/data/train.jsonl" +def test_pipeline_config_allows_missing_drafter_path_for_fresh_training() -> None: + args = [ + item + for item in _training_args() + if not item.startswith("actor_rollout_ref.rollout.drafter.model_path=") + ] + + config = resolve_pipeline_config(args, environ={}) + + assert config.model_path == "/models/Qwen3-8B" + + +def test_pipeline_config_accepts_multiple_vllm_endpoints() -> None: + config = resolve_pipeline_config( + _training_args(), + environ={ + "SPECO_VLLM_ENDPOINTS": ( + "[http://127.0.0.1:8000/v1,http://127.0.0.1:8001/v1]" + ) + }, + ) + + assert config.vllm_endpoints == ( + "http://127.0.0.1:8000/v1", + "http://127.0.0.1:8001/v1", + ) + + def test_target_final_layer_id_uses_local_model_config(tmp_path) -> None: (tmp_path / "config.json").write_text( json.dumps({"text_config": {"num_hidden_layers": 48}}), @@ -130,7 +158,7 @@ def test_pipeline_commands_hide_and_replace_tq_overrides() -> None: assert commands.vllm is not None assert commands.vllm[:3] == ["vllm", "serve", "/models/Qwen3-8B"] assert "ExampleHiddenStatesConnector" in " ".join(commands.vllm) - assert commands.vllm_endpoint == "http://127.0.0.1:8000/v1" + assert commands.vllm_endpoints == ("http://127.0.0.1:8000/v1",) assert commands.producer[:3] == [ "python", "-m", @@ -148,6 +176,53 @@ def test_pipeline_commands_hide_and_replace_tq_overrides() -> None: "vllm_endpoints=[http://127.0.0.1:8000/v1]" in item for item in commands.producer ) + assert any("max_samples=40" in item for item in commands.producer) + + +def test_pipeline_commands_pass_all_external_vllm_endpoints_to_producer() -> None: + endpoints = "[http://127.0.0.1:8000/v1,http://127.0.0.1:8001/v1]" + config = resolve_pipeline_config( + _training_args(), environ={"SPECO_VLLM_ENDPOINTS": endpoints} + ) + + commands = build_pipeline_commands( + config, + _training_args(), + ray_address="127.0.0.1:6379", + python_executable="python", + ) + + # Multiple services are started by the dedicated shell script. The unified + # launcher only verifies them and passes both URLs into the Producer pool. + assert commands.vllm is None + assert any(f"vllm_endpoints={endpoints}" in item for item in commands.producer) + + +def test_pipeline_commands_forward_producer_tuning_overrides() -> None: + args = [ + *_training_args(), + "speco.standalone_tq_producer.max_inflight_requests=32", + "speco.standalone_tq_producer.per_endpoint_concurrency=8", + "speco.standalone_tq_producer.max_feature_length=384", + ] + config = resolve_pipeline_config(args, environ={}) + + commands = build_pipeline_commands( + config, + args, + ray_address="127.0.0.1:6379", + python_executable="python", + ) + + assert ( + "speco.standalone_tq_producer.max_inflight_requests=32" + in commands.producer + ) + assert ( + "speco.standalone_tq_producer.per_endpoint_concurrency=8" + in commands.producer + ) + assert "speco.standalone_tq_producer.max_feature_length=384" in commands.producer class _FakeRuntimeContext: diff --git a/tests/unit/test_tq_producer.py b/tests/unit/test_tq_producer.py index 6ac2fd43..9bd7f586 100644 --- a/tests/unit/test_tq_producer.py +++ b/tests/unit/test_tq_producer.py @@ -148,11 +148,13 @@ def __init__(self, root: Path, *, close_error: BaseException | None = None): self.started = False self.closed = False self.generate_calls = 0 + self.prefill_calls = 0 async def start(self) -> None: self.started = True async def prefill(self, request: Any) -> RawVllmFeature: + self.prefill_calls += 1 path = self.root / f"{request.sample_id}.safetensors" path.write_bytes(b"temporary") self.paths.append(path) @@ -242,6 +244,35 @@ def test_run_producer_publishes_samples_then_eos(tmp_path: Path) -> None: assert tuple(first_fields["hidden_states"].shape) == (3, 6) +def test_run_producer_restarts_input_until_max_samples(tmp_path: Path) -> None: + input_path = tmp_path / "input.jsonl" + _write_input(input_path) + config = _config(input_path) + config["speco"]["standalone_tq_producer"]["max_samples"] = 5 + transport = _Transport() + pool = _Pool(tmp_path) + + stats = asyncio.run( + run_producer( + config, + transport=transport, + tokenizer=_Tokenizer(), + client_pool=pool, + ) + ) + + sample_tags = [ + tag + for tag in transport.records.values() + if tag.get("record_type") == "sample" + ] + eos_tags = [tag for tag in transport.records.values() if tag.get("status") == "eos"] + assert stats.input_count == stats.published_count == 5 + assert sorted(tag["sequence_no"] for tag in sample_tags) == [0, 1, 2, 3, 4] + assert pool.prefill_calls == 5 + assert eos_tags[0]["total_samples"] == 5 + + def test_run_producer_generates_response_for_verl_chat_prompt(tmp_path: Path) -> None: input_path = tmp_path / "dapo.jsonl" input_path.write_text( @@ -274,6 +305,7 @@ def test_run_producer_generates_response_for_verl_chat_prompt(tmp_path: Path) -> ] assert stats.input_count == stats.published_count == 1 assert pool.generate_calls == 1 + assert pool.prefill_calls == 1 assert len(sample_keys) == 1 fields = transport.payloads[sample_keys[0]] assert fields["input_ids"].tolist() == [10, 11] diff --git a/tests/unit/test_vllm_feature_client.py b/tests/unit/test_vllm_feature_client.py index b09bb601..d748f1de 100644 --- a/tests/unit/test_vllm_feature_client.py +++ b/tests/unit/test_vllm_feature_client.py @@ -23,7 +23,7 @@ ) -def test_request_generate_includes_output_tokens_in_hidden_states() -> None: +def test_request_generate_only_requests_generated_token_ids() -> None: calls = [] class Completions: @@ -54,7 +54,4 @@ async def create(self, **kwargs): assert response.generated_token_ids == (3, 4) assert calls[0]["max_tokens"] == 128 - assert calls[0]["extra_body"] == { - "return_token_ids": True, - "kv_transfer_params": {"include_output_tokens": True}, - } + assert calls[0]["extra_body"] == {"return_token_ids": True} diff --git a/verl_speco/config/speco_base.yaml b/verl_speco/config/speco_base.yaml index 2c0e5aa9..196e985e 100644 --- a/verl_speco/config/speco_base.yaml +++ b/verl_speco/config/speco_base.yaml @@ -41,6 +41,9 @@ speco: input_queue_size: 32 publish_queue_size: 16 max_pending_samples: 1024 + # Zero means one input-file pass. The unified launcher sets this to the + # exact number of samples required by max_steps and the global batch size. + max_samples: 0 pending_poll_interval_seconds: 0.5 owner_ready_timeout_seconds: 120 max_sequence_length: 8192 diff --git a/verl_speco/producer/input_reader.py b/verl_speco/producer/input_reader.py index 6cc2e433..69e4bdf8 100644 --- a/verl_speco/producer/input_reader.py +++ b/verl_speco/producer/input_reader.py @@ -75,8 +75,8 @@ def iter_input_records(path: str | os.PathLike[str]) -> Iterator[InputRecord]: raise ValueError( f"Producer input at {location} must be a JSON-style object" ) - prompt = _normalize_prompt(payload.get("prompt"), location) - response = payload.get("response") + prompt_value, response = _prompt_response_from_payload(payload, location) + prompt = _normalize_prompt(prompt_value, location) if response is not None and (not isinstance(response, str) or not response): raise ValueError( f"Producer input at {location} field 'response' must be a non-empty " @@ -90,7 +90,8 @@ def iter_input_records(path: str | os.PathLike[str]) -> Iterator[InputRecord]: source_metadata = { key: value for key, value in payload.items() - if key not in {"prompt", "response", "sample_id"} + if key + not in {"prompt", "response", "conversation", "conversations", "sample_id"} } yield InputRecord( sequence_no=sequence_no, @@ -102,6 +103,82 @@ def iter_input_records(path: str | os.PathLike[str]) -> Iterator[InputRecord]: sequence_no += 1 +def _prompt_response_from_payload( + payload: Mapping[str, Any], location: str +) -> tuple[Any, Any]: + """Normalize supported row schemas to Producer ``prompt``/``response``.""" + + if "prompt" in payload: + return payload.get("prompt"), payload.get("response") + + conversation = payload.get("conversation") + if conversation is not None: + messages = _normalize_conversation_messages( + conversation, + location, + role_key="role", + content_key="content", + ) + return _split_final_assistant(messages) + + conversations = payload.get("conversations") + if conversations is not None: + messages = _normalize_conversation_messages( + conversations, + location, + role_key="from", + content_key="value", + ) + return _split_final_assistant(messages) + + return None, payload.get("response") + + +def _normalize_conversation_messages( + value: Any, + location: str, + *, + role_key: str, + content_key: str, +) -> tuple[dict[str, str], ...]: + if not isinstance(value, (list, tuple)) or not value: + raise ValueError(f"Producer input at {location} conversation must be non-empty") + role_mapping = {"human": "user", "gpt": "assistant"} + messages: list[dict[str, str]] = [] + for index, item in enumerate(value): + if not isinstance(item, Mapping): + raise ValueError( + f"Producer input at {location} conversation item {index} must be an object" + ) + role = item.get(role_key) + content = item.get(content_key) + if not isinstance(role, str) or not role: + raise ValueError( + f"Producer input at {location} conversation item {index} requires " + f"string field {role_key!r}" + ) + if not isinstance(content, str) or not content: + raise ValueError( + f"Producer input at {location} conversation item {index} requires " + f"non-empty string field {content_key!r}" + ) + messages.append( + {"role": role_mapping.get(role.strip().lower(), role), "content": content} + ) + return tuple(messages) + + +def _split_final_assistant( + messages: tuple[dict[str, str], ...], +) -> tuple[tuple[dict[str, str], ...], str | None]: + if messages[-1]["role"] != "assistant": + return messages, None + prompt = messages[:-1] + if not prompt: + raise ValueError("Conversation cannot contain only an assistant response") + return prompt, messages[-1]["content"] + + def _normalize_prompt(value: Any, location: str) -> str | tuple[dict[str, str], ...]: if isinstance(value, str): return value @@ -332,6 +409,40 @@ def finalize_generated_request( ) +def prepare_generated_prefill_request( + request: GenerationRequest, + response_token_ids: Any, + config: Mapping[str, Any] | Any, +) -> TokenizedRequest: + """Build the full-sequence prefill request after target generation. + + The final sampled token has not itself passed through a model forward, so + target features are requested for ``prompt + completion[:-1]`` while the + full completion remains in ``input_ids`` as the next-token label sequence. + This matches the existing non-TQ vLLM replay path and does not require the + connector to capture decode-step hidden states. + """ + + response_ids = _token_ids(response_token_ids) + if not response_ids: + raise ValueError( + f"vLLM generated no response tokens for sample {request.sample_id!r}" + ) + prompt_ids = list(request.prompt_token_ids) + full_ids = [*prompt_ids, *response_ids] + hidden_input_ids = full_ids[:-1] + return _build_tokenized_request( + sequence_no=request.sequence_no, + sample_id=request.sample_id, + prompt_length=len(prompt_ids), + full_ids=full_ids, + source_metadata=request.source_metadata, + config=config, + vllm_prompt_token_ids=hidden_input_ids, + feature_end_limit=len(hidden_input_ids), + ) + + def _prompt_ids(prompt: str | tuple[dict[str, str], ...], tokenizer: Any) -> list[int]: if isinstance(prompt, str): return _token_ids(tokenizer(prompt, add_special_tokens=False)) @@ -364,12 +475,6 @@ def _build_tokenized_request( if int(input_ids.numel()) <= 0: raise ValueError(f"Producer sample {sample_id!r} produced no input tokens") - max_sequence_length = int(_config_value(config, "max_sequence_length", 0) or 0) - if max_sequence_length > 0 and int(input_ids.numel()) > max_sequence_length: - raise ValueError( - f"Producer sample {sample_id!r} has {int(input_ids.numel())} tokens, " - f"exceeding max_sequence_length={max_sequence_length}" - ) loss_mask = build_loss_mask(input_ids, prompt_length) position_ids = torch.arange(int(input_ids.numel()), dtype=torch.int64) @@ -382,6 +487,23 @@ def _build_tokenized_request( raise ValueError("max_feature_length must be 0 or at least 2") if max_feature_length > 1: feature_end = min(feature_start + max_feature_length, feature_end) + request_prompt_token_ids = ( + list(vllm_prompt_token_ids) + if vllm_prompt_token_ids is not None + else full_ids[:feature_end] + ) + max_sequence_length = int(_config_value(config, "max_sequence_length", 0) or 0) + if ( + max_sequence_length > 0 + and len(request_prompt_token_ids) > max_sequence_length + ): + raise ValueError( + f"Producer sample {sample_id!r} requires a vLLM prefill of " + f"{len(request_prompt_token_ids)} tokens after selecting its training " + f"window, exceeding max_sequence_length={max_sequence_length} " + f"(full_sequence_length={int(input_ids.numel())}, " + f"prompt_length={prompt_length})" + ) feature_positions = torch.arange(feature_start, feature_end, dtype=torch.int64) if int(feature_positions.numel()) <= 0: raise ValueError(f"Producer sample {sample_id!r} has an empty feature window") @@ -395,11 +517,7 @@ def _build_tokenized_request( feature_positions=feature_positions, draft_position_ids=draft_position_ids, source_metadata=dict(source_metadata), - vllm_prompt_token_ids=tuple( - vllm_prompt_token_ids - if vllm_prompt_token_ids is not None - else full_ids[:feature_end] - ), + vllm_prompt_token_ids=tuple(request_prompt_token_ids), ) @@ -439,5 +557,6 @@ def _config_value(config: Any, key: str, default: Any = None) -> Any: "finalize_generated_request", "iter_input_records", "prepare_generation_request", + "prepare_generated_prefill_request", "tokenize_record", ] diff --git a/verl_speco/producer/vllm_feature_client.py b/verl_speco/producer/vllm_feature_client.py index b57961dc..45c55ffb 100644 --- a/verl_speco/producer/vllm_feature_client.py +++ b/verl_speco/producer/vllm_feature_client.py @@ -101,16 +101,13 @@ async def request_generate( max_tokens: int, timeout: float, ) -> VllmResponse: - """Generate a response and request hidden states for prompt and output tokens.""" + """Generate response tokens; hidden states are collected by a later prefill.""" response = await client.completions.create( model=model, prompt=prompt_token_ids, max_tokens=max_tokens, - extra_body={ - "return_token_ids": True, - "kv_transfer_params": {"include_output_tokens": True}, - }, + extra_body={"return_token_ids": True}, timeout=timeout, ) choices = getattr(response, "choices", None) or [] diff --git a/verl_speco/standalone_tq_producer.py b/verl_speco/standalone_tq_producer.py index 71e5e9d6..6f3c4947 100644 --- a/verl_speco/standalone_tq_producer.py +++ b/verl_speco/standalone_tq_producer.py @@ -17,7 +17,7 @@ import asyncio import logging -from dataclasses import dataclass +from dataclasses import dataclass, replace from typing import Any, Mapping import torch @@ -29,9 +29,9 @@ from verl_speco.producer.input_reader import ( GenerationRequest, TokenizedRequest, - finalize_generated_request, iter_input_records, prepare_generation_request, + prepare_generated_prefill_request, tokenize_record, ) from verl_speco.producer.vllm_feature_client import ( @@ -145,6 +145,10 @@ def validate_producer_config(config: Any) -> None: raise ValueError(f"standalone_tq_producer fields must be positive: {invalid}") +def _should_log_sample_progress(count: int) -> bool: + return count <= 3 or count % 50 == 0 + + async def run_producer( config: Any, *, @@ -161,24 +165,43 @@ async def run_producer( connected = False pool = client_pool try: + logger.info( + "Standalone TQ Producer starting run_id=%s input=%s endpoints=%s", + run_id, + producer_cfg["input_path"], + producer_cfg["vllm_endpoints"], + ) if not transport.configure_transfer_queue(tq_cfg): raise RuntimeError("Standalone TQ Producer requires TransferQueue==0.1.7") ray_cfg = tq_cfg["ray"] + logger.info( + "Standalone TQ Producer connecting Ray address=%s namespace=%s", + ray_cfg["address"], + ray_cfg.get("namespace"), + ) transport.connect_ray_cluster( str(ray_cfg["address"]), str(ray_cfg["namespace"]) if ray_cfg.get("namespace") else None, ) + logger.info("Standalone TQ Producer connected Ray; initializing TQ client") transport.connect_transfer_queue_client() connected = True + logger.info("Standalone TQ Producer connected TQ; waiting for owner_ready") await _wait_for_owner_ready( transport, run_id, timeout=float(producer_cfg["owner_ready_timeout_seconds"]), poll_interval=float(producer_cfg["pending_poll_interval_seconds"]), ) + logger.info("Standalone TQ Producer observed owner_ready run_id=%s", run_id) if tokenizer is None: + logger.info( + "Standalone TQ Producer loading tokenizer path=%s", + producer_cfg["tokenizer_path"], + ) tokenizer = await asyncio.to_thread(_load_tokenizer, producer_cfg) + logger.info("Standalone TQ Producer tokenizer loaded") if pool is None: endpoint_concurrency = int(producer_cfg["per_endpoint_concurrency"]) pool = VllmFeatureClientPool( @@ -191,6 +214,7 @@ async def run_producer( request_timeout=float(producer_cfg["request_timeout"]), ) await pool.start() + logger.info("Standalone TQ Producer vLLM client pool started") feature_contract = FeatureContract( algorithm="DSPARK", @@ -213,16 +237,56 @@ async def run_producer( ) async def read_inputs() -> None: - for record in iter_input_records(str(producer_cfg["input_path"])): - request = ( - prepare_generation_request(record, tokenizer, producer_cfg) - if record.response is None - else tokenize_record(record, tokenizer, producer_cfg) + max_samples = int(producer_cfg.get("max_samples", 0) or 0) + epoch = 0 + while True: + epoch_count = 0 + for source_record in iter_input_records( + str(producer_cfg["input_path"]) + ): + if max_samples > 0 and stats.input_count >= max_samples: + break + # iter_input_records restarts sequence_no at zero on every + # pass. TQ keys require a run-global sequence number so a + # repeated sample never overwrites an earlier pending copy. + record = replace( + source_record, + sequence_no=stats.input_count, + ) + request = ( + prepare_generation_request(record, tokenizer, producer_cfg) + if record.response is None + else tokenize_record(record, tokenizer, producer_cfg) + ) + await input_queue.put(request) + stats.input_count += 1 + epoch_count += 1 + if _should_log_sample_progress(stats.input_count): + logger.info( + "Standalone TQ Producer queued input count=%s epoch=%s " + "sample_id=%s has_response=%s", + stats.input_count, + epoch, + record.sample_id, + record.response is not None, + ) + if epoch_count == 0 and stats.input_count == 0: + raise ValueError("Standalone TQ Producer input contains no samples") + if max_samples <= 0 or stats.input_count >= max_samples: + break + epoch += 1 + logger.info( + "Standalone TQ Producer restarting input epoch=%s " + "queued=%s target=%s", + epoch, + stats.input_count, + max_samples, ) - await input_queue.put(request) - stats.input_count += 1 for _ in range(worker_count): await input_queue.put(_INPUT_DONE) + logger.info( + "Standalone TQ Producer input exhausted total=%s", stats.input_count + ) async def request_worker() -> None: while True: @@ -236,14 +300,30 @@ async def request_worker() -> None: max_pending_samples=int(producer_cfg["max_pending_samples"]), poll_interval=float(producer_cfg["pending_poll_interval_seconds"]), ) - if isinstance(request, GenerationRequest): - raw = await pool.generate(request) - request = finalize_generated_request( - request, - raw.payload.get("token_ids"), - producer_cfg, - expected_response_token_ids=raw.generated_token_ids, + if _should_log_sample_progress(int(request.sequence_no) + 1): + logger.info( + "Standalone TQ Producer requesting vLLM sequence_no=%s " + "sample_id=%s mode=%s", + request.sequence_no, + request.sample_id, + "generate_then_prefill" + if isinstance(request, GenerationRequest) + else "prefill", ) + if isinstance(request, GenerationRequest): + generated = await pool.generate(request) + try: + request = prepare_generated_prefill_request( + request, + generated.generated_token_ids, + producer_cfg, + ) + finally: + # The generation request may still produce a prompt-only + # connector file. It is not the training payload; the + # following full-sequence prefill produces that payload. + await asyncio.to_thread(delete_temporary_result, generated) + raw = await pool.prefill(request) else: raw = await pool.prefill(request) stats.pending_bytes += int(raw.byte_size) @@ -268,6 +348,14 @@ async def publish_results() -> None: continue await publish_one(result, transport) stats.published_count += 1 + if _should_log_sample_progress(stats.published_count): + logger.info( + "Standalone TQ Producer published count=%s sequence_no=%s " + "sample_id=%s", + stats.published_count, + result.request.sequence_no, + result.request.sample_id, + ) stats.pending_bytes = max( stats.pending_bytes - int(result.raw.byte_size), 0 ) diff --git a/verl_speco/standalone_tq_training_launcher.py b/verl_speco/standalone_tq_training_launcher.py index f0a5e403..95e1ae8e 100644 --- a/verl_speco/standalone_tq_training_launcher.py +++ b/verl_speco/standalone_tq_training_launcher.py @@ -52,9 +52,40 @@ _DSPARK_LAYER_IDS_KEY = ( "actor_rollout_ref.rollout.drafter.training.dspark_target_layer_ids" ) +_MAX_STEPS_KEY = "actor_rollout_ref.rollout.drafter.training.max_steps" +_BATCH_SIZE_PER_GPU_KEY = ( + "actor_rollout_ref.rollout.drafter.training.batch_size_per_gpu" +) +_NPROC_KEYS = ( + "speco.draft_training.nproc_per_node", + "speco.draft_training.num_gpus_per_node", + "actor_rollout_ref.rollout.drafter.training.nproc_per_node", + "actor_rollout_ref.rollout.drafter.training.num_gpus_per_node", +) +_NNODES_KEYS = ( + "speco.draft_training.nnodes", + "speco.draft_training.num_nodes", + "actor_rollout_ref.rollout.drafter.training.nnodes", + "actor_rollout_ref.rollout.drafter.training.num_nodes", +) _TQ_PREFIX = "actor_rollout_ref.rollout.drafter.training.transfer_queue" _FEATURE_STORE_PREFIX = "actor_rollout_ref.rollout.drafter.training.feature_store" +_PRODUCER_PREFIX = "speco.standalone_tq_producer" +_PRODUCER_TUNING_KEYS = frozenset( + { + f"{_PRODUCER_PREFIX}.request_timeout", + f"{_PRODUCER_PREFIX}.max_inflight_requests", + f"{_PRODUCER_PREFIX}.per_endpoint_concurrency", + f"{_PRODUCER_PREFIX}.input_queue_size", + f"{_PRODUCER_PREFIX}.publish_queue_size", + f"{_PRODUCER_PREFIX}.max_pending_samples", + f"{_PRODUCER_PREFIX}.pending_poll_interval_seconds", + f"{_PRODUCER_PREFIX}.max_sequence_length", + f"{_PRODUCER_PREFIX}.max_feature_length", + f"{_PRODUCER_PREFIX}.generation_max_tokens", + } +) _INTERNAL_OVERRIDE_KEYS = frozenset( { f"{_FEATURE_STORE_PREFIX}.type", @@ -88,14 +119,14 @@ class PipelineConfig: tokenizer_path: str algorithm: str target_layer_ids: tuple[int, ...] - vllm_endpoint: str + vllm_endpoints: tuple[str, ...] run_id: str @dataclass(frozen=True) class PipelineCommands: vllm: list[str] | None - vllm_endpoint: str + vllm_endpoints: tuple[str, ...] owner: list[str] producer: list[str] consumer: list[str] @@ -125,6 +156,14 @@ def _find_override(overrides: Sequence[str], key: str) -> str | None: return None +def _find_first_override(overrides: Sequence[str], keys: Sequence[str]) -> str | None: + for key in keys: + value = _find_override(overrides, key) + if value is not None: + return value + return None + + def _strip_quotes(value: str) -> str: value = value.strip() if len(value) >= 2 and value[0] == value[-1] and value[0] in {"'", '"'}: @@ -161,6 +200,36 @@ def _parse_layer_ids(value: str | None) -> tuple[int, ...]: return result +def _parse_vllm_endpoints(env: Mapping[str, str]) -> tuple[str, ...]: + """Read a Hydra-style endpoint list while preserving the singular fallback.""" + + configured = str(env.get("SPECO_VLLM_ENDPOINTS", "")).strip() + if configured: + text = _strip_quotes(configured) + if not (text.startswith("[") and text.endswith("]")): + raise ValueError( + "SPECO_VLLM_ENDPOINTS must be a Hydra-style list, for example " + "[http://127.0.0.1:8000/v1,http://127.0.0.1:8001/v1]" + ) + endpoints = tuple( + _strip_quotes(item).rstrip("/") + for item in text[1:-1].split(",") + if item.strip() + ) + else: + endpoint = str( + env.get("SPECO_VLLM_ENDPOINT", _DEFAULT_VLLM_ENDPOINT) + ).strip() + endpoints = (endpoint.rstrip("/"),) if endpoint else () + if not endpoints: + raise ValueError("At least one hidden-state vLLM endpoint is required") + for endpoint in endpoints: + parsed = urlparse(endpoint) + if parsed.scheme not in {"http", "https"} or not parsed.hostname: + raise ValueError(f"Invalid vLLM endpoint: {endpoint!r}") + return endpoints + + def _required_override(overrides: Sequence[str], key: str) -> str: value = _find_override(overrides, key) normalized = _strip_quotes(value or "") @@ -169,6 +238,33 @@ def _required_override(overrides: Sequence[str], key: str) -> str: return normalized +def _positive_int_override( + overrides: Sequence[str], keys: Sequence[str], *, default: int +) -> int: + raw = _find_first_override(overrides, keys) + value = int(_strip_quotes(raw)) if raw is not None else int(default) + if value <= 0: + raise ValueError(f"{keys[0]} must be positive, got {value}") + return value + + +def _producer_max_samples(training_args: Sequence[str]) -> int: + """Return samples needed for exactly max_steps complete global batches.""" + + raw_max_steps = _find_override(training_args, _MAX_STEPS_KEY) + max_steps = int(_strip_quotes(raw_max_steps)) if raw_max_steps is not None else 1000 + if max_steps <= 0: + # An unbounded training run cannot have a finite Producer target. Keep + # the direct Producer's one-pass behavior instead of looping forever. + return 0 + batch_size = _positive_int_override( + training_args, (_BATCH_SIZE_PER_GPU_KEY,), default=4 + ) + nproc = _positive_int_override(training_args, _NPROC_KEYS, default=1) + nnodes = _positive_int_override(training_args, _NNODES_KEYS, default=1) + return max_steps * batch_size * nproc * nnodes + + def _stable_path_identity(kind: str, path: str) -> str: digest = hashlib.sha256(path.encode("utf-8")).hexdigest() return f"{kind}-path-sha256-{digest}" @@ -204,7 +300,6 @@ def resolve_pipeline_config( env = os.environ if environ is None else environ model_path = _required_override(training_args, _MODEL_PATH_KEY) - _required_override(training_args, _DRAFTER_PATH_KEY) input_path = _single_train_file(_find_override(training_args, _TRAIN_FILES_KEY)) tokenizer_path = _strip_quotes( _find_override(training_args, _TOKENIZER_PATH_KEY) or model_path @@ -217,16 +312,14 @@ def resolve_pipeline_config( target_layer_ids = _parse_layer_ids( _find_override(training_args, _DSPARK_LAYER_IDS_KEY) ) - endpoint = str(env.get("SPECO_VLLM_ENDPOINT", _DEFAULT_VLLM_ENDPOINT)).strip() - if not endpoint: - raise ValueError("SPECO_VLLM_ENDPOINT must not be empty") + endpoints = _parse_vllm_endpoints(env) return PipelineConfig( input_path=input_path, model_path=model_path, tokenizer_path=tokenizer_path, algorithm=algorithm, target_layer_ids=target_layer_ids, - vllm_endpoint=endpoint.rstrip("/"), + vllm_endpoints=endpoints, run_id=f"dspark-{uuid.uuid4().hex}", ) @@ -303,9 +396,7 @@ def build_pipeline_commands( f"{_TQ_PREFIX}.backend.SimpleStorage.total_storage_size=17179869184", f"{_TQ_PREFIX}.backend.SimpleStorage.num_data_storage_units=8", ] - parsed_endpoint = urlparse(config.vllm_endpoint) - if parsed_endpoint.scheme not in {"http", "https"} or not parsed_endpoint.hostname: - raise ValueError(f"Invalid vLLM endpoint: {config.vllm_endpoint!r}") + parsed_endpoint = urlparse(config.vllm_endpoints[0]) vllm_port = parsed_endpoint.port or ( 443 if parsed_endpoint.scheme == "https" else 80 ) @@ -334,7 +425,11 @@ def build_pipeline_commands( }, } vllm = None - if parsed_endpoint.hostname in {"127.0.0.1", "localhost", "0.0.0.0"}: + if len(config.vllm_endpoints) == 1 and parsed_endpoint.hostname in { + "127.0.0.1", + "localhost", + "0.0.0.0", + }: vllm = [ "vllm", "serve", @@ -357,12 +452,19 @@ def build_pipeline_commands( "verl_speco.tq_owner", *tq_overrides, ] + producer_tuning_overrides = [ + item + for item in training_args + if (parsed := _split_override(item)) is not None + and parsed[0] in _PRODUCER_TUNING_KEYS + ] producer = [ python_executable, "-m", "verl_speco.standalone_tq_producer", f"{_ALGORITHM_KEY}={config.algorithm}", *tq_overrides, + *producer_tuning_overrides, f"speco.standalone_tq_producer.input_path={config.input_path}", f"speco.standalone_tq_producer.tokenizer_path={config.tokenizer_path}", "speco.standalone_tq_producer.tokenizer_fingerprint=" @@ -373,8 +475,10 @@ def build_pipeline_commands( "speco.standalone_tq_producer.target_layer_ids=" + _hydra_list(config.target_layer_ids), "speco.standalone_tq_producer.vllm_endpoints=" - + _hydra_list((config.vllm_endpoint,)), + + _hydra_list(config.vllm_endpoints), f"speco.standalone_tq_producer.vllm_model={config.model_path}", + "speco.standalone_tq_producer.max_samples=" + + str(_producer_max_samples(training_args)), ] consumer_internal = [ f"{_FEATURE_STORE_PREFIX}.type=tq", @@ -392,7 +496,7 @@ def build_pipeline_commands( ] return PipelineCommands( vllm=vllm, - vllm_endpoint=config.vllm_endpoint, + vllm_endpoints=config.vllm_endpoints, owner=owner, producer=producer, consumer=consumer, @@ -489,13 +593,18 @@ def run_pipeline( prefix="speco-vllm-hidden-states-" ) hidden_states_dir = Path(hidden_states_temp.name) - config_endpoint = commands.vllm_endpoint - if not endpoint_ready(config_endpoint): + unavailable_endpoints = [ + endpoint + for endpoint in commands.vllm_endpoints + if not endpoint_ready(endpoint) + ] + if unavailable_endpoints: if commands.vllm is None: raise RuntimeError( - "The configured remote hidden-state vLLM is unavailable at " - f"{config_endpoint}" + "The configured hidden-state vLLM endpoints are unavailable: " + + ", ".join(unavailable_endpoints) ) + config_endpoint = commands.vllm_endpoints[0] vllm_command = [ part.replace(_VLLM_HIDDEN_STATES_DIR, str(hidden_states_dir)) for part in commands.vllm @@ -517,10 +626,17 @@ def run_pipeline( timeout_seconds=vllm_ready_timeout_seconds, endpoint_ready=endpoint_ready, ) - elif not endpoint_ready(config_endpoint): - raise RuntimeError( - f"hidden-state vLLM became unavailable at {config_endpoint}" - ) + else: + unavailable_endpoints = [ + endpoint + for endpoint in commands.vllm_endpoints + if not endpoint_ready(endpoint) + ] + if unavailable_endpoints: + raise RuntimeError( + "hidden-state vLLM became unavailable at: " + + ", ".join(unavailable_endpoints) + ) logger.info("Starting standalone DSpark Consumer") consumer = popen(commands.consumer, env=base_env) logger.info("Starting standalone vLLM Producer") @@ -600,7 +716,10 @@ def main(argv: Sequence[str] | None = None) -> int: ] for role, command in printable_commands: if command is None: - print(f"{role}: external service at {commands.vllm_endpoint}") + print( + f"{role}: external services at " + + ", ".join(commands.vllm_endpoints) + ) continue print(f"{role}: {_format_command(command)}") return 0 From d72d589267039db2ad3c05b7563f2ea995ff20d8 Mon Sep 17 00:00:00 2001 From: xxyyrr598 <1615681145@qq.com> Date: Wed, 26 Aug 2026 09:58:06 +0800 Subject: [PATCH 38/50] feat(tq): bump drafter sample protocol to schema v2 and add ready tag parsing --- tests/unit/test_drafter_sample_protocol.py | 157 +++---- .../test_standalone_tq_training_launcher.py | 28 ++ tests/unit/test_tq_consumer.py | 23 +- tests/unit/test_tq_producer.py | 52 ++- tools/tq_connection_smoke.py | 30 +- tools/tq_delayed_test_producer.py | 20 +- verl_speco/config/speco_base.yaml | 2 +- verl_speco/standalone_tq_producer.py | 67 +-- verl_speco/standalone_tq_training_launcher.py | 48 +- verl_speco/tq_owner.py | 3 +- verl_speco/trainer/target_feature_replay.py | 22 +- verl_speco/trainer/tq_feature_store.py | 17 +- verl_speco/transport/__init__.py | 4 + .../transport/drafter_sample_protocol.py | 411 +++++++++--------- 14 files changed, 472 insertions(+), 412 deletions(-) diff --git a/tests/unit/test_drafter_sample_protocol.py b/tests/unit/test_drafter_sample_protocol.py index 393b41ab..3a93e018 100644 --- a/tests/unit/test_drafter_sample_protocol.py +++ b/tests/unit/test_drafter_sample_protocol.py @@ -1,16 +1,5 @@ # Copyright 2026 Bytedance Ltd. and/or its affiliates -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. +# Licensed under the Apache License, Version 2.0 from __future__ import annotations from dataclasses import replace @@ -20,120 +9,136 @@ from verl_speco.trainer.feature_store import DraftFeatureSample from verl_speco.transport.drafter_sample_protocol import ( + PROTOCOL_SCHEMA_VERSION, ExpectedFeatureConfig, SampleMetadata, decode_sample, encode_sample, + is_ready_sample_tag, make_eos_record, make_ready_tag, make_sample_key, + parse_ready_tag, ) def _metadata() -> SampleMetadata: return SampleMetadata( - schema_version=1, + schema_version=PROTOCOL_SCHEMA_VERSION, run_id="run-a", sample_id="train-000017", sequence_no=17, - algorithm="DSPARK", - target_model_id="/models/Qwen3-8B", - target_model_revision="rev-a", - tokenizer_fingerprint="sha256:tokenizer", - target_layer_ids=[2, 8, 14, -1], - hidden_states_layout="dflash_aux_plus_last", - hidden_dtype="bfloat16", - hidden_shape=[4, 16], - feature_length=4, - full_sequence_length=10, - feature_start=6, - feature_end=10, - use_logits=False, ) -def _sample() -> DraftFeatureSample: +def _sample(*, algorithm: str = "DSPARK") -> DraftFeatureSample: return DraftFeatureSample( - algorithm="DSPARK", + algorithm=algorithm, input_ids=torch.tensor([10, 11, 12, 13]), loss_mask=torch.tensor([0.0, 1.0, 1.0, 1.0]), position_ids=torch.tensor([6, 7, 8, 9]), hidden_states=torch.arange(64, dtype=torch.bfloat16).reshape(4, 16), - metadata={"ignored_at_wire_boundary": True}, + last_hidden_states=torch.arange(32, dtype=torch.float32).reshape(4, 8), + metadata={ + "hidden_states_layout": "dflash_aux_plus_last", + "target_layer_ids": [2, 8, 14], + "hidden_positions": torch.tensor([6, 7, 8, 9]), + "nested": {"pair": ("a", 2)}, + }, ) -def _expected() -> ExpectedFeatureConfig: - meta = _metadata() - return ExpectedFeatureConfig( - run_id=meta.run_id, - target_model_id=meta.target_model_id, - target_model_revision=meta.target_model_revision, - tokenizer_fingerprint=meta.tokenizer_fingerprint, - target_layer_ids=meta.target_layer_ids, - hidden_states_layout=meta.hidden_states_layout, - hidden_dtype=meta.hidden_dtype, - ) - - -def test_sample_round_trip() -> None: +def test_sample_round_trip_preserves_complete_sample() -> None: meta = _metadata() key = make_sample_key(meta) fields = encode_sample(_sample(), meta) - restored = decode_sample(key, make_ready_tag(meta), fields, _expected()) + restored = decode_sample( + key, + make_ready_tag(meta), + fields, + ExpectedFeatureConfig(run_id=meta.run_id), + ) - assert key == "drafter:v1:run-a:000000000017:train-000017" - assert tuple(restored.hidden_states.shape) == (4, 16) - assert restored.hidden_states.dtype == torch.bfloat16 - assert restored.input_ids.tolist() == [10, 11, 12, 13] - assert restored.metadata["sequence_no"] == 17 - assert restored.metadata["hidden_states_layout"] == "dflash_aux_plus_last" - assert fields["metadata_json"].dtype == torch.uint8 + assert key == "drafter:v2:run-a:000000000017:train-000017" + assert restored.algorithm == "DSPARK" + assert torch.equal(restored.input_ids, _sample().input_ids) + assert torch.equal(restored.hidden_states, _sample().hidden_states) + assert torch.equal(restored.last_hidden_states, _sample().last_hidden_states) + assert torch.equal( + restored.metadata["hidden_positions"], _sample().metadata["hidden_positions"] + ) + assert restored.metadata["nested"]["pair"] == ("a", 2) + assert fields["sample__manifest_json"].dtype == torch.uint8 -def test_decode_rejects_identity_mismatch() -> None: +def test_hidden_state_tensor_list_round_trip() -> None: + sample = replace( + _sample(algorithm="EAGLE3"), + hidden_states=[torch.ones(4, 3), torch.zeros(4, 5)], + ) meta = _metadata() - fields = encode_sample(_sample(), meta) - bad_tag = {**make_ready_tag(meta), "sample_id": "wrong"} - with pytest.raises(ValueError, match="tag mismatch for sample_id"): - decode_sample(make_sample_key(meta), bad_tag, fields, _expected()) + restored = decode_sample( + make_sample_key(meta), + make_ready_tag(meta), + encode_sample(sample, meta), + ExpectedFeatureConfig(run_id="run-a"), + ) + assert restored.algorithm == "EAGLE3" + assert isinstance(restored.hidden_states, list) + assert [tuple(value.shape) for value in restored.hidden_states] == [(4, 3), (4, 5)] -def test_decode_rejects_consumer_contract_mismatch() -> None: +def test_ready_parser_is_the_shared_discovery_contract() -> None: meta = _metadata() - fields = encode_sample(_sample(), meta) - expected = replace(_expected(), target_model_revision="different") - with pytest.raises(ValueError, match="target_model_revision"): - decode_sample(make_sample_key(meta), make_ready_tag(meta), fields, expected) + tag = make_ready_tag(meta) + assert parse_ready_tag(tag, run_id="run-a") == meta + assert is_ready_sample_tag(tag, run_id="run-a") + assert not is_ready_sample_tag(tag, run_id="another-run") + assert parse_ready_tag({**tag, "sequence_no": "bad"}, run_id="run-a") is None + assert parse_ready_tag({**tag, "schema_version": 1}, run_id="run-a") is None -def test_encode_rejects_shape_mismatch() -> None: - meta = replace(_metadata(), hidden_shape=[4, 32]) - with pytest.raises(ValueError, match="hidden_states shape mismatch"): - encode_sample(_sample(), meta) +def test_decode_rejects_identity_mismatch() -> None: + meta = _metadata() + fields = encode_sample(_sample(), meta) + bad_tag = {**make_ready_tag(meta), "sample_id": "wrong"} + with pytest.raises(ValueError, match="key mismatch"): + decode_sample( + make_sample_key(meta), + bad_tag, + fields, + ExpectedFeatureConfig(run_id="run-a"), + ) -def test_protocol_algorithm_is_not_hardcoded_to_dspark() -> None: - meta = replace(_metadata(), algorithm="EAGLE3") - sample = replace(_sample(), algorithm="EAGLE3") - expected = replace(_expected(), algorithm="EAGLE3") +def test_metadata_codec_rejects_lossy_unknown_values() -> None: + sample = replace(_sample(), metadata={"unsupported": object()}) + with pytest.raises(TypeError, match="metadata.unsupported"): + encode_sample(sample, _metadata()) - key = make_sample_key(meta) - fields = encode_sample(sample, meta) - restored = decode_sample(key, make_ready_tag(meta), fields, expected) - assert make_ready_tag(meta)["algorithm"] == "EAGLE3" - assert restored.algorithm == "EAGLE3" +@pytest.mark.parametrize("algorithm", ["EAGLE3", "DFLASH", "DSPARK", "DOMINO"]) +def test_protocol_algorithm_is_not_hardcoded(algorithm: str) -> None: + meta = _metadata() + sample = _sample(algorithm=algorithm) + restored = decode_sample( + make_sample_key(meta), + make_ready_tag(meta), + encode_sample(sample, meta), + ExpectedFeatureConfig(run_id="run-a"), + ) + assert restored.algorithm == algorithm + assert "algorithm" not in make_ready_tag(meta) def test_eos_record_is_control_only() -> None: key, fields, tag = make_eos_record("run-a", 18) - assert key == "control:v1:run-a:eos" + assert key == "control:v2:run-a:eos" assert fields["marker"].tolist() == [1] assert tag == { "record_type": "control", "status": "eos", - "schema_version": 1, + "schema_version": 2, "run_id": "run-a", "total_samples": 18, } diff --git a/tests/unit/test_standalone_tq_training_launcher.py b/tests/unit/test_standalone_tq_training_launcher.py index 5e6baf4c..f0921c0e 100644 --- a/tests/unit/test_standalone_tq_training_launcher.py +++ b/tests/unit/test_standalone_tq_training_launcher.py @@ -53,6 +53,34 @@ def test_pipeline_config_derives_transport_identity_from_training_args() -> None assert config.run_id.startswith("dspark-") +def test_pipeline_config_reads_non_dspark_algorithm_from_training_args() -> None: + args = [ + item.replace("speculative_algorithm=DSPARK", "speculative_algorithm=DFLASH") + for item in _training_args() + ] + args.append( + "actor_rollout_ref.rollout.drafter.training.dflash_target_layer_ids=[2,10,20]" + ) + + config = resolve_pipeline_config(args, environ={}) + + assert config.algorithm == "DFLASH" + assert config.target_layer_ids == (2, 10, 20) + assert config.run_id.startswith("dflash-") + + +def test_pipeline_config_prefers_generic_producer_layer_ids() -> None: + args = [ + *_training_args(), + "speco.standalone_tq_producer.target_layer_ids=[3,11,21]", + "actor_rollout_ref.rollout.drafter.training.dspark_target_layer_ids=[1,9,17]", + ] + + config = resolve_pipeline_config(args, environ={}) + + assert config.target_layer_ids == (3, 11, 21) + + def test_pipeline_config_accepts_one_hydra_list_train_file() -> None: args = _training_args() args[0] = "data.train_files=['/data/train.jsonl']" diff --git a/tests/unit/test_tq_consumer.py b/tests/unit/test_tq_consumer.py index bce05a96..59b5736b 100644 --- a/tests/unit/test_tq_consumer.py +++ b/tests/unit/test_tq_consumer.py @@ -29,6 +29,7 @@ build_assignments, ) from verl_speco.transport.drafter_sample_protocol import ( + PROTOCOL_SCHEMA_VERSION, SampleMetadata, encode_sample, make_ready_tag, @@ -42,29 +43,16 @@ def _config() -> dict: "ray": {"address": "ray-head:6379", "namespace": "speco-drafter"}, "partition_id": "speco_drafter_features", "run_id": "run-a", - "schema_version": 1, + "schema_version": PROTOCOL_SCHEMA_VERSION, } def _metadata(sequence_no: int = 0) -> SampleMetadata: return SampleMetadata( - schema_version=1, + schema_version=PROTOCOL_SCHEMA_VERSION, run_id="run-a", sample_id=f"sample-{sequence_no}", sequence_no=sequence_no, - algorithm="DSPARK", - target_model_id="producer-model-is-not-strictly-checked", - target_model_revision="producer-revision", - tokenizer_fingerprint="producer-tokenizer", - target_layer_ids=[2, 8, 14], - hidden_states_layout="dflash_aux_plus_last", - hidden_dtype="float32", - hidden_shape=[3, 4], - feature_length=3, - full_sequence_length=3, - feature_start=0, - feature_end=3, - use_logits=False, ) @@ -75,6 +63,7 @@ def _sample() -> DraftFeatureSample: loss_mask=torch.tensor([1.0, 1.0, 1.0]), position_ids=torch.tensor([0, 1, 2]), hidden_states=torch.arange(12, dtype=torch.float32).reshape(3, 4), + metadata={"target_model_revision": "producer-revision"}, ) @@ -123,10 +112,10 @@ def test_tq_store_connect_filter_sort_and_minimal_decode(monkeypatch) -> None: entries[0].key: entries[0].tag, unrelated.key: unrelated.tag, entries[1].key: entries[1].tag, - "control:v1:run-a:eos": { + "control:v2:run-a:eos": { "record_type": "control", "status": "eos", - "schema_version": 1, + "schema_version": PROTOCOL_SCHEMA_VERSION, "run_id": "run-a", "total_samples": 2, }, diff --git a/tests/unit/test_tq_producer.py b/tests/unit/test_tq_producer.py index 9bd7f586..5744d2c8 100644 --- a/tests/unit/test_tq_producer.py +++ b/tests/unit/test_tq_producer.py @@ -24,6 +24,7 @@ from verl_speco.producer.vllm_feature_client import RawVllmFeature from verl_speco.standalone_tq_producer import run_producer, validate_producer_config +from verl_speco.transport.drafter_sample_protocol import PROTOCOL_SCHEMA_VERSION def _config(input_path: Path) -> dict[str, Any]: @@ -69,7 +70,7 @@ def _config(input_path: Path) -> dict[str, Any]: }, "partition_id": "speco_drafter_features", "run_id": "run-a", - "schema_version": 1, + "schema_version": PROTOCOL_SCHEMA_VERSION, }, }, } @@ -102,10 +103,10 @@ class _Transport: def __init__(self, *, fail_sample_put: bool = False): self.fail_sample_put = fail_sample_put self.records: dict[str, dict[str, Any]] = { - "control:v1:run-a:owner-ready": { + "control:v2:run-a:owner-ready": { "record_type": "control", "status": "owner_ready", - "schema_version": 1, + "schema_version": PROTOCOL_SCHEMA_VERSION, "run_id": "run-a", } } @@ -194,6 +195,13 @@ async def close(self) -> None: raise self.close_error +class _MisalignedPool(_Pool): + async def prefill(self, request: Any) -> RawVllmFeature: + raw = await super().prefill(request) + raw.payload["hidden_states"] = raw.payload["hidden_states"][-1:] + return raw + + def _write_input(path: Path) -> None: records = [ {"sample_id": "sample-1", "prompt": "Q1: ", "response": "A1"}, @@ -227,13 +235,13 @@ def test_run_producer_publishes_samples_then_eos(tmp_path: Path) -> None: ] eos_tags = [tag for tag in transport.records.values() if tag.get("status") == "eos"] assert stats.input_count == stats.published_count == 2 - assert stats.failed_count == stats.pending_bytes == 0 + assert stats.failed_count == stats.dropped_count == stats.pending_bytes == 0 assert len(sample_keys) == 2 assert eos_tags == [ { "record_type": "control", "status": "eos", - "schema_version": 1, + "schema_version": PROTOCOL_SCHEMA_VERSION, "run_id": "run-a", "total_samples": 2, } @@ -241,7 +249,7 @@ def test_run_producer_publishes_samples_then_eos(tmp_path: Path) -> None: assert all(not path.exists() for path in pool.paths) assert pool.started and pool.closed and transport.closed first_fields = transport.payloads[sorted(sample_keys)[0]] - assert tuple(first_fields["hidden_states"].shape) == (3, 6) + assert tuple(first_fields["sample__hidden_states"].shape) == (3, 6) def test_run_producer_restarts_input_until_max_samples(tmp_path: Path) -> None: @@ -308,8 +316,36 @@ def test_run_producer_generates_response_for_verl_chat_prompt(tmp_path: Path) -> assert pool.prefill_calls == 1 assert len(sample_keys) == 1 fields = transport.payloads[sample_keys[0]] - assert fields["input_ids"].tolist() == [10, 11] - assert fields["loss_mask"].tolist() == [0.0, 1.0] + assert fields["sample__input_ids"].tolist() == [10, 11] + assert fields["sample__loss_mask"].tolist() == [0.0, 1.0] + + +def test_run_producer_drops_misaligned_hidden_states_and_writes_eos( + tmp_path: Path, +) -> None: + input_path = tmp_path / "input.jsonl" + _write_input(input_path) + transport = _Transport() + pool = _MisalignedPool(tmp_path) + + stats = asyncio.run( + run_producer( + _config(input_path), + transport=transport, + tokenizer=_Tokenizer(), + client_pool=pool, + ) + ) + + assert stats.input_count == 2 + assert stats.published_count == 0 + assert stats.dropped_count == 2 + assert not any( + tag.get("record_type") == "sample" for tag in transport.records.values() + ) + eos = next(tag for tag in transport.records.values() if tag.get("status") == "eos") + assert eos["total_samples"] == 0 + assert all(not path.exists() for path in pool.paths) def test_run_producer_put_failure_keeps_temporary_file_and_omits_eos( diff --git a/tools/tq_connection_smoke.py b/tools/tq_connection_smoke.py index d1878681..ec2869be 100644 --- a/tools/tq_connection_smoke.py +++ b/tools/tq_connection_smoke.py @@ -36,7 +36,8 @@ from verl_speco.trainer.feature_store import DraftFeatureSample from verl_speco.trainer.tq_feature_store import TQFeatureStore from verl_speco.trainer.tq_sample_source import TQFeatureDataLoader -from verl_speco.transport.drafter_sample_protocol import ( +from verl_speco.transport.drafter_sample_protocol import ( + PROTOCOL_SCHEMA_VERSION, SampleMetadata, encode_sample, make_eos_record, @@ -52,7 +53,7 @@ def _config(args) -> dict: "ray": {"address": args.ray_address, "namespace": args.namespace}, "partition_id": "speco_drafter_features", "run_id": args.run_id, - "schema_version": 1, + "schema_version": PROTOCOL_SCHEMA_VERSION, "controller": {"polling_mode": True}, "backend": { "storage_backend": "SimpleStorage", @@ -66,23 +67,10 @@ def _config(args) -> dict: def _record(run_id: str, sequence_no: int): meta = SampleMetadata( - schema_version=1, - run_id=run_id, - sample_id=f"smoke-{sequence_no:04d}", - sequence_no=sequence_no, - algorithm="DSPARK", - target_model_id="smoke-target", - target_model_revision="smoke-revision", - tokenizer_fingerprint="smoke-tokenizer", - target_layer_ids=[0], - hidden_states_layout="dflash_aux", - hidden_dtype="float32", - hidden_shape=[3, 4], - feature_length=3, - full_sequence_length=3, - feature_start=0, - feature_end=3, - use_logits=False, + schema_version=PROTOCOL_SCHEMA_VERSION, + run_id=run_id, + sample_id=f"smoke-{sequence_no:04d}", + sequence_no=sequence_no, ) sample = DraftFeatureSample( algorithm="DSPARK", @@ -114,14 +102,14 @@ def run_owner(args) -> None: records = [_record(args.run_id, sequence_no) for sequence_no in range(2)] keys = [make_sample_key(meta) for meta, _ in records] try: - owner_ready_key = f"control:v1:{args.run_id}:owner-ready" + owner_ready_key = f"control:v{PROTOCOL_SCHEMA_VERSION}:{args.run_id}:owner-ready" put_sample( owner_ready_key, {"marker": torch.tensor([1], dtype=torch.uint8)}, tag={ "record_type": "control", "status": "owner_ready", - "schema_version": 1, + "schema_version": PROTOCOL_SCHEMA_VERSION, "run_id": args.run_id, }, ) diff --git a/tools/tq_delayed_test_producer.py b/tools/tq_delayed_test_producer.py index 6c89a500..60f98b3c 100644 --- a/tools/tq_delayed_test_producer.py +++ b/tools/tq_delayed_test_producer.py @@ -32,6 +32,7 @@ ) from verl_speco.trainer.feature_store import DraftFeatureSample from verl_speco.transport.drafter_sample_protocol import ( + PROTOCOL_SCHEMA_VERSION, SampleMetadata, encode_sample, make_eos_record, @@ -47,7 +48,7 @@ def _tq_config(args: argparse.Namespace) -> dict[str, Any]: "ray": {"address": args.ray_address, "namespace": args.namespace}, "partition_id": args.partition_id, "run_id": args.run_id, - "schema_version": 1, + "schema_version": PROTOCOL_SCHEMA_VERSION, "controller": {"polling_mode": True}, "backend": { "storage_backend": "SimpleStorage", @@ -71,7 +72,7 @@ def _model_dimensions(model_path: str) -> tuple[int, int, int]: def _wait_for_owner(args: argparse.Namespace) -> None: - owner_ready_key = f"control:v1:{args.run_id}:owner-ready" + owner_ready_key = f"control:v{PROTOCOL_SCHEMA_VERSION}:{args.run_id}:owner-ready" deadline = time.monotonic() + args.timeout while time.monotonic() < deadline: tag = list_samples().get(owner_ready_key) @@ -128,23 +129,10 @@ def _sample( ) sample_id = f"consumer-test-{sequence_no:08d}" metadata = SampleMetadata( - schema_version=1, + schema_version=PROTOCOL_SCHEMA_VERSION, run_id=args.run_id, sample_id=sample_id, sequence_no=sequence_no, - algorithm="DSPARK", - target_model_id=str(Path(args.model_path).resolve()), - target_model_revision="local-test", - tokenizer_fingerprint="synthetic-consumer-test", - target_layer_ids=target_layer_ids, - hidden_states_layout="dflash_aux_plus_last", - hidden_dtype="bfloat16", - hidden_shape=[args.sequence_length, hidden_dim], - feature_length=args.sequence_length, - full_sequence_length=args.sequence_length, - feature_start=0, - feature_end=args.sequence_length, - use_logits=False, ) sample = DraftFeatureSample( algorithm="DSPARK", diff --git a/verl_speco/config/speco_base.yaml b/verl_speco/config/speco_base.yaml index 196e985e..94588d20 100644 --- a/verl_speco/config/speco_base.yaml +++ b/verl_speco/config/speco_base.yaml @@ -233,7 +233,7 @@ actor_rollout_ref: namespace: speco-drafter partition_id: speco_drafter_features run_id: null - schema_version: 1 + schema_version: 2 connect_timeout_seconds: 120 poll_interval_seconds: 0.5 drop_last: true diff --git a/verl_speco/standalone_tq_producer.py b/verl_speco/standalone_tq_producer.py index 6f3c4947..d749d8b9 100644 --- a/verl_speco/standalone_tq_producer.py +++ b/verl_speco/standalone_tq_producer.py @@ -43,6 +43,7 @@ from verl_speco.trainer.feature_store import DraftFeatureSample from verl_speco.trainer.target_feature_replay import ( FeatureContract, + HiddenStateAlignmentError, feature_from_vllm_payload, ) from verl_speco.transport.drafter_sample_protocol import ( @@ -50,6 +51,7 @@ PROTOCOL_SCHEMA_VERSION, SampleMetadata, encode_sample, + is_ready_sample_tag, make_eos_record, make_ready_tag, make_sample_key, @@ -66,6 +68,7 @@ class ProducerStats: input_count: int = 0 published_count: int = 0 failed_count: int = 0 + dropped_count: int = 0 pending_bytes: int = 0 @@ -111,8 +114,9 @@ def validate_producer_config(config: Any) -> None: raise ValueError( "standalone_tq_producer.target_layer_ids must be a non-empty list" ) - if str(training_cfg.get("speculative_algorithm", "DSPARK")).upper() != "DSPARK": - raise ValueError("Standalone TQ Producer currently supports only DSPARK") + algorithm = str(training_cfg.get("speculative_algorithm", "") or "").strip() + if not algorithm: + raise ValueError("drafter.speculative_algorithm must not be empty") if bool(training_cfg.get("use_logits", False)): raise ValueError("Standalone TQ Producer does not support use_logits=true") if int(tq_cfg.get("schema_version", 0)) != PROTOCOL_SCHEMA_VERSION: @@ -216,17 +220,19 @@ async def run_producer( await pool.start() logger.info("Standalone TQ Producer vLLM client pool started") + algorithm = str(drafter_cfg["speculative_algorithm"]).strip().upper() feature_contract = FeatureContract( - algorithm="DSPARK", + algorithm=algorithm, target_layer_ids=[int(value) for value in producer_cfg["target_layer_ids"]], hidden_states_layout=resolve_drafter_hidden_states_layout( - "DSPARK", drafter_cfg + algorithm, drafter_cfg ), dtype=_parse_dtype(producer_cfg["hidden_dtype"]), target_model_id=str(producer_cfg["target_model_id"]), target_model_revision=str(producer_cfg["target_model_revision"]), tokenizer_fingerprint=str(producer_cfg["tokenizer_fingerprint"]), use_logits=False, + require_full_alignment=True, ) worker_count = int(producer_cfg["max_inflight_requests"]) input_queue: asyncio.Queue[Any] = asyncio.Queue( @@ -327,7 +333,25 @@ async def request_worker() -> None: else: raw = await pool.prefill(request) stats.pending_bytes += int(raw.byte_size) - sample = feature_from_vllm_payload(raw, request, feature_contract) + try: + sample = feature_from_vllm_payload( + raw, request, feature_contract + ) + except HiddenStateAlignmentError as exc: + stats.dropped_count += 1 + stats.pending_bytes = max( + stats.pending_bytes - int(raw.byte_size), 0 + ) + await asyncio.to_thread(delete_temporary_result, raw) + logger.warning( + "Standalone TQ Producer dropped misaligned sample " + "sequence_no=%s sample_id=%s dropped=%s reason=%s", + request.sequence_no, + request.sample_id, + stats.dropped_count, + exc, + ) + continue await publish_queue.put( PreparedFeature( request=request, @@ -378,9 +402,10 @@ async def publish_results() -> None: eos_key, eos_fields, eos_tag = make_eos_record(run_id, stats.published_count) await asyncio.to_thread(transport.put_sample, eos_key, eos_fields, tag=eos_tag) logger.info( - "Standalone TQ Producer completed inputs=%s published=%s", + "Standalone TQ Producer completed inputs=%s published=%s dropped=%s", stats.input_count, stats.published_count, + stats.dropped_count, ) return stats finally: @@ -428,9 +453,11 @@ async def _wait_for_pending_capacity( ready_count = sum( 1 for tag in records.values() - if tag.get("record_type") == "sample" - and tag.get("status") == "ready" - and tag.get("run_id") == run_id + if is_ready_sample_tag( + tag, + run_id=run_id, + schema_version=PROTOCOL_SCHEMA_VERSION, + ) ) if ready_count < max_pending_samples: return @@ -444,32 +471,12 @@ def _sample_metadata( run_id: str, tq_cfg: Mapping[str, Any], ) -> SampleMetadata: - hidden = sample.hidden_states - if not torch.is_tensor(hidden): - raise TypeError("Standalone TQ Producer requires dense hidden_states") - feature_start = int(sample.metadata["feature_start"]) - feature_end = int(sample.metadata["feature_end"]) - wire_layer_ids = list(contract.target_layer_ids) - if contract.hidden_states_layout == "dflash_aux_plus_last": - wire_layer_ids.append(-1) + del sample, contract return SampleMetadata( schema_version=int(tq_cfg["schema_version"]), run_id=run_id, sample_id=request.sample_id, sequence_no=request.sequence_no, - algorithm=contract.algorithm, - target_model_id=contract.target_model_id, - target_model_revision=str(contract.target_model_revision or ""), - tokenizer_fingerprint=contract.tokenizer_fingerprint, - target_layer_ids=wire_layer_ids, - hidden_states_layout=contract.hidden_states_layout, - hidden_dtype=str(hidden.dtype).removeprefix("torch."), - hidden_shape=[int(value) for value in hidden.shape], - feature_length=int(hidden.size(0)), - full_sequence_length=int(request.input_ids.numel()), - feature_start=feature_start, - feature_end=feature_end, - use_logits=contract.use_logits, ) diff --git a/verl_speco/standalone_tq_training_launcher.py b/verl_speco/standalone_tq_training_launcher.py index 95e1ae8e..879bb6e1 100644 --- a/verl_speco/standalone_tq_training_launcher.py +++ b/verl_speco/standalone_tq_training_launcher.py @@ -49,9 +49,12 @@ _TOKENIZER_PATH_KEY = ( "actor_rollout_ref.rollout.drafter.training.feature_store.tokenizer_path" ) -_DSPARK_LAYER_IDS_KEY = ( - "actor_rollout_ref.rollout.drafter.training.dspark_target_layer_ids" -) +_PRODUCER_TARGET_LAYER_IDS_KEY = "speco.standalone_tq_producer.target_layer_ids" +_ALGORITHM_TARGET_LAYER_IDS_KEYS = { + "DFLASH": "actor_rollout_ref.rollout.drafter.training.dflash_target_layer_ids", + "DSPARK": "actor_rollout_ref.rollout.drafter.training.dspark_target_layer_ids", + "DOMINO": "actor_rollout_ref.rollout.drafter.training.domino_target_layer_ids", +} _MAX_STEPS_KEY = "actor_rollout_ref.rollout.drafter.training.max_steps" _BATCH_SIZE_PER_GPU_KEY = ( "actor_rollout_ref.rollout.drafter.training.batch_size_per_gpu" @@ -185,21 +188,36 @@ def _single_train_file(value: str | None) -> str: return text -def _parse_layer_ids(value: str | None) -> tuple[int, ...]: +def _parse_layer_ids(value: str | None, *, config_key: str) -> tuple[int, ...]: if value is None or _strip_quotes(value).lower() in {"", "null", "none"}: return _DEFAULT_TARGET_LAYER_IDS text = _strip_quotes(value) if not (text.startswith("[") and text.endswith("]")): - raise ValueError(f"{_DSPARK_LAYER_IDS_KEY} must be a Hydra integer list") + raise ValueError(f"{config_key} must be a Hydra integer list") try: result = tuple(int(item.strip()) for item in text[1:-1].split(",")) except ValueError as exc: - raise ValueError(f"{_DSPARK_LAYER_IDS_KEY} must contain only integers") from exc + raise ValueError(f"{config_key} must contain only integers") from exc if not result or any(value < 0 for value in result): - raise ValueError(f"{_DSPARK_LAYER_IDS_KEY} must contain non-negative IDs") + raise ValueError(f"{config_key} must contain non-negative IDs") return result +def _resolve_target_layer_ids( + training_args: Sequence[str], algorithm: str +) -> tuple[int, ...]: + """Resolve Producer layers without making the launcher DSpark-specific.""" + + algorithm_key = _ALGORITHM_TARGET_LAYER_IDS_KEYS.get(algorithm) + candidate_keys = ( + (_PRODUCER_TARGET_LAYER_IDS_KEY, algorithm_key) + if algorithm_key is not None + else (_PRODUCER_TARGET_LAYER_IDS_KEY,) + ) + raw = _find_first_override(training_args, candidate_keys) + return _parse_layer_ids(raw, config_key=candidate_keys[0]) + + def _parse_vllm_endpoints(env: Mapping[str, str]) -> tuple[str, ...]: """Read a Hydra-style endpoint list while preserving the singular fallback.""" @@ -307,11 +325,9 @@ def resolve_pipeline_config( algorithm = _strip_quotes( _find_override(training_args, _ALGORITHM_KEY) or "DSPARK" ).upper() - if algorithm != "DSPARK": - raise ValueError("Standalone Producer/TQ/Consumer launcher requires DSPARK") - target_layer_ids = _parse_layer_ids( - _find_override(training_args, _DSPARK_LAYER_IDS_KEY) - ) + if not algorithm: + raise ValueError(f"{_ALGORITHM_KEY} must not be empty") + target_layer_ids = _resolve_target_layer_ids(training_args, algorithm) endpoints = _parse_vllm_endpoints(env) return PipelineConfig( input_path=input_path, @@ -320,7 +336,7 @@ def resolve_pipeline_config( algorithm=algorithm, target_layer_ids=target_layer_ids, vllm_endpoints=endpoints, - run_id=f"dspark-{uuid.uuid4().hex}", + run_id=f"{algorithm.lower()}-{uuid.uuid4().hex}", ) @@ -486,8 +502,12 @@ def build_pipeline_commands( f"{_FEATURE_STORE_PREFIX}.shuffle=false", f"{_FEATURE_STORE_PREFIX}.repeat=false", *tq_overrides, - f"{_DSPARK_LAYER_IDS_KEY}={_hydra_list(config.target_layer_ids)}", ] + algorithm_layer_ids_key = _ALGORITHM_TARGET_LAYER_IDS_KEYS.get(config.algorithm) + if algorithm_layer_ids_key is not None: + consumer_internal.append( + f"{algorithm_layer_ids_key}={_hydra_list(config.target_layer_ids)}" + ) consumer = [ python_executable, "-m", diff --git a/verl_speco/tq_owner.py b/verl_speco/tq_owner.py index cae5552b..c9128848 100644 --- a/verl_speco/tq_owner.py +++ b/verl_speco/tq_owner.py @@ -33,6 +33,7 @@ put_sample, start_transfer_queue_owner, ) +from verl_speco.transport.drafter_sample_protocol import PROTOCOL_SCHEMA_VERSION logger = logging.getLogger(__name__) @@ -97,7 +98,7 @@ def run_owner(config: Any, *, stop_event: threading.Event | None = None) -> int: started = True ready_key = publish_owner_ready( str(tq_cfg.get("run_id") or ""), - int(tq_cfg.get("schema_version", 1)), + int(tq_cfg.get("schema_version", PROTOCOL_SCHEMA_VERSION)), ) logger.info("TQ owner ready key=%s", ready_key) ready_file = os.environ.get("SPECO_TQ_OWNER_READY_FILE") diff --git a/verl_speco/trainer/target_feature_replay.py b/verl_speco/trainer/target_feature_replay.py index 9e3d50e5..3f020d5b 100644 --- a/verl_speco/trainer/target_feature_replay.py +++ b/verl_speco/trainer/target_feature_replay.py @@ -38,6 +38,10 @@ logger = logging.getLogger(__name__) +class HiddenStateAlignmentError(ValueError): + """The returned hidden states cannot cover the requested training sample.""" + + @dataclass(frozen=True) class MooncakeReplayDescriptor: sample: DraftReplaySample @@ -59,6 +63,7 @@ class FeatureContract: use_logits: bool = False target_config_fingerprint: str | None = None source: str = "standalone_tq_producer" + require_full_alignment: bool = False @dataclass @@ -397,7 +402,9 @@ def feature_from_vllm_payload( else request.input_ids[:feature_end_for_request].detach().cpu().long().tolist() ) if token_ids.detach().cpu().long().tolist() != expected_prompt_ids: - raise ValueError("vLLM hidden-states token_ids do not match replay input") + raise HiddenStateAlignmentError( + "vLLM hidden-states token_ids do not match replay input" + ) if hidden.dim() != 3: raise ValueError( "vLLM hidden_states must have shape [seq, layers, hidden], " @@ -425,7 +432,7 @@ def feature_from_vllm_payload( } required_layers = len(target_layer_ids) + (1 if include_final else 0) if int(hidden.size(1)) < required_layers: - raise ValueError( + raise HiddenStateAlignmentError( "vLLM hidden_states layer count is too small: " f"got {int(hidden.size(1))}, need at least {required_layers}. " "Start vLLM with target layer ids plus the final layer when the " @@ -435,6 +442,15 @@ def feature_from_vllm_payload( keep_mask = (relative_positions >= 0) & (relative_positions < int(hidden.size(0))) filtered = not bool(keep_mask.all().item()) if filtered: + if feature_config.require_full_alignment: + raise HiddenStateAlignmentError( + "vLLM hidden-state rows do not cover the complete feature window: " + f"hidden_rows={int(hidden.size(0))}, " + f"hidden_position_offset={hidden_position_offset}, " + f"dropped={int((~keep_mask).sum().item())}, " + f"feature_min={int(feature_positions.min().item())}, " + f"feature_max={int(feature_positions.max().item())}" + ) logger.warning( "Dropping vLLM feature positions outside hidden rows dropped=%s " "hidden_rows=%s hidden_offset=%s feature_min=%s feature_max=%s", @@ -447,7 +463,7 @@ def feature_from_vllm_payload( feature_positions = feature_positions[keep_mask] relative_positions = relative_positions[keep_mask] if int(feature_positions.numel()) <= 0: - raise ValueError( + raise HiddenStateAlignmentError( "vLLM hidden_states contain no rows for replay feature positions: " f"hidden_rows={int(hidden.size(0))}, " f"hidden_position_offset={hidden_position_offset}" diff --git a/verl_speco/trainer/tq_feature_store.py b/verl_speco/trainer/tq_feature_store.py index e59ccc4c..ff5ae570 100644 --- a/verl_speco/trainer/tq_feature_store.py +++ b/verl_speco/trainer/tq_feature_store.py @@ -31,6 +31,7 @@ from verl_speco.transport.drafter_sample_protocol import ( ExpectedFeatureConfig, decode_sample, + parse_ready_tag, ) @@ -101,16 +102,12 @@ def list_ready(self, run_id: str | None = None) -> list[ReadyEntry]: ready: list[ReadyEntry] = [] for key, raw_tag in list_samples().items(): tag = dict(raw_tag) - if tag.get("record_type") != "sample" or tag.get("status") != "ready": - continue - if str(tag.get("run_id") or "") != expected_run_id: - continue - if int(tag.get("schema_version", -1)) != self.schema_version: - continue - try: - int(tag["sequence_no"]) - str(tag["sample_id"]) - except (KeyError, TypeError, ValueError): + meta = parse_ready_tag( + tag, + run_id=expected_run_id, + schema_version=self.schema_version, + ) + if meta is None: continue ready.append(ReadyEntry(key=str(key), tag=tag)) ready.sort(key=lambda entry: (int(entry.tag["sequence_no"]), entry.key)) diff --git a/verl_speco/transport/__init__.py b/verl_speco/transport/__init__.py index 4acd2a06..6f8b6dd2 100644 --- a/verl_speco/transport/__init__.py +++ b/verl_speco/transport/__init__.py @@ -20,9 +20,11 @@ SampleMetadata, decode_sample, encode_sample, + is_ready_sample_tag, make_eos_record, make_ready_tag, make_sample_key, + parse_ready_tag, ) __all__ = [ @@ -32,7 +34,9 @@ "SampleMetadata", "decode_sample", "encode_sample", + "is_ready_sample_tag", "make_eos_record", "make_ready_tag", "make_sample_key", + "parse_ready_tag", ] diff --git a/verl_speco/transport/drafter_sample_protocol.py b/verl_speco/transport/drafter_sample_protocol.py index 174c03cf..5ffa1ac4 100644 --- a/verl_speco/transport/drafter_sample_protocol.py +++ b/verl_speco/transport/drafter_sample_protocol.py @@ -11,12 +11,7 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. -"""Wire protocol for standalone drafter samples stored in TransferQueue. - -One TQ key represents one training sample. Tensor payloads live in TQ fields; -small discovery attributes live in the TQ tag; richer metadata is JSON encoded -as a uint8 tensor so Producer and Consumer use one versioned contract. -""" +"""Algorithm-neutral TransferQueue codec for ``DraftFeatureSample``.""" from __future__ import annotations @@ -29,46 +24,31 @@ from verl_speco.trainer.feature_store import DraftFeatureSample -PROTOCOL_SCHEMA_VERSION = 1 +PROTOCOL_SCHEMA_VERSION = 2 DRAFTER_TQ_PARTITION = "speco_drafter_features" -_REQUIRED_FIELDS = ( - "input_ids", - "loss_mask", - "position_ids", - "hidden_states", - "metadata_json", -) +_MANIFEST_FIELD = "sample__manifest_json" +_REQUIRED_SAMPLE_FIELDS = ("input_ids", "loss_mask", "hidden_states") _OPTIONAL_TENSOR_FIELDS = ( "last_hidden_states", "target", "target_logprobs", + "position_ids", ) @dataclass(frozen=True) class SampleMetadata: + """Small control-plane envelope; training metadata belongs to the sample.""" + schema_version: int run_id: str sample_id: str sequence_no: int - algorithm: str - target_model_id: str - target_model_revision: str - tokenizer_fingerprint: str - target_layer_ids: list[int] - hidden_states_layout: str - hidden_dtype: str - hidden_shape: list[int] - feature_length: int - full_sequence_length: int - feature_start: int - feature_end: int - use_logits: bool def validate(self) -> None: if self.schema_version != PROTOCOL_SCHEMA_VERSION: raise ValueError( - f"Unsupported drafter sample schema_version={self.schema_version}; " + f"Unsupported drafter protocol schema_version={self.schema_version}; " f"expected {PROTOCOL_SCHEMA_VERSION}" ) if not self.run_id: @@ -77,35 +57,10 @@ def validate(self) -> None: raise ValueError("SampleMetadata.sample_id must not be empty") if self.sequence_no < 0: raise ValueError("SampleMetadata.sequence_no must be non-negative") - if not self.algorithm.strip(): - raise ValueError("SampleMetadata.algorithm must not be empty") - if len(self.hidden_shape) != 2: - raise ValueError( - f"SampleMetadata.hidden_shape must be [rows, hidden_dim], got {self.hidden_shape!r}" - ) - if self.feature_length <= 0: - raise ValueError("SampleMetadata.feature_length must be positive") - if self.hidden_shape[0] != self.feature_length: - raise ValueError( - "SampleMetadata hidden_shape/feature_length mismatch: " - f"{self.hidden_shape[0]} vs {self.feature_length}" - ) - if not (0 <= self.feature_start < self.feature_end <= self.full_sequence_length): - raise ValueError( - "SampleMetadata feature window must satisfy " - "0 <= feature_start < feature_end <= full_sequence_length" - ) - if self.feature_end - self.feature_start != self.feature_length: - raise ValueError( - "SampleMetadata feature window length does not match feature_length: " - f"{self.feature_end - self.feature_start} vs {self.feature_length}" - ) def to_dict(self) -> dict[str, Any]: self.validate() - payload = asdict(self) - payload["algorithm"] = self.algorithm.strip().upper() - return payload + return asdict(self) @classmethod def from_dict(cls, payload: Mapping[str, Any]) -> "SampleMetadata": @@ -115,39 +70,19 @@ def from_dict(cls, payload: Mapping[str, Any]) -> "SampleMetadata": run_id=str(payload["run_id"]), sample_id=str(payload["sample_id"]), sequence_no=int(payload["sequence_no"]), - algorithm=str(payload["algorithm"]), - target_model_id=str(payload["target_model_id"]), - target_model_revision=str(payload["target_model_revision"]), - tokenizer_fingerprint=str(payload["tokenizer_fingerprint"]), - target_layer_ids=[int(v) for v in payload["target_layer_ids"]], - hidden_states_layout=str(payload["hidden_states_layout"]), - hidden_dtype=str(payload["hidden_dtype"]), - hidden_shape=[int(v) for v in payload["hidden_shape"]], - feature_length=int(payload["feature_length"]), - full_sequence_length=int(payload["full_sequence_length"]), - feature_start=int(payload["feature_start"]), - feature_end=int(payload["feature_end"]), - use_logits=bool(payload["use_logits"]), ) except KeyError as exc: - raise ValueError(f"metadata_json missing required field {exc.args[0]!r}") from exc + raise ValueError(f"ready tag missing required field {exc.args[0]!r}") from exc meta.validate() return meta @dataclass(frozen=True) class ExpectedFeatureConfig: - """Consumer-side contract. ``None`` fields are intentionally unchecked.""" + """Consumer-side transport contract.""" run_id: str schema_version: int = PROTOCOL_SCHEMA_VERSION - algorithm: str | None = None - target_model_id: str | None = None - target_model_revision: str | None = None - tokenizer_fingerprint: str | None = None - target_layer_ids: list[int] | None = None - hidden_states_layout: str | None = None - hidden_dtype: str | None = None def make_sample_key(meta: SampleMetadata) -> str: @@ -160,21 +95,45 @@ def make_sample_key(meta: SampleMetadata) -> str: def make_ready_tag(meta: SampleMetadata) -> dict[str, Any]: meta.validate() - return { - "record_type": "sample", - "status": "ready", - "schema_version": meta.schema_version, - "run_id": meta.run_id, - "sequence_no": meta.sequence_no, - "sample_id": meta.sample_id, - "algorithm": meta.algorithm.strip().upper(), - } + return {"record_type": "sample", "status": "ready", **meta.to_dict()} + + +def parse_ready_tag( + tag: Mapping[str, Any], + *, + run_id: str | None = None, + schema_version: int = PROTOCOL_SCHEMA_VERSION, +) -> SampleMetadata | None: + """Return one valid ready envelope, or ``None`` when it is not consumable.""" + + if tag.get("record_type") != "sample" or tag.get("status") != "ready": + return None + try: + meta = SampleMetadata.from_dict(tag) + except (TypeError, ValueError): + return None + if meta.schema_version != int(schema_version): + return None + if run_id is not None and meta.run_id != str(run_id): + return None + return meta + + +def is_ready_sample_tag( + tag: Mapping[str, Any], + *, + run_id: str, + schema_version: int = PROTOCOL_SCHEMA_VERSION, +) -> bool: + return parse_ready_tag( + tag, run_id=run_id, schema_version=schema_version + ) is not None def encode_sample( sample: DraftFeatureSample | Mapping[str, Any], meta: SampleMetadata ) -> dict[str, torch.Tensor]: - """Encode one normalized feature sample into TQ tensor fields.""" + """Losslessly encode one normalized feature sample into TQ tensor fields.""" meta.validate() normalized = ( @@ -183,29 +142,48 @@ def encode_sample( else DraftFeatureSample.from_dict(dict(sample), strict=True) ) normalized.validate(strict=True) - if isinstance(normalized.hidden_states, (list, tuple)): - raise TypeError("TQ drafter protocol requires hidden_states to be one dense tensor") - hidden = _cpu_contiguous(normalized.hidden_states) - input_ids = _cpu_contiguous(normalized.input_ids, dtype=torch.int64).reshape(-1) - loss_mask = _cpu_contiguous(normalized.loss_mask, dtype=torch.float32).reshape(-1) - if normalized.position_ids is None: - position_ids = torch.arange(input_ids.numel(), dtype=torch.int64) - else: - position_ids = _cpu_contiguous(normalized.position_ids, dtype=torch.int64).reshape(-1) - - _validate_primary_tensors(input_ids, loss_mask, position_ids, hidden, meta) - metadata_json = _json_to_tensor(meta.to_dict()) - fields: dict[str, torch.Tensor] = { - "input_ids": input_ids, - "loss_mask": loss_mask, - "position_ids": position_ids, - "hidden_states": hidden, - "metadata_json": metadata_json, + payload = normalized.to_dict() + fields: dict[str, torch.Tensor] = {} + manifest: dict[str, Any] = { + "draft_feature_schema_version": int(normalized.schema_version), + "algorithm": str(normalized.algorithm), + "present_fields": [], } - for field_name in _OPTIONAL_TENSOR_FIELDS: - value = getattr(normalized, field_name) - if value is not None: + + for name in ("input_ids", "loss_mask", *_OPTIONAL_TENSOR_FIELDS): + value = payload.get(name) + if value is None: + continue + if not torch.is_tensor(value): + raise TypeError(f"DraftFeatureSample.{name} must be a torch.Tensor") + fields[f"sample__{name}"] = _cpu_contiguous(value) + manifest["present_fields"].append(name) + + hidden = payload["hidden_states"] + if torch.is_tensor(hidden): + fields["sample__hidden_states"] = _cpu_contiguous(hidden) + manifest["hidden_states_kind"] = "tensor" + elif isinstance(hidden, (list, tuple)): + hidden_fields: list[str] = [] + for index, value in enumerate(hidden): + if not torch.is_tensor(value): + raise TypeError( + f"DraftFeatureSample.hidden_states[{index}] must be a torch.Tensor" + ) + field_name = f"sample__hidden_states__{index:06d}" fields[field_name] = _cpu_contiguous(value) + hidden_fields.append(field_name) + if not hidden_fields: + raise ValueError("DraftFeatureSample.hidden_states list must not be empty") + manifest["hidden_states_kind"] = "list" + manifest["hidden_states_fields"] = hidden_fields + else: + raise TypeError("DraftFeatureSample.hidden_states must be a tensor or tensor list") + manifest["present_fields"].append("hidden_states") + manifest["metadata"] = _encode_metadata_tree( + payload.get("metadata", {}), fields, path="metadata" + ) + fields[_MANIFEST_FIELD] = _json_to_tensor(manifest) return fields @@ -215,41 +193,47 @@ def decode_sample( fields: Mapping[str, Any], expected_config: ExpectedFeatureConfig | Mapping[str, Any], ) -> DraftFeatureSample: - """Validate a TQ record and restore the existing training sample type.""" + """Validate one queue record and restore the complete training sample.""" expected = ( expected_config if isinstance(expected_config, ExpectedFeatureConfig) else ExpectedFeatureConfig(**dict(expected_config)) ) - missing = [name for name in _REQUIRED_FIELDS if name not in fields] - if missing: - raise ValueError(f"TQ sample {key!r} missing required fields: {missing}") - metadata = SampleMetadata.from_dict(_tensor_to_json(fields["metadata_json"])) - expected_key = make_sample_key(metadata) + meta = parse_ready_tag( + tag, run_id=expected.run_id, schema_version=expected.schema_version + ) + if meta is None: + raise ValueError(f"TQ sample {key!r} has an invalid or unexpected ready tag") + expected_key = make_sample_key(meta) if key != expected_key: raise ValueError(f"TQ sample key mismatch: got {key!r}, expected {expected_key!r}") - _validate_tag(tag, metadata) - _validate_expected(metadata, expected) - - input_ids = _require_tensor(fields, "input_ids").detach().cpu().to(torch.int64).reshape(-1) - loss_mask = _require_tensor(fields, "loss_mask").detach().cpu().to(torch.float32).reshape(-1) - position_ids = _require_tensor(fields, "position_ids").detach().cpu().to(torch.int64).reshape(-1) - hidden = _require_tensor(fields, "hidden_states").detach().cpu().contiguous() - _validate_primary_tensors(input_ids, loss_mask, position_ids, hidden, metadata) + if _MANIFEST_FIELD not in fields: + raise ValueError(f"TQ sample {key!r} missing field {_MANIFEST_FIELD!r}") + manifest = _tensor_to_json(fields[_MANIFEST_FIELD], name=_MANIFEST_FIELD) + present = manifest.get("present_fields") + if not isinstance(present, list): + raise ValueError("sample manifest present_fields must be a list") + missing = [name for name in _REQUIRED_SAMPLE_FIELDS if name not in present] + if missing: + raise ValueError(f"TQ sample {key!r} missing required sample fields: {missing}") + try: + sample_schema_version = int(manifest["draft_feature_schema_version"]) + algorithm = str(manifest["algorithm"]) + except KeyError as exc: + raise ValueError(f"sample manifest missing field {exc.args[0]!r}") from exc payload: dict[str, Any] = { - "schema_version": metadata.schema_version, - "algorithm": metadata.algorithm, - "input_ids": input_ids, - "loss_mask": loss_mask, - "position_ids": position_ids, - "hidden_states": hidden, - "metadata": metadata.to_dict(), + "schema_version": sample_schema_version, + "algorithm": algorithm, + "metadata": _decode_metadata_tree( + manifest.get("metadata", {}), fields, path="metadata" + ), } - for field_name in _OPTIONAL_TENSOR_FIELDS: - if field_name in fields and fields[field_name] is not None: - payload[field_name] = _require_tensor(fields, field_name).detach().cpu().contiguous() + for name in ("input_ids", "loss_mask", *_OPTIONAL_TENSOR_FIELDS): + if name in present: + payload[name] = _require_tensor(fields, f"sample__{name}") + payload["hidden_states"] = _decode_hidden_states(fields, manifest) return DraftFeatureSample.from_dict(payload, strict=True) @@ -272,75 +256,78 @@ def make_eos_record( return key, fields, tag -def _validate_tag(tag: Mapping[str, Any], meta: SampleMetadata) -> None: - expected = make_ready_tag(meta) - for name, expected_value in expected.items(): - if tag.get(name) != expected_value: - raise ValueError( - f"TQ sample tag mismatch for {name}: got {tag.get(name)!r}, " - f"expected {expected_value!r}" - ) - - -def _validate_expected(meta: SampleMetadata, expected: ExpectedFeatureConfig) -> None: - checks = { - "run_id": expected.run_id, - "schema_version": expected.schema_version, - "algorithm": ( - expected.algorithm.strip().upper() - if expected.algorithm is not None - else None - ), - "target_model_id": expected.target_model_id, - "target_model_revision": expected.target_model_revision, - "tokenizer_fingerprint": expected.tokenizer_fingerprint, - "target_layer_ids": expected.target_layer_ids, - "hidden_states_layout": expected.hidden_states_layout, - "hidden_dtype": expected.hidden_dtype, - } - for name, expected_value in checks.items(): - if expected_value is None: - continue - actual = getattr(meta, name) - if name == "algorithm": - actual = str(actual).strip().upper() - if actual != expected_value: - raise ValueError( - f"TQ sample metadata mismatch for {name}: got {actual!r}, " - f"expected {expected_value!r}" - ) +def _encode_metadata_tree( + value: Any, fields: dict[str, torch.Tensor], *, path: str +) -> Any: + if torch.is_tensor(value): + field_name = f"sample__metadata_tensor__{len(fields):06d}" + fields[field_name] = _cpu_contiguous(value) + return {"__tq_tensor_ref__": field_name} + if value is None or isinstance(value, (str, int, float, bool)): + return value + if isinstance(value, Mapping): + encoded: dict[str, Any] = {} + for key, item in value.items(): + if not isinstance(key, str): + raise TypeError(f"DraftFeatureSample metadata key at {path} must be str") + encoded[key] = _encode_metadata_tree(item, fields, path=f"{path}.{key}") + return {"__tq_mapping__": encoded} + if isinstance(value, (list, tuple)): + items = [ + _encode_metadata_tree(item, fields, path=f"{path}[{index}]") + for index, item in enumerate(value) + ] + return { + "__tq_sequence__": "tuple" if isinstance(value, tuple) else "list", + "items": items, + } + raise TypeError( + f"Unsupported DraftFeatureSample metadata value at {path}: {type(value).__name__}" + ) -def _validate_primary_tensors( - input_ids: torch.Tensor, - loss_mask: torch.Tensor, - position_ids: torch.Tensor, - hidden: torch.Tensor, - meta: SampleMetadata, -) -> None: - if hidden.dim() == 3 and hidden.size(0) == 1: - hidden = hidden.squeeze(0) - if hidden.dim() != 2: - raise ValueError(f"hidden_states must have shape [L,D], got {tuple(hidden.shape)}") - lengths = { - "input_ids": int(input_ids.numel()), - "loss_mask": int(loss_mask.numel()), - "position_ids": int(position_ids.numel()), - "hidden_states": int(hidden.size(0)), - } - if any(value != meta.feature_length for value in lengths.values()): - raise ValueError( - f"TQ sample tensor lengths must equal feature_length={meta.feature_length}: {lengths}" - ) - if list(hidden.shape) != meta.hidden_shape: - raise ValueError( - f"hidden_states shape mismatch: got {list(hidden.shape)}, expected {meta.hidden_shape}" - ) - actual_dtype = _dtype_name(hidden.dtype) - if actual_dtype != meta.hidden_dtype: - raise ValueError( - f"hidden_states dtype mismatch: got {actual_dtype!r}, expected {meta.hidden_dtype!r}" - ) +def _decode_metadata_tree(value: Any, fields: Mapping[str, Any], *, path: str) -> Any: + if value is None or isinstance(value, (str, int, float, bool)): + return value + if not isinstance(value, Mapping): + raise ValueError(f"Invalid metadata manifest node at {path}") + if "__tq_tensor_ref__" in value: + return _require_tensor(fields, str(value["__tq_tensor_ref__"])) + if "__tq_mapping__" in value: + mapping = value["__tq_mapping__"] + if not isinstance(mapping, Mapping): + raise ValueError(f"Invalid metadata mapping node at {path}") + return { + str(key): _decode_metadata_tree(item, fields, path=f"{path}.{key}") + for key, item in mapping.items() + } + if "__tq_sequence__" in value: + items = value.get("items") + if not isinstance(items, list): + raise ValueError(f"Invalid metadata sequence node at {path}") + decoded = [ + _decode_metadata_tree(item, fields, path=f"{path}[{index}]") + for index, item in enumerate(items) + ] + if value["__tq_sequence__"] == "tuple": + return tuple(decoded) + if value["__tq_sequence__"] == "list": + return decoded + raise ValueError(f"Unknown metadata manifest node at {path}") + + +def _decode_hidden_states( + fields: Mapping[str, Any], manifest: Mapping[str, Any] +) -> torch.Tensor | list[torch.Tensor]: + kind = manifest.get("hidden_states_kind") + if kind == "tensor": + return _require_tensor(fields, "sample__hidden_states") + if kind == "list": + names = manifest.get("hidden_states_fields") + if not isinstance(names, list) or not names: + raise ValueError("sample manifest hidden_states_fields must be non-empty") + return [_require_tensor(fields, str(name)) for name in names] + raise ValueError(f"Unsupported hidden_states_kind={kind!r}") def _json_to_tensor(payload: Mapping[str, Any]) -> torch.Tensor: @@ -348,17 +335,16 @@ def _json_to_tensor(payload: Mapping[str, Any]) -> torch.Tensor: return torch.tensor(list(raw), dtype=torch.uint8) -def _tensor_to_json(value: Any) -> dict[str, Any]: - tensor = value - if not torch.is_tensor(tensor): - raise TypeError("metadata_json must be a torch.Tensor") - tensor = tensor.detach().cpu().to(torch.uint8).reshape(-1) +def _tensor_to_json(value: Any, *, name: str) -> dict[str, Any]: + if not torch.is_tensor(value): + raise TypeError(f"{name} must be a torch.Tensor") + tensor = value.detach().cpu().to(torch.uint8).reshape(-1) try: decoded = json.loads(bytes(tensor.tolist()).decode("utf-8")) except (UnicodeDecodeError, json.JSONDecodeError) as exc: - raise ValueError("metadata_json is not valid UTF-8 JSON") from exc + raise ValueError(f"{name} is not valid UTF-8 JSON") from exc if not isinstance(decoded, dict): - raise ValueError("metadata_json must decode to a JSON object") + raise ValueError(f"{name} must decode to a JSON object") return decoded @@ -366,20 +352,13 @@ def _require_tensor(fields: Mapping[str, Any], name: str) -> torch.Tensor: value = fields.get(name) if not torch.is_tensor(value): raise TypeError(f"TQ field {name!r} must be a torch.Tensor") - return value + return value.detach().cpu().contiguous() -def _cpu_contiguous(value: torch.Tensor, *, dtype: torch.dtype | None = None) -> torch.Tensor: +def _cpu_contiguous(value: torch.Tensor) -> torch.Tensor: if not torch.is_tensor(value): raise TypeError(f"Expected torch.Tensor, got {type(value)!r}") - result = value.detach().cpu() - if dtype is not None: - result = result.to(dtype) - return result.contiguous() - - -def _dtype_name(dtype: torch.dtype) -> str: - return str(dtype).removeprefix("torch.") + return value.detach().cpu().contiguous() __all__ = [ @@ -389,7 +368,9 @@ def _dtype_name(dtype: torch.dtype) -> str: "SampleMetadata", "decode_sample", "encode_sample", + "is_ready_sample_tag", "make_eos_record", "make_ready_tag", "make_sample_key", + "parse_ready_tag", ] From 9b1c7ad647866c821adf9414ccdac19830b1409a Mon Sep 17 00:00:00 2001 From: 755651978 <755651978@qq.com> Date: Wed, 26 Aug 2026 15:11:44 +0800 Subject: [PATCH 39/50] Fix standalone EAGLE3 training and checkpoint export compatibility for target model configurations such as Qwen3. --- .../test_eagle3_aux_hidden_contract.py | 18 ++++++ tests/unit/test_draft_training_loop.py | 64 +++++++++++++++++++ verl_speco/backends/eagle3_trainer_backend.py | 2 + verl_speco/models/eagle/llama_eagle.py | 11 ++-- verl_speco/trainer/standalone_checkpoint.py | 28 +++++++- 5 files changed, 117 insertions(+), 6 deletions(-) diff --git a/tests/integration/test_eagle3_aux_hidden_contract.py b/tests/integration/test_eagle3_aux_hidden_contract.py index 3a3f3b18..f6c8474c 100644 --- a/tests/integration/test_eagle3_aux_hidden_contract.py +++ b/tests/integration/test_eagle3_aux_hidden_contract.py @@ -84,3 +84,21 @@ def test_eagle3_model_uses_dynamic_aux_hidden_count() -> None: with pytest.raises(ValueError, match="num_aux_hidden_states=5"): model.project_hidden_states(torch.randn(2, 3, 12)) + + +def test_eagle3_model_defaults_missing_pretraining_tp() -> None: + torch = pytest.importorskip("torch") + pytest.importorskip("transformers") + from verl_speco.models.auto import AutoDraftModelConfig + from verl_speco.models.eagle.llama_eagle import LlamaMLP + + raw_config = _minimal_eagle3_config() + raw_config.pop("pretraining_tp") + config = AutoDraftModelConfig._config_mapping["LlamaForCausalLMEagle3"].from_dict( + raw_config + ) + mlp = LlamaMLP(config) + + output = mlp(torch.randn(2, 3, config.hidden_size)) + + assert output.shape == (2, 3, config.hidden_size) diff --git a/tests/unit/test_draft_training_loop.py b/tests/unit/test_draft_training_loop.py index 684d0aec..3cd3deaf 100644 --- a/tests/unit/test_draft_training_loop.py +++ b/tests/unit/test_draft_training_loop.py @@ -564,6 +564,70 @@ def test_standalone_block_checkpoint_uses_target_model_type_without_source_confi assert saved_training_config == training_config +def test_standalone_eagle3_checkpoint_exports_vllm_llama_runtime_config(tmp_path): + checkpoint_dir = tmp_path / "draft_step_5" + checkpoint_dir.mkdir() + target_dir = tmp_path / "target_qwen3" + target_dir.mkdir() + missing_source_dir = tmp_path / "missing_source_eagle3" + (target_dir / "config.json").write_text( + json.dumps( + { + "model_type": "qwen3", + "hidden_size": 4096, + "head_dim": 128, + "rope_theta": 1000000, + "max_position_embeddings": 40960, + } + ), + encoding="utf-8", + ) + training_config = { + "model_type": "qwen3", + "architectures": ["LlamaForCausalLMEagle3"], + "num_hidden_layers": 1, + "hidden_size": 4096, + "vocab_size": 151936, + "tie_word_embeddings": False, + } + (checkpoint_dir / "config.json").write_text( + json.dumps(training_config), encoding="utf-8" + ) + trainer = SimpleNamespace( + backend=SimpleNamespace(model_type="eagle3"), + config=SimpleNamespace( + model=SimpleNamespace(path=str(target_dir)), + rollout=SimpleNamespace( + drafter=SimpleNamespace(model_path=str(missing_source_dir)) + ), + ), + ) + + source_model_path = rewrite_standalone_runtime_config(trainer, str(checkpoint_dir)) + + runtime_config = json.loads( + (checkpoint_dir / "config.json").read_text(encoding="utf-8") + ) + saved_training_config = json.loads( + (checkpoint_dir / "speco_training_config.json").read_text(encoding="utf-8") + ) + assert source_model_path == str(missing_source_dir) + assert runtime_config["model_type"] == "llama" + assert runtime_config["architectures"] == ["LlamaForCausalLMEagle3"] + assert runtime_config["draft_model_type"] == "eagle3" + assert runtime_config["speculative_algorithm"] == "EAGLE3" + assert runtime_config["speco_training_model_type"] == "eagle3" + assert runtime_config["pretraining_tp"] == 1 + assert runtime_config["num_hidden_layers"] == 1 + assert runtime_config["tie_word_embeddings"] is False + assert runtime_config["draft_vocab_size"] == 151936 + assert runtime_config["target_hidden_size"] == 4096 + assert runtime_config["head_dim"] == 128 + assert runtime_config["rope_theta"] == 1000000 + assert runtime_config["max_position_embeddings"] == 40960 + assert saved_training_config == training_config + + def test_standalone_dflash_checkpoint_preserves_source_lm_head(tmp_path): safetensors_torch = pytest.importorskip("safetensors.torch") checkpoint_dir = tmp_path / "draft_step_5" diff --git a/verl_speco/backends/eagle3_trainer_backend.py b/verl_speco/backends/eagle3_trainer_backend.py index 427677ea..edcbc6f4 100644 --- a/verl_speco/backends/eagle3_trainer_backend.py +++ b/verl_speco/backends/eagle3_trainer_backend.py @@ -689,6 +689,8 @@ def build_model(self): drafter_config.tie_word_embeddings = False drafter_config.architectures = ["LlamaForCausalLMEagle3"] + if not hasattr(drafter_config, "pretraining_tp"): + drafter_config.pretraining_tp = 1 if not hasattr(drafter_config, "draft_vocab_size"): drafter_config.draft_vocab_size = drafter_config.vocab_size if not hasattr(drafter_config, "target_hidden_size"): diff --git a/verl_speco/models/eagle/llama_eagle.py b/verl_speco/models/eagle/llama_eagle.py index 932eaf37..2dcfed0b 100644 --- a/verl_speco/models/eagle/llama_eagle.py +++ b/verl_speco/models/eagle/llama_eagle.py @@ -1318,8 +1318,9 @@ def __init__(self, config): self.act_fn = ACT2FN[config.hidden_act] def forward(self, x): - if self.config.pretraining_tp > 1: - slice = self.intermediate_size // self.config.pretraining_tp + pretraining_tp = int(getattr(self.config, "pretraining_tp", 1)) + if pretraining_tp > 1: + slice = self.intermediate_size // pretraining_tp gate_proj_slices = self.gate_proj.weight.split(slice, dim=0) up_proj_slices = self.up_proj.weight.split(slice, dim=0) down_proj_slices = self.down_proj.weight.split(slice, dim=1) @@ -1327,14 +1328,14 @@ def forward(self, x): gate_proj = torch.cat( [ F.linear(x, gate_proj_slices[i]) - for i in range(self.config.pretraining_tp) + for i in range(pretraining_tp) ], dim=-1, ) up_proj = torch.cat( [ F.linear(x, up_proj_slices[i]) - for i in range(self.config.pretraining_tp) + for i in range(pretraining_tp) ], dim=-1, ) @@ -1342,7 +1343,7 @@ def forward(self, x): intermediate_states = (self.act_fn(gate_proj) * up_proj).split(slice, dim=2) down_proj = [ F.linear(intermediate_states[i], down_proj_slices[i]) - for i in range(self.config.pretraining_tp) + for i in range(pretraining_tp) ] down_proj = sum(down_proj) else: diff --git a/verl_speco/trainer/standalone_checkpoint.py b/verl_speco/trainer/standalone_checkpoint.py index 6eae949e..04fbc8f3 100644 --- a/verl_speco/trainer/standalone_checkpoint.py +++ b/verl_speco/trainer/standalone_checkpoint.py @@ -161,6 +161,30 @@ def _normalize_dspark_runtime_architecture( runtime_config["architectures"] = [architecture] +def _normalize_eagle3_runtime_config( + runtime_config: dict[str, Any], + target_runtime_config: dict[str, Any] | None, +) -> None: + runtime_config["model_type"] = "llama" + runtime_config["architectures"] = ["LlamaForCausalLMEagle3"] + runtime_config.setdefault("draft_model_type", "eagle3") + runtime_config.setdefault("speculative_algorithm", "EAGLE3") + runtime_config.setdefault("pretraining_tp", 1) + runtime_config["num_hidden_layers"] = 1 + runtime_config["tie_word_embeddings"] = False + if ( + "draft_vocab_size" not in runtime_config + and runtime_config.get("vocab_size") is not None + ): + runtime_config["draft_vocab_size"] = int(runtime_config["vocab_size"]) + if ( + "target_hidden_size" not in runtime_config + and isinstance(target_runtime_config, dict) + and target_runtime_config.get("hidden_size") is not None + ): + runtime_config["target_hidden_size"] = int(target_runtime_config["hidden_size"]) + + _VARIANT_RUNTIME_ALIASES: dict[str, tuple[str, tuple[str, ...]]] = { "domino": ( "dflash_config", @@ -209,7 +233,7 @@ def rewrite_standalone_runtime_config( """ backend_type = getattr(getattr(trainer, "backend", None), "model_type", None) - if backend_type not in {"dflash", "dspark", "domino"}: + if backend_type not in {"dflash", "dspark", "domino", "eagle3"}: return None if completed_future is not None: @@ -261,6 +285,8 @@ def rewrite_standalone_runtime_config( _normalize_block_runtime_model_type(runtime_config, target_model_type, backend_type) if backend_type == "dspark": _normalize_dspark_runtime_architecture(runtime_config, target_model_type) + elif backend_type == "eagle3": + _normalize_eagle3_runtime_config(runtime_config, target_runtime_config) runtime_config["speco_training_model_type"] = backend_type target_runtime_keys = ("head_dim", "rope_theta", "max_position_embeddings") From 2258be29a869ecb71eb66c820c42f54a78459aa5 Mon Sep 17 00:00:00 2001 From: vx120 <893600387@qq.com> Date: Wed, 26 Aug 2026 17:06:43 +0800 Subject: [PATCH 40/50] feat: improve EAGLE3 standalone training workflow Signed-off-by: vx120 <893600387@qq.com> --- ...en3-8b_drafter_eagle3_separate_training.sh | 164 ++++++++++++++++++ verl_speco/backends/eagle3_trainer_backend.py | 19 +- verl_speco/config/speco_base.yaml | 2 + verl_speco/models/auto.py | 31 ++++ verl_speco/producer/input_reader.py | 48 +++++ 5 files changed, 256 insertions(+), 8 deletions(-) create mode 100644 examples/run_qwen3-8b_drafter_eagle3_separate_training.sh diff --git a/examples/run_qwen3-8b_drafter_eagle3_separate_training.sh b/examples/run_qwen3-8b_drafter_eagle3_separate_training.sh new file mode 100644 index 00000000..83c70fac --- /dev/null +++ b/examples/run_qwen3-8b_drafter_eagle3_separate_training.sh @@ -0,0 +1,164 @@ +#!/usr/bin/env bash +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +set -euo pipefail +set -x + +# Standalone EAGLE3 drafter training. Start +# examples/run_qwen3-8b_drafter_hidden_state_vllm.sh in another terminal first. +# The target-model vLLM and the drafter trainer can therefore use disjoint GPUs. +# +# EAGLE3 can initialize its drafter structure from the target model config, so +# no pre-initialized drafter directory is required. Set model_path only when +# loading an existing drafter checkpoint/config is desired. +# +# The vLLM hidden-state layer IDs must be EAGLE3_TARGET_LAYER_IDS followed by +# the target model's final layer. For Qwen3-8B the default is: +# [1,9,17,25,33,36] +# The EAGLE3 drafter config must have the same number (five) of aux states. + +project_name=${PROJECT_NAME:-verl_eagle3_drafter} +exp_name=${EXP_NAME:-qwen3_8b_eagle3_separate_training} + +draft_train_gpus_per_node=${TRAIN_GPUS:-2} +MODEL_PATH=${MODEL_PATH:-/path/to/Qwen3-4B} +TRAIN_FILE=${TRAIN_FILE:-/path/to/data} +DRAFT_CKPTS_DIR=${DRAFT_CKPTS_DIR:-/path/to/ckpt} + +PYTHON_BIN=${PYTHON_BIN:-python3} +DEVICE_ENV=${DEVICE_ENV:-CUDA_VISIBLE_DEVICES} +TRAIN_DEVICES=${TRAIN_DEVICES:-6,7} +SPECO_VLLM_ENDPOINTS=${SPECO_VLLM_ENDPOINTS:-'[http://127.0.0.1:8000/v1]'} +VLLM_READY_TIMEOUT_SECONDS=${VLLM_READY_TIMEOUT_SECONDS:-120} + +# These IDs must equal the auxiliary prefix of VLLM_HIDDEN_STATE_LAYER_IDS in +# run_qwen3-8b_drafter_hidden_state_vllm.sh. Do not include the final layer. +EAGLE3_TARGET_LAYER_IDS=${EAGLE3_TARGET_LAYER_IDS:-'[1,9,17,25,33]'} + +# Producer throughput and bounded queues. +VLLM_REQUEST_TIMEOUT=${VLLM_REQUEST_TIMEOUT:-120} +VLLM_MAX_INFLIGHT_REQUESTS=${VLLM_MAX_INFLIGHT_REQUESTS:-16} +VLLM_PER_ENDPOINT_CONCURRENCY=${VLLM_PER_ENDPOINT_CONCURRENCY:-4} +PRODUCER_INPUT_QUEUE_SIZE=${PRODUCER_INPUT_QUEUE_SIZE:-32} +PRODUCER_PUBLISH_QUEUE_SIZE=${PRODUCER_PUBLISH_QUEUE_SIZE:-16} +PRODUCER_MAX_PENDING_SAMPLES=${PRODUCER_MAX_PENDING_SAMPLES:-1024} +PRODUCER_PENDING_POLL_INTERVAL=${PRODUCER_PENDING_POLL_INTERVAL:-0.5} +PRODUCER_MAX_SEQUENCE_LENGTH=${PRODUCER_MAX_SEQUENCE_LENGTH:-8192} +PRODUCER_MAX_FEATURE_LENGTH=${PRODUCER_MAX_FEATURE_LENGTH:-512} +PRODUCER_GENERATION_MAX_TOKENS=${PRODUCER_GENERATION_MAX_TOKENS:-511} + +# Standalone EAGLE3 trainer settings. +MAX_STEPS=${MAX_STEPS:-1000} +SAVE_INTERVAL_STEPS=${SAVE_INTERVAL_STEPS:-100} +SAVE_FINAL_CHECKPOINT=${SAVE_FINAL_CHECKPOINT:-true} +BATCH_SIZE_PER_GPU=${BATCH_SIZE_PER_GPU:-2} +LEARNING_RATE=${LEARNING_RATE:-1e-5} +LR_WARMUP_STEPS=${LR_WARMUP_STEPS:-0} +LR_SCHEDULER_TYPE=${LR_SCHEDULER_TYPE:-constant} +LR_DECAY_STEPS=${LR_DECAY_STEPS:-1000} +MIN_LR_RATIO=${MIN_LR_RATIO:-0.1} +PARAM_OFFLOAD=${PARAM_OFFLOAD:-true} +OPTIMIZER_OFFLOAD=${OPTIMIZER_OFFLOAD:-true} + +if [[ "${MODEL_PATH}" == /path/to/* || "${TRAIN_FILE}" == /path/to/* ]]; then + echo "Set MODEL_PATH and TRAIN_FILE before starting training." >&2 + exit 2 +fi + +export "${DEVICE_ENV}=${TRAIN_DEVICES}" +export SPECO_VLLM_ENDPOINTS + +# Avoid the launcher's localhost fallback vLLM: this job must consume the +# separately managed hidden-state services, which keep target inference off the +# training devices. +"${PYTHON_BIN}" - "${SPECO_VLLM_ENDPOINTS}" "${VLLM_READY_TIMEOUT_SECONDS}" <<'PY' +import sys +import time +from urllib.error import URLError +from urllib.request import urlopen + +raw_endpoints = sys.argv[1].strip() +if not (raw_endpoints.startswith("[") and raw_endpoints.endswith("]")): + raise SystemExit("SPECO_VLLM_ENDPOINTS must use [url0,url1] syntax") +endpoints = [ + item.strip().strip("'\"").rstrip("/") + for item in raw_endpoints[1:-1].split(",") + if item.strip() +] +if not endpoints: + raise SystemExit("SPECO_VLLM_ENDPOINTS must contain at least one URL") +deadline = time.monotonic() + float(sys.argv[2]) +pending = set(endpoints) +while pending: + for endpoint in list(pending): + try: + with urlopen(f"{endpoint}/models", timeout=2) as response: + if 200 <= response.status < 300: + print(f"EXTERNAL_VLLM_READY endpoint={endpoint}", flush=True) + pending.remove(endpoint) + except (OSError, URLError): + pass + if pending and time.monotonic() >= deadline: + raise SystemExit( + "external hidden-state vLLM is not ready at: " + + ", ".join(sorted(pending)) + + "; start examples/run_qwen3-8b_drafter_hidden_state_vllm.sh first" + ) + if pending: + time.sleep(1) +PY + +PYTHONUNBUFFERED=1 "${PYTHON_BIN}" -m verl_speco.standalone_tq_training_launcher \ + speco.draft_training.num_gpus_per_node=${draft_train_gpus_per_node} \ + speco.draft_training.nnodes=1 \ + speco.draft_training.standalone=True \ + data.train_files=${TRAIN_FILE} \ + actor_rollout_ref.model.path=${MODEL_PATH} \ + actor_rollout_ref.actor.strategy=fsdp2 \ + actor_rollout_ref.actor.fsdp_config.param_offload=${PARAM_OFFLOAD} \ + actor_rollout_ref.actor.fsdp_config.optimizer_offload=${OPTIMIZER_OFFLOAD} \ + actor_rollout_ref.rollout.tensor_model_parallel_size=1 \ + actor_rollout_ref.rollout.drafter.enable=true \ + actor_rollout_ref.rollout.drafter.enable_drafter_training=true \ + actor_rollout_ref.rollout.drafter.checkpoint_path=${DRAFT_CKPTS_DIR} \ + actor_rollout_ref.rollout.drafter.speculative_algorithm=EAGLE3 \ + actor_rollout_ref.rollout.drafter.rollout.spec_steps=3 \ + actor_rollout_ref.rollout.drafter.rollout.spec_topk=1 \ + actor_rollout_ref.rollout.drafter.rollout.spec_verify_tokens=4 \ + actor_rollout_ref.rollout.drafter.training.mode=offline \ + actor_rollout_ref.rollout.drafter.training.max_steps=${MAX_STEPS} \ + actor_rollout_ref.rollout.drafter.training.save_interval_steps=${SAVE_INTERVAL_STEPS} \ + actor_rollout_ref.rollout.drafter.training.save_final_checkpoint=${SAVE_FINAL_CHECKPOINT} \ + actor_rollout_ref.rollout.drafter.training.batch_size_per_gpu=${BATCH_SIZE_PER_GPU} \ + actor_rollout_ref.rollout.drafter.training.lr=${LEARNING_RATE} \ + actor_rollout_ref.rollout.drafter.training.lr_warmup_steps=${LR_WARMUP_STEPS} \ + actor_rollout_ref.rollout.drafter.training.lr_scheduler_type=${LR_SCHEDULER_TYPE} \ + actor_rollout_ref.rollout.drafter.training.lr_decay_steps=${LR_DECAY_STEPS} \ + actor_rollout_ref.rollout.drafter.training.min_lr_ratio=${MIN_LR_RATIO} \ + actor_rollout_ref.rollout.drafter.training.use_logits=false \ + actor_rollout_ref.rollout.drafter.training.eagle3_target_layer_ids=${EAGLE3_TARGET_LAYER_IDS} \ + speco.standalone_tq_producer.target_layer_ids=${EAGLE3_TARGET_LAYER_IDS} \ + speco.standalone_tq_producer.request_timeout=${VLLM_REQUEST_TIMEOUT} \ + speco.standalone_tq_producer.max_inflight_requests=${VLLM_MAX_INFLIGHT_REQUESTS} \ + speco.standalone_tq_producer.per_endpoint_concurrency=${VLLM_PER_ENDPOINT_CONCURRENCY} \ + speco.standalone_tq_producer.input_queue_size=${PRODUCER_INPUT_QUEUE_SIZE} \ + speco.standalone_tq_producer.publish_queue_size=${PRODUCER_PUBLISH_QUEUE_SIZE} \ + speco.standalone_tq_producer.max_pending_samples=${PRODUCER_MAX_PENDING_SAMPLES} \ + speco.standalone_tq_producer.pending_poll_interval_seconds=${PRODUCER_PENDING_POLL_INTERVAL} \ + speco.standalone_tq_producer.max_sequence_length=${PRODUCER_MAX_SEQUENCE_LENGTH} \ + speco.standalone_tq_producer.max_feature_length=${PRODUCER_MAX_FEATURE_LENGTH} \ + speco.standalone_tq_producer.generation_max_tokens=${PRODUCER_GENERATION_MAX_TOKENS} \ + trainer.project_name=${project_name} \ + trainer.experiment_name=${exp_name} \ + "$@" diff --git a/verl_speco/backends/eagle3_trainer_backend.py b/verl_speco/backends/eagle3_trainer_backend.py index edcbc6f4..25843b0e 100644 --- a/verl_speco/backends/eagle3_trainer_backend.py +++ b/verl_speco/backends/eagle3_trainer_backend.py @@ -13,7 +13,6 @@ # limitations under the License. import logging import os -from copy import deepcopy from typing import Any, Optional, cast import torch @@ -23,7 +22,11 @@ from verl.utils.device import get_device_id, get_device_name from verl_speco.backends.lr_scheduler import build_drafter_lr_scheduler -from verl_speco.models.auto import AutoDraftModelConfig, AutoEagle3DraftModel +from verl_speco.models.auto import ( + AutoDraftModelConfig, + AutoEagle3DraftModel, + eagle3_draft_config_from_target, +) from verl_speco.models.eagle.llama_eagle import resolve_eagle3_num_aux_hidden_states from verl_speco.models.target.target_head import TargetHead from verl_speco.trainer.checkpoint import log_drafter_checkpoint_step @@ -678,16 +681,17 @@ def build_model(self): spec_model_path = self.config.rollout.drafter.model_path config_path = os.path.join(spec_model_path, "config.json") target_hf_config = self._get_target_hf_config() + training_cfg = self.config.rollout.drafter.training # 1. Load config if os.path.exists(config_path): drafter_config = AutoDraftModelConfig.from_file(config_path) else: - drafter_config = deepcopy(target_hf_config) - drafter_config.num_hidden_layers = 1 - drafter_config.torch_dtype = torch.bfloat16 - drafter_config.tie_word_embeddings = False - drafter_config.architectures = ["LlamaForCausalLMEagle3"] + drafter_config = eagle3_draft_config_from_target( + target_hf_config, + training_cfg.get("eagle3_target_layer_ids"), + ) + drafter_config.dtype = torch.bfloat16 if not hasattr(drafter_config, "pretraining_tp"): drafter_config.pretraining_tp = 1 @@ -740,7 +744,6 @@ def build_model(self): drafter_module.load_embedding(target_model_path) drafter_module.freeze_embedding() - training_cfg = self.config.rollout.drafter.training if drafter_module.draft_vocab_size != drafter_module.vocab_size: if checkpoint_has_vocab_mapping and self._has_valid_vocab_mapping( drafter_module diff --git a/verl_speco/config/speco_base.yaml b/verl_speco/config/speco_base.yaml index 94588d20..74c1e8aa 100644 --- a/verl_speco/config/speco_base.yaml +++ b/verl_speco/config/speco_base.yaml @@ -166,6 +166,8 @@ actor_rollout_ref: dspark_debug_log: false dspark_debug_log_first_n: 2 dspark_debug_log_interval: 100 + # EAGLE3 target layers whose hidden states are concatenated as drafter input. + eagle3_target_layer_ids: null # P-EAGLE (parallel drafting: COD-subsampled multi-depth forward + KL loss). peagle_num_draft_layers: 4 peagle_num_aux_hidden_states: 3 diff --git a/verl_speco/models/auto.py b/verl_speco/models/auto.py index 985ecb36..9ad6a73b 100644 --- a/verl_speco/models/auto.py +++ b/verl_speco/models/auto.py @@ -160,6 +160,37 @@ class AutoEagle3DraftModel(AutoDraftModel): } +def eagle3_draft_config_from_target( + target_config: PretrainedConfig, target_layer_ids=None +) -> LlamaConfig: + """Convert a target-model config to the Llama-compatible EAGLE3 config.""" + if not isinstance(target_config, PretrainedConfig): + raise TypeError( + "EAGLE3 target config must be a transformers.PretrainedConfig, got " + f"{type(target_config)!r}" + ) + + config = target_config.to_dict() + config.update( + { + "architectures": ["LlamaForCausalLMEagle3"], + "model_type": "llama", + "num_hidden_layers": 1, + "pretraining_tp": int(config.get("pretraining_tp") or 1), + "target_hidden_size": int(target_config.hidden_size), + "tie_word_embeddings": False, + } + ) + if target_layer_ids is not None: + layer_ids = _normalize_int_list(target_layer_ids) + if not layer_ids: + raise ValueError("EAGLE3 target_layer_ids must not be empty") + config["target_hidden_layer_ids"] = layer_ids + config["eagle_aux_hidden_state_layer_ids"] = layer_ids + + return LlamaConfig.from_dict(_normalize_eagle3_config_dict(config)) + + class AutoDraftModelConfig: _config_mapping = { "LlamaForCausalLMEagle3": LlamaConfig, diff --git a/verl_speco/producer/input_reader.py b/verl_speco/producer/input_reader.py index 69e4bdf8..c68c3245 100644 --- a/verl_speco/producer/input_reader.py +++ b/verl_speco/producer/input_reader.py @@ -296,6 +296,14 @@ def tokenize_record( full_ids = _token_ids( tokenizer(record.prompt + record.response, add_special_tokens=False) ) + if full_ids[: len(prompt_ids)] != prompt_ids: + # Some tokenizers merge text across the prompt/response boundary. + # Keep that boundary explicit so the response-only loss mask and the + # exact token IDs sent to vLLM remain aligned. + response_ids = _token_ids( + tokenizer(record.response, add_special_tokens=False) + ) + full_ids = [*prompt_ids, *response_ids] else: full_ids = _token_ids( tokenizer.apply_chat_template( @@ -304,6 +312,12 @@ def tokenize_record( add_generation_prompt=False, ) ) + if full_ids[: len(prompt_ids)] != prompt_ids: + prompt_ids, full_ids = _tokenize_chat_response_with_explicit_boundary( + record.prompt, + record.response, + tokenizer, + ) if full_ids[: len(prompt_ids)] != prompt_ids: raise ValueError( f"Producer sample {record.sample_id!r} has an unstable tokenizer boundary " @@ -460,6 +474,40 @@ def _prompt_ids(prompt: str | tuple[dict[str, str], ...], tokenizer: Any) -> lis ) +def _tokenize_chat_response_with_explicit_boundary( + prompt: tuple[dict[str, str], ...], + response: str, + tokenizer: Any, +) -> tuple[list[int], list[int]]: + """Render a chat response while keeping its loss boundary deterministic. + + Qwen-family templates may render a generation prompt differently from an + existing assistant message (for example by inserting a thinking preamble). + A marker lets us retain the template's assistant header and suffix while + tokenizing the response as a separate loss-bearing region. + """ + + marker = "__VERL_SPECO_ASSISTANT_RESPONSE_BOUNDARY_8F7C2D91__" + while marker in response or any(marker in message["content"] for message in prompt): + marker += "_" + rendered = tokenizer.apply_chat_template( + [*prompt, {"role": "assistant", "content": marker}], + tokenize=False, + add_generation_prompt=False, + ) + if not isinstance(rendered, str) or rendered.count(marker) != 1: + raise ValueError( + "Tokenizer chat template did not preserve the assistant response marker" + ) + prompt_text, suffix_text = rendered.split(marker, 1) + explicit_prompt_ids = _token_ids( + tokenizer(prompt_text, add_special_tokens=False) + ) + response_ids = _token_ids(tokenizer(response, add_special_tokens=False)) + suffix_ids = _token_ids(tokenizer(suffix_text, add_special_tokens=False)) + return explicit_prompt_ids, [*explicit_prompt_ids, *response_ids, *suffix_ids] + + def _build_tokenized_request( *, sequence_no: int, From c99ad697ae0be62c631ac9699a1521da2b4bda97 Mon Sep 17 00:00:00 2001 From: vx120 <893600387@qq.com> Date: Thu, 27 Aug 2026 10:38:13 +0800 Subject: [PATCH 41/50] change the version of the tq Signed-off-by: vx120 <893600387@qq.com> --- pyproject.toml | 2 +- tests/unit/test_tq_producer.py | 2 +- tests/unit/test_transferqueue_bridge.py | 2 +- tools/tq_connection_smoke.py | 4 ++-- tools/tq_delayed_test_producer.py | 2 +- verl_speco/config/speco_base.yaml | 6 +++--- .../integration/transferqueue_bridge.py | 20 +++++++++---------- verl_speco/standalone_tq_producer.py | 6 +++--- verl_speco/standalone_tq_training_launcher.py | 2 +- verl_speco/tq_owner.py | 2 +- verl_speco/trainer/tq_feature_store.py | 2 +- 11 files changed, 25 insertions(+), 25 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index d87e9806..72672830 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -17,7 +17,7 @@ dependencies = [ ] [project.optional-dependencies] -transfer-queue = ["TransferQueue==0.1.7"] +transfer-queue = ["TransferQueue==0.1.10"] [project.urls] Repository = "https://github.com/verl-project/verl-SpeCo" diff --git a/tests/unit/test_tq_producer.py b/tests/unit/test_tq_producer.py index 5744d2c8..1204b8c1 100644 --- a/tests/unit/test_tq_producer.py +++ b/tests/unit/test_tq_producer.py @@ -63,7 +63,7 @@ def _config(input_path: Path) -> dict[str, Any]: "dspark_l1_loss_alpha": 0.9, "transfer_queue": { "enable": True, - "package_version": "0.1.7", + "package_version": "0.1.10", "ray": { "address": "ray-head:6379", "namespace": "speco-drafter", diff --git a/tests/unit/test_transferqueue_bridge.py b/tests/unit/test_transferqueue_bridge.py index f792819c..ff6fa694 100644 --- a/tests/unit/test_transferqueue_bridge.py +++ b/tests/unit/test_transferqueue_bridge.py @@ -106,7 +106,7 @@ def fake_runtime(monkeypatch): def _config(): return { "enable": True, - "package_version": "0.1.7", + "package_version": "0.1.10", "partition_id": "speco_drafter_features", "run_id": "run-a", "schema_version": 1, diff --git a/tools/tq_connection_smoke.py b/tools/tq_connection_smoke.py index ec2869be..40dbdee0 100644 --- a/tools/tq_connection_smoke.py +++ b/tools/tq_connection_smoke.py @@ -11,7 +11,7 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. -"""Real two-process smoke test for Ray + TransferQueue 0.1.7. +"""Real two-process smoke test for Ray + TransferQueue 0.1.10. Start a Ray head first, then run ``owner`` and ``client`` in separate shells. The owner publishes one protocol-valid sample; the client list/get/decodes and @@ -49,7 +49,7 @@ def _config(args) -> dict: return { "enable": True, - "package_version": "0.1.7", + "package_version": "0.1.10", "ray": {"address": args.ray_address, "namespace": args.namespace}, "partition_id": "speco_drafter_features", "run_id": args.run_id, diff --git a/tools/tq_delayed_test_producer.py b/tools/tq_delayed_test_producer.py index 60f98b3c..af9e3af1 100644 --- a/tools/tq_delayed_test_producer.py +++ b/tools/tq_delayed_test_producer.py @@ -44,7 +44,7 @@ def _tq_config(args: argparse.Namespace) -> dict[str, Any]: return { "enable": True, - "package_version": "0.1.7", + "package_version": "0.1.10", "ray": {"address": args.ray_address, "namespace": args.namespace}, "partition_id": args.partition_id, "run_id": args.run_id, diff --git a/verl_speco/config/speco_base.yaml b/verl_speco/config/speco_base.yaml index 74c1e8aa..d68987a5 100644 --- a/verl_speco/config/speco_base.yaml +++ b/verl_speco/config/speco_base.yaml @@ -223,11 +223,11 @@ actor_rollout_ref: # rollout server writes hidden states directly to TQ and the drafter # worker reads them by key, bypassing the RayPPOTrainer driver and the # Ray object store for the dominant tensor. Default off -> unchanged - # Ray path. Requires `pip install TransferQueue==0.1.7`. + # Ray path. Requires `pip install TransferQueue==0.1.10`. transfer_queue: enable: false - package_version: "0.1.7" - # TQ 0.1.7 discovers its named TransferQueueController through Ray. + package_version: "0.1.10" + # TQ 0.1.10 discovers its named TransferQueueController through Ray. # Standalone owner, Producer and every torchrun rank must use the same # address and namespace. PR #48 Ray actors already have this context. ray: diff --git a/verl_speco/integration/transferqueue_bridge.py b/verl_speco/integration/transferqueue_bridge.py index 497e76ed..15bca804 100644 --- a/verl_speco/integration/transferqueue_bridge.py +++ b/verl_speco/integration/transferqueue_bridge.py @@ -38,7 +38,7 @@ finer-grained leader-clears-after-barrier is future work. Note: the TQ call sites follow the documented KV API (``kv_put`` / -``kv_batch_get`` / ``kv_close``) of TransferQueue 0.1.7. The exact signatures +``kv_batch_get`` / ``kv_close``) of TransferQueue 0.1.10. The exact signatures (keyword names, return shapes) must be verified against the installed TQ version on first run; the bridge fails loud, never silently. """ @@ -77,7 +77,7 @@ def __getattr__(self, name: str) -> Any: def _raise(*args: Any, **kwargs: Any) -> Any: raise RuntimeError( f"transfer_queue is not installed. Cannot call tq.{name}(). " - "Install with `pip install TransferQueue==0.1.7` or disable " + "Install with `pip install TransferQueue==0.1.10` or disable " "actor_rollout_ref.rollout.drafter.training.transfer_queue.enable." ) @@ -174,14 +174,14 @@ def connect_ray_cluster( PR #48 workers are already Ray actors and therefore do not call this function. Standalone owner, Producer and torchrun ranks must call it - before ``tq.init`` so TQ 0.1.7 can discover its named Controller actor. + before ``tq.init`` so TQ 0.1.10 can discover its named Controller actor. """ try: import ray except ImportError as exc: # pragma: no cover - depends on optional package raise RuntimeError( - "Ray is required by TransferQueue 0.1.7. Install TransferQueue==0.1.7 " + "Ray is required by TransferQueue 0.1.10. Install TransferQueue==0.1.10 " "and connect all standalone processes to the same Ray cluster." ) from exc @@ -207,7 +207,7 @@ def start_transfer_queue_owner(tq_config: Any) -> None: if not bool(plain.get("enable", True)): raise ValueError("TransferQueue owner requires transfer_queue.enable=true") if not _TQ_IMPORTABLE: - raise RuntimeError("TransferQueue==0.1.7 is required to start the TQ owner") + raise RuntimeError("TransferQueue==0.1.10 is required to start the TQ owner") with _state_lock: if _state["initialized"]: raise RuntimeError("TransferQueue is already initialized in this process") @@ -225,7 +225,7 @@ def connect_transfer_queue_client() -> None: """Attach this process to the named TQ Controller on its Ray cluster.""" if not _TQ_IMPORTABLE: - raise RuntimeError("TransferQueue==0.1.7 is required to connect a TQ client") + raise RuntimeError("TransferQueue==0.1.10 is required to connect a TQ client") if not bool(_state["enabled"]): raise RuntimeError( "configure_transfer_queue() must enable TQ before client connect" @@ -269,7 +269,7 @@ def _drafter_training_cfg(config: Any) -> Any: def _ensure_initialized() -> None: """Lazily ``tq.init(config)`` once per worker process. - TransferQueue 0.1.7 first tries to discover the named controller and ignores + TransferQueue 0.1.10 first tries to discover the named controller and ignores the supplied configuration when one already exists. Supplying the same native configuration in every process is therefore safe for ordinary clients and also prevents an unexpectedly early client from creating a @@ -308,7 +308,7 @@ def _native_tq_config(tq_cfg: Mapping[str, Any]) -> dict[str, Any]: def _as_tq_config(value: Mapping[str, Any]) -> Any: - """TQ 0.1.7 annotates its config as DictConfig; keep tests dependency-light.""" + """TQ 0.1.10 annotates its config as DictConfig; keep tests dependency-light.""" try: from omegaconf import OmegaConf @@ -361,7 +361,7 @@ def put_sample( _ensure_initialized() # Pass a plain single-sample dict of columns. TQ's kv_put adds its required # batch dimension internally; constructing a scalar TensorDict here would be - # incorrect. (Exact kwarg names verified against TQ 0.1.7 on first run.) + # incorrect. (Exact kwarg names verified against TQ 0.1.10 on first run.) tq.kv_put( key=key, partition_id=_partition_id(), @@ -401,7 +401,7 @@ def list_samples() -> dict[str, dict[str, Any]]: return {} if not isinstance(result, Mapping): raise TypeError(f"tq.kv_list returned unsupported type {type(result)!r}") - # 0.1.7 returns key -> tag when partition_id is supplied. Accept the + # 0.1.10 returns key -> tag when partition_id is supplied. Accept the # partition -> (key -> tag) wrapper as well to keep the bridge version-safe. nested = result.get(_partition_id()) if isinstance(nested, Mapping) and all( diff --git a/verl_speco/standalone_tq_producer.py b/verl_speco/standalone_tq_producer.py index d749d8b9..e06fbae1 100644 --- a/verl_speco/standalone_tq_producer.py +++ b/verl_speco/standalone_tq_producer.py @@ -123,8 +123,8 @@ def validate_producer_config(config: Any) -> None: raise ValueError( f"transfer_queue.schema_version must be {PROTOCOL_SCHEMA_VERSION}" ) - if tq_cfg.get("package_version") != "0.1.7": - raise ValueError("transfer_queue.package_version must be '0.1.7'") + if tq_cfg.get("package_version") != "0.1.10": + raise ValueError("transfer_queue.package_version must be '0.1.10'") if tq_cfg.get("partition_id") != DRAFTER_TQ_PARTITION: raise ValueError( f"transfer_queue.partition_id must be {DRAFTER_TQ_PARTITION!r}" @@ -176,7 +176,7 @@ async def run_producer( producer_cfg["vllm_endpoints"], ) if not transport.configure_transfer_queue(tq_cfg): - raise RuntimeError("Standalone TQ Producer requires TransferQueue==0.1.7") + raise RuntimeError("Standalone TQ Producer requires TransferQueue==0.1.10") ray_cfg = tq_cfg["ray"] logger.info( "Standalone TQ Producer connecting Ray address=%s namespace=%s", diff --git a/verl_speco/standalone_tq_training_launcher.py b/verl_speco/standalone_tq_training_launcher.py index 879bb6e1..2c1907ac 100644 --- a/verl_speco/standalone_tq_training_launcher.py +++ b/verl_speco/standalone_tq_training_launcher.py @@ -354,7 +354,7 @@ def start_ray_session( ray_runtime = importlib.import_module("ray") except ImportError as exc: raise RuntimeError( - "Standalone TQ training requires Ray and TransferQueue==0.1.7" + "Standalone TQ training requires Ray and TransferQueue==0.1.10" ) from exc # ``ray.init()`` consults RAY_ADDRESS when no explicit address is supplied. # This launcher owns the complete Producer/TQ/Consumer lifetime, so an diff --git a/verl_speco/tq_owner.py b/verl_speco/tq_owner.py index c9128848..16b1b10b 100644 --- a/verl_speco/tq_owner.py +++ b/verl_speco/tq_owner.py @@ -79,7 +79,7 @@ def run_owner(config: Any, *, stop_event: threading.Event | None = None) -> int: # and enable only this process's copied configuration. tq_cfg["enable"] = True if not configure_transfer_queue(tq_cfg): - raise RuntimeError("Standalone TQ owner requires TransferQueue==0.1.7") + raise RuntimeError("Standalone TQ owner requires TransferQueue==0.1.10") ray_cfg = tq_cfg.get("ray", {}) ray_address = ray_cfg.get("address") if not ray_address: diff --git a/verl_speco/trainer/tq_feature_store.py b/verl_speco/trainer/tq_feature_store.py index ff5ae570..4447a173 100644 --- a/verl_speco/trainer/tq_feature_store.py +++ b/verl_speco/trainer/tq_feature_store.py @@ -90,7 +90,7 @@ def connect(self) -> None: return if not configure_transfer_queue(self.config): raise RuntimeError( - "TQ Consumer requires transfer_queue.enable=true and TransferQueue==0.1.7" + "TQ Consumer requires transfer_queue.enable=true and TransferQueue==0.1.10" ) connect_ray_cluster(self.ray_address, self.ray_namespace) connect_transfer_queue_client() From b6dcd472d5f0f8914d4e41cb220b8a8c0dee7085 Mon Sep 17 00:00:00 2001 From: xxyyrr598 <1615681145@qq.com> Date: Thu, 27 Aug 2026 12:00:19 +0800 Subject: [PATCH 42/50] refactor(tq): reorganize standalone TQ scripts and fix lint/type checks - Move hidden_state_vllm script from examples/ to tools/, drop redundant tools scripts - Adjust example script naming checks - Fix mypy errors (None guards, annotations, missing return) and doc Last updated info --- README.md | 2 - docs/standalone_tq_consumer_implementation.md | 36 +-- ...standalone_tq_foundation_implementation.md | 23 +- docs/standalone_tq_producer.md | 6 +- ...standalone_vllm_tq_dspark_training_plan.md | 9 +- examples/run_dspark_tq_producer.sh | 46 ---- .../run_qwen3-8b_drafter_separate_training.sh | 51 +--- tests/examples/test_example_scripts.py | 16 +- tests/special_sanity/check_example_naming.py | 7 +- tools/run_dspark_tq_consumer.sh | 63 ----- tools/run_dspark_tq_consumer_test.sh | 84 ------- tools/run_dspark_tq_e2e_test.sh | 170 -------------- tools/run_dspark_tq_owner.sh | 29 --- tools/run_qwen3-4b_drafter_dspark_mooncake.sh | 72 ------ .../run_qwen3-8b_drafter_hidden_state_vllm.sh | 83 +++++-- tools/run_tq_connection_smoke.sh | 53 ----- tools/tq_connection_smoke.py | 175 -------------- tools/tq_delayed_test_producer.py | 222 ------------------ tools/wait_for_vllm_endpoints.py | 92 ++++++++ verl_speco/backends/dflash_trainer_backend.py | 8 +- verl_speco/draft_train_launcher.py | 8 +- verl_speco/inspect_jsonl_samples.py | 4 +- .../mooncake_hidden_states_connector.py | 4 +- verl_speco/integration/oldlogprob_runtime.py | 14 +- verl_speco/integration/sglang_runtime.py | 22 +- verl_speco/integration/task_runner.py | 5 +- verl_speco/models/eagle/llama_eagle.py | 10 +- verl_speco/producer/input_reader.py | 9 +- verl_speco/producer/vllm_feature_client.py | 78 +++++- verl_speco/standalone_tq_producer.py | 4 +- verl_speco/standalone_tq_training_launcher.py | 4 +- verl_speco/trainer/base_trainer.py | 4 +- verl_speco/trainer/draft_training_loop.py | 24 +- verl_speco/trainer/feature_store.py | 31 +-- verl_speco/trainer/mooncake_transfer.py | 4 +- verl_speco/trainer/speco_ray_trainer.py | 4 +- verl_speco/trainer/standalone_checkpoint.py | 80 +++++-- verl_speco/trainer/target_feature_pipeline.py | 4 +- verl_speco/trainer/target_feature_replay.py | 10 +- verl_speco/trainer/tq_feature_store.py | 9 +- verl_speco/trainer/tq_sample_source.py | 15 +- .../transport/drafter_sample_protocol.py | 28 ++- verl_speco/vllm_hidden_states_generate.py | 4 +- verl_speco/workers/speco_worker.py | 6 +- 44 files changed, 453 insertions(+), 1179 deletions(-) delete mode 100644 examples/run_dspark_tq_producer.sh delete mode 100644 tools/run_dspark_tq_consumer.sh delete mode 100644 tools/run_dspark_tq_consumer_test.sh delete mode 100644 tools/run_dspark_tq_e2e_test.sh delete mode 100644 tools/run_dspark_tq_owner.sh delete mode 100644 tools/run_qwen3-4b_drafter_dspark_mooncake.sh rename {examples => tools}/run_qwen3-8b_drafter_hidden_state_vllm.sh (50%) delete mode 100644 tools/run_tq_connection_smoke.sh delete mode 100644 tools/tq_connection_smoke.py delete mode 100644 tools/tq_delayed_test_producer.py create mode 100644 tools/wait_for_vllm_endpoints.py diff --git a/README.md b/README.md index 767a7a81..11d6fdb5 100644 --- a/README.md +++ b/README.md @@ -400,8 +400,6 @@ only while vLLM throughput rises. `prefetch_depth` of 2 normally hides transfer latency without retaining too many large hidden-state batches. The standalone metrics include producer queue depth, consumer wait time, vLLM request time, and Mooncake GET time. -The complete three-process launcher is -[`tools/run_qwen3-4b_drafter_dspark_mooncake.sh`](./tools/run_qwen3-4b_drafter_dspark_mooncake.sh). ## Configuration diff --git a/docs/standalone_tq_consumer_implementation.md b/docs/standalone_tq_consumer_implementation.md index ef01c311..78075686 100644 --- a/docs/standalone_tq_consumer_implementation.md +++ b/docs/standalone_tq_consumer_implementation.md @@ -27,8 +27,7 @@ Last updated: 08/21/2026 |---|---|---| | `verl_speco/trainer/tq_feature_store.py` | `TQFeatureStore`、`ReadyEntry`、`EosMetadata` | 将公共 TQ bridge 包装成 Consumer 数据访问层,负责连接、发现、批量读取、解码、删除和读取 EOS | | `verl_speco/trainer/tq_sample_source.py` | `TQFeatureDataLoader`、`TQLocalBatch`、`build_assignments()` | 实现多 rank 流式取数:rank 0 发现样本并分配 key,各 rank 自己从 TQ 取 Tensor | -| `tools/run_dspark_tq_consumer.sh` | Consumer 测试启动工具 | 给出一套完整的 DSpark、offline、TQ 配置和 `torchrun` 启动方式 | -| `tests/unit/test_tq_consumer.py` | Consumer 单元测试 | 覆盖 store 构建、ready 过滤排序、协议解码、rank 分配、EOS、清理和非 rank 0 行为 | +| `tests/unit/test_tq_consumer.py` | Consumer 单元测试 | 覆盖 store 构建、ready 过滤排序、协议解码、rank 分配、EOS、清理和非 rank 0 行为 | ### 2.2 本次修改的既有文件 @@ -38,8 +37,7 @@ Last updated: 08/21/2026 | `verl_speco/trainer/draft_training_loop.py` | 接入 TQ store/loader、跨 rank 连接检查、训练成功后清理 | 将流式取数接入原训练循环,同时保留原 DSpark trainer、loss、optimizer、metric 和 checkpoint 逻辑 | | `verl_speco/draft_train_launcher.py` | 增加 TQ 启动参数的 fail-fast 检查 | 在启动多个 torchrun 子进程前检查 `enable`、Ray address 和 `run_id`,避免各 rank 启动后才失败 | | `verl_speco/config/speco_base.yaml` | 标注 `feature_store.type=tq` 为无路径流式数据源 | 保留统一 Hydra 配置入口;TQ 的公共配置仍位于 sibling `training.transfer_queue` | -| `tools/tq_connection_smoke.py` | 将原连接 smoke 扩展为真实 Consumer 路径测试 | 验证 owner、真实 TQ、Consumer 读取、EOS、清理和仅关闭本地 client | -| `tests/unit/test_draft_train_launcher.py` | 增加 TQ 参数检查测试 | 验证必要配置缺失时 launcher 直接拒绝启动 | +| `tests/unit/test_draft_train_launcher.py` | 增加 TQ 参数检查测试 | 验证必要配置缺失时 launcher 直接拒绝启动 | | `tests/unit/test_draft_training_loop.py` | 增加连接和 clear 时序测试 | 验证只由 rank 0 清理、clear 失败会报告、连接失败会传播 | ### 2.3 直接复用的公共基础 @@ -612,29 +610,11 @@ TQ 会保留 Producer 写入的 `SampleMetadata.algorithm`,但不使用它选 ## 11. 如何启动和检查 -推荐启动顺序: - -1. 启动 Ray head。 -2. 启动 TQ Owner,并保持该进程存活。 -3. 启动 Producer,使用相同 Ray address、namespace、partition 和 run ID。 -4. 启动 `tools/run_dspark_tq_consumer.sh`。 -5. Producer 完成全部样本后发布 EOS。 -6. Consumer 消费完成并退出后,再停止 Owner/Ray。 - -示例: - -```bash -MODEL_PATH=/models/Qwen3-8B \ -DRAFTER_PATH=/models/dspark-drafter \ -DRAFT_CKPTS_DIR=/checkpoints/dspark-tq \ -TRAIN_DEVICES=0,1,2,3 \ -TRAIN_GPUS=4 \ -RAY_ADDRESS=127.0.0.1:6379 \ -SPECO_TQ_RUN_ID=dspark-run-001 \ -bash tools/run_dspark_tq_consumer.sh -``` - -脚本中的 namespace 固定为 `speco-drafter`,默认 partition 来自公共配置 `speco_drafter_features`。Owner 和 Producer 必须使用相同值。 +正式运行统一使用端到端 launcher;它负责 Ray、TQ Owner、Producer 和 Consumer 的启动与清理: + +```bash +bash examples/run_qwen3-8b_drafter_separate_training.sh +``` ## 12. 已完成的测试 @@ -656,7 +636,7 @@ python -m pytest \ ### 12.2 真实 TQ 0.1.7 跨进程 smoke -`tools/tq_connection_smoke.py` 使用真实 Ray + TQ Owner 和另一个 Consumer 进程验证了: +早期临时跨进程工具验证过以下行为,当前回归由协议、bridge、Consumer和launcher单元测试承担: - Owner 发布 owner-ready; - 两条 sample 写入 TQ; diff --git a/docs/standalone_tq_foundation_implementation.md b/docs/standalone_tq_foundation_implementation.md index fd01e7ff..b3e0dea3 100644 --- a/docs/standalone_tq_foundation_implementation.md +++ b/docs/standalone_tq_foundation_implementation.md @@ -8,12 +8,10 @@ Last updated: 08/21/2026 ```text verl_speco/transport/drafter_sample_protocol.py -verl_speco/integration/transferqueue_bridge.py -verl_speco/config/speco_base.yaml -verl_speco/tq_owner.py -tools/run_dspark_tq_owner.sh -tools/tq_connection_smoke.py -tests/unit/test_drafter_sample_protocol.py +verl_speco/integration/transferqueue_bridge.py +verl_speco/config/speco_base.yaml +verl_speco/tq_owner.py +tests/unit/test_drafter_sample_protocol.py tests/unit/test_transferqueue_bridge.py pyproject.toml ``` @@ -871,9 +869,7 @@ Owner 命令: verl-speco-tq-owner ``` -也可以使用 `tools/run_dspark_tq_owner.sh`。 - -## 20. 单元测试 +## 20. 单元测试 协议测试 `tests/unit/test_drafter_sample_protocol.py` 覆盖: @@ -905,13 +901,8 @@ python -m pytest \ ## 21. 真实双进程 smoke test -程序: - -```text -tools/tq_connection_smoke.py -``` - -它使用真实 `TransferQueue==0.1.7`、Ray、SimpleStorage、两个独立 Python进程和两个 batch samples。 +早期用于该验证的临时双进程工具已经移除;正式入口统一由 +`verl_speco.standalone_tq_training_launcher` 管理 Owner、Producer 和 Consumer 生命周期。 Owner 路径: diff --git a/docs/standalone_tq_producer.md b/docs/standalone_tq_producer.md index e651a04f..bb766907 100644 --- a/docs/standalone_tq_producer.md +++ b/docs/standalone_tq_producer.md @@ -176,11 +176,7 @@ vLLM 请求会暂停,等待 Consumer 清理已成功训练的 key。 | `schema_version` | `1` | | `run_id`、Ray address、Ray namespace | 三个进程必须相同 | -使用 [run_dspark_tq_producer.sh](../examples/run_dspark_tq_producer.sh) 启动。它要求 -显式提供 `RAY_ADDRESS`、`SPECO_TQ_RUN_ID`、输入、target/tokenizer 身份、layers 和 -vLLM endpoints,并且只启动 Producer。 - -也可以通过安装后的命令入口运行: +单独调试时可以通过安装后的命令入口运行;正式训练由统一launcher启动Producer: ```bash verl-speco-tq-producer \ diff --git a/docs/standalone_vllm_tq_dspark_training_plan.md b/docs/standalone_vllm_tq_dspark_training_plan.md index e22c8660..7a865b8b 100644 --- a/docs/standalone_vllm_tq_dspark_training_plan.md +++ b/docs/standalone_vllm_tq_dspark_training_plan.md @@ -430,7 +430,6 @@ transfer_queue: ```text verl_speco/tq_owner.py -tools/run_dspark_tq_owner.sh ``` `tq_owner.py` 建议明确实现: @@ -557,7 +556,6 @@ InputReader | `verl_speco/producer/input_reader.py` | 流式读输入、分配 `sequence_no/sample_id`、tokenize、构造 loss mask | | `verl_speco/producer/vllm_feature_client.py` | endpoint pool、并发控制、HTTP 请求、临时结果读取和删除 | | `verl_speco/trainer/target_feature_replay.py` | 提取可复用的 vLLM payload → `DraftFeatureSample` 纯转换函数 | -| `examples/run_dspark_tq_producer.sh` | Producer 配置和启动命令 | | `tests/unit/test_drafter_sample_protocol.py` | 协议编码、字段和 shape 测试(两人共同) | | `tests/integration/test_tq_producer_smoke.py` | fake/短 vLLM → TQ sample + EOS | @@ -644,9 +642,9 @@ def feature_from_vllm_payload( 它不进行 HTTP、不访问 TQ、不删除文件,只完成 shape/layout/layer 校验、feature-position 选择、aux layer flatten 和可选 L1 final hidden 拼接。现有 file replay backend 和新 Producer 都调用这一函数,避免两套转换规则。 -#### `examples/run_dspark_tq_producer.sh` +#### Producer 启动配置 -负责提供同一套: +统一 launcher 向 `verl_speco.standalone_tq_producer` 提供同一套: ```text RAY_ADDRESS / Ray namespace @@ -657,7 +655,7 @@ vLLM endpoint 列表 max_inflight_requests / per_endpoint_concurrency ``` -脚本只启动 Producer,不启动 TQ owner 或 Consumer,便于两位开发者独立调试。 +正式运行不再保留单独的角色 shell wrapper。 ## 5. Consumer 要实现什么 @@ -875,7 +873,6 @@ checkpoint 仍使用现有 `save_interval_steps` 和 final checkpoint;第一 | `verl_speco/trainer/draft_training_loop.py` | 保持 `mode=offline`;`type=tq` 时跳过 `feature_store.path` 必填检查并选择专用 loader;调用现有 prepare/train,global success 后 clear,final checkpoint | | `verl_speco/draft_train_launcher.py` | `feature_store.type=tq` 时不要求 feature-store path,透传 torchrun/TQ 配置 | | `verl_speco/config/speco_base.yaml` | 新增 TQ connection、run、poll、batch 等默认配置 | -| `tools/run_dspark_tq_consumer.sh` | Consumer GPU、batch、checkpoint 和共享 TQ 配置 | | `tests/unit/test_tq_sample_source.py` | key 过滤/排序/分 rank、EOS、decode 调用测试 | | `tests/integration/test_dspark_tq_consumer_smoke.py` | 两 rank 取不同样本,完成训练并 clear | diff --git a/examples/run_dspark_tq_producer.sh b/examples/run_dspark_tq_producer.sh deleted file mode 100644 index 63c19384..00000000 --- a/examples/run_dspark_tq_producer.sh +++ /dev/null @@ -1,46 +0,0 @@ -#!/usr/bin/env bash -# Copyright 2026 Bytedance Ltd. and/or its affiliates -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -set -euo pipefail - -: "${RAY_ADDRESS:?Set RAY_ADDRESS to the running Ray head}" -: "${SPECO_TQ_RUN_ID:?Set SPECO_TQ_RUN_ID to the owner/consumer run id}" -: "${PRODUCER_INPUT_PATH:?Set PRODUCER_INPUT_PATH to prompt/response JSONL}" -: "${TARGET_MODEL_PATH:?Set TARGET_MODEL_PATH to the target model id/path}" -: "${TARGET_MODEL_REVISION:?Set TARGET_MODEL_REVISION to a revision or checksum}" -: "${TOKENIZER_PATH:?Set TOKENIZER_PATH to the tokenizer id/path}" -: "${TOKENIZER_FINGERPRINT:?Set TOKENIZER_FINGERPRINT to a verified fingerprint}" -: "${TARGET_LAYER_IDS:?Set TARGET_LAYER_IDS as a Hydra list, for example '[2,8,14,20,26]'}" -: "${VLLM_ENDPOINTS:?Set VLLM_ENDPOINTS as a Hydra list, for example '[http://node0:8000/v1]'}" -PYTHON_BIN=${PYTHON_BIN:-python3} -TQ_NAMESPACE=${TQ_NAMESPACE:-speco-drafter} -TQ_PARTITION_ID=${TQ_PARTITION_ID:-speco_drafter_features} -VLLM_MODEL=${VLLM_MODEL:-${TARGET_MODEL_PATH}} - -exec "${PYTHON_BIN}" -m verl_speco.standalone_tq_producer \ - actor_rollout_ref.rollout.drafter.speculative_algorithm=DSPARK \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.enable=true \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.address="${RAY_ADDRESS}" \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.namespace="${TQ_NAMESPACE}" \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.partition_id="${TQ_PARTITION_ID}" \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.run_id="${SPECO_TQ_RUN_ID}" \ - speco.standalone_tq_producer.input_path="${PRODUCER_INPUT_PATH}" \ - speco.standalone_tq_producer.target_model_id="${TARGET_MODEL_PATH}" \ - speco.standalone_tq_producer.target_model_revision="${TARGET_MODEL_REVISION}" \ - speco.standalone_tq_producer.tokenizer_path="${TOKENIZER_PATH}" \ - speco.standalone_tq_producer.tokenizer_fingerprint="${TOKENIZER_FINGERPRINT}" \ - speco.standalone_tq_producer.target_layer_ids="${TARGET_LAYER_IDS}" \ - speco.standalone_tq_producer.vllm_endpoints="${VLLM_ENDPOINTS}" \ - speco.standalone_tq_producer.vllm_model="${VLLM_MODEL}" \ - "$@" diff --git a/examples/run_qwen3-8b_drafter_separate_training.sh b/examples/run_qwen3-8b_drafter_separate_training.sh index 5a49d844..1a08a43b 100644 --- a/examples/run_qwen3-8b_drafter_separate_training.sh +++ b/examples/run_qwen3-8b_drafter_separate_training.sh @@ -15,8 +15,12 @@ set -euo pipefail set -x +script_dir=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd) +repo_root=$(cd -- "${script_dir}/.." && pwd) +cd "${repo_root}" + # Standalone DSpark draft-model training using an already-running hidden-state -# vLLM. Start run_qwen3-8b_drafter_hidden_state_vllm.sh in another terminal +# vLLM. Start tools/run_qwen3-8b_drafter_hidden_state_vllm.sh in another terminal # first. This process owns Ray/TQ, Producer and Consumer, but it must not own # the target vLLM so inference and training can use different accelerators. @@ -51,7 +55,7 @@ PRODUCER_MAX_PENDING_SAMPLES=${PRODUCER_MAX_PENDING_SAMPLES:-1024} PRODUCER_PENDING_POLL_INTERVAL=${PRODUCER_PENDING_POLL_INTERVAL:-0.5} PRODUCER_MAX_SEQUENCE_LENGTH=${PRODUCER_MAX_SEQUENCE_LENGTH:-8192} PRODUCER_MAX_FEATURE_LENGTH=${PRODUCER_MAX_FEATURE_LENGTH:-512} -PRODUCER_GENERATION_MAX_TOKENS=${PRODUCER_GENERATION_MAX_TOKENS:-511} +PRODUCER_GENERATION_MAX_TOKENS=${PRODUCER_GENERATION_MAX_TOKENS:-512} # Standalone trainer. MAX_STEPS=${MAX_STEPS:-10} @@ -95,43 +99,12 @@ export SPECO_VLLM_ENDPOINTS # Fail before entering the unified launcher when the separately managed vLLM # is absent. Otherwise a localhost endpoint would make the launcher start its # fallback vLLM inside the training process and on the training devices. -"${PYTHON_BIN}" - "${SPECO_VLLM_ENDPOINTS}" "${VLLM_READY_TIMEOUT_SECONDS}" <<'PY' -import sys -import time -from urllib.error import URLError -from urllib.request import urlopen - -raw_endpoints = sys.argv[1].strip() -if not (raw_endpoints.startswith("[") and raw_endpoints.endswith("]")): - raise SystemExit("SPECO_VLLM_ENDPOINTS must use [url0,url1] syntax") -endpoints = [ - item.strip().strip("'\"").rstrip("/") - for item in raw_endpoints[1:-1].split(",") - if item.strip() -] -if not endpoints: - raise SystemExit("SPECO_VLLM_ENDPOINTS must contain at least one URL") -timeout_seconds = float(sys.argv[2]) -deadline = time.monotonic() + timeout_seconds -pending = set(endpoints) -while pending: - for endpoint in list(pending): - try: - with urlopen(f"{endpoint}/models", timeout=2) as response: - if 200 <= response.status < 300: - print(f"EXTERNAL_VLLM_READY endpoint={endpoint}", flush=True) - pending.remove(endpoint) - except (OSError, URLError): - pass - if pending and time.monotonic() >= deadline: - raise SystemExit( - "external hidden-state vLLM is not ready at: " - + ", ".join(sorted(pending)) - + "; start examples/run_qwen3-8b_drafter_hidden_state_vllm.sh first" - ) - if pending: - time.sleep(1) -PY +if ! "${PYTHON_BIN}" tools/wait_for_vllm_endpoints.py \ + --endpoints "${SPECO_VLLM_ENDPOINTS}" \ + --timeout-seconds "${VLLM_READY_TIMEOUT_SECONDS}"; then + echo "Start tools/run_qwen3-8b_drafter_hidden_state_vllm.sh first" >&2 + exit 1 +fi PYTHONUNBUFFERED=1 "${PYTHON_BIN}" -m verl_speco.standalone_tq_training_launcher \ speco.draft_training.num_gpus_per_node=${draft_train_gpus_per_node} \ diff --git a/tests/examples/test_example_scripts.py b/tests/examples/test_example_scripts.py index 0699f5d0..663b95c2 100644 --- a/tests/examples/test_example_scripts.py +++ b/tests/examples/test_example_scripts.py @@ -26,8 +26,6 @@ script for script in EXAMPLES if not script.name.endswith("_separate_training.sh") - and script.name != "run_dspark_tq_producer.sh" - and script.name != "run_qwen3-8b_drafter_hidden_state_vllm.sh" ] @@ -88,11 +86,11 @@ def test_standalone_tq_training_example_uses_unified_launcher() -> None: def test_standalone_tq_hidden_state_vllm_uses_separate_devices() -> None: source = ( - ROOT / "examples" / "run_qwen3-8b_drafter_hidden_state_vllm.sh" + ROOT / "tools" / "run_qwen3-8b_drafter_hidden_state_vllm.sh" ).read_text(encoding="utf-8") - assert 'VLLM_DEVICES_0=${VLLM_DEVICES_0:-0}' in source - assert 'VLLM_DEVICES_1=${VLLM_DEVICES_1:-1}' in source + assert 'VLLM_DEVICES=${VLLM_DEVICES:-0,1,2,3,4,5}' in source + assert "service_count=$((device_count / VLLM_TP))" in source assert 'env "${DEVICE_ENV}=${devices}" vllm serve "${MODEL_PATH}"' in source assert '--tensor-parallel-size "${VLLM_TP}"' in source assert '--max-num-seqs "${VLLM_MAX_NUM_SEQS}"' in source @@ -108,14 +106,6 @@ def test_standalone_tq_compatibility_example_delegates_to_formal_entry() -> None assert 'run_qwen3-8b_drafter_separate_training.sh" "$@"' in source -def test_standalone_tq_producer_example_uses_producer_entrypoint() -> None: - source = (ROOT / "examples" / "run_dspark_tq_producer.sh").read_text( - encoding="utf-8" - ) - - assert "-m verl_speco.standalone_tq_producer" in source - - def test_vllm_eagle3_example_keeps_runtime_agnostic_training_switches() -> None: source = (ROOT / "examples" / "run_qwen3-8b_drafter_eagle3_vllm.sh").read_text( encoding="utf-8" diff --git a/tests/special_sanity/check_example_naming.py b/tests/special_sanity/check_example_naming.py index ac375191..da49fb0c 100644 --- a/tests/special_sanity/check_example_naming.py +++ b/tests/special_sanity/check_example_naming.py @@ -43,12 +43,7 @@ STANDALONE_SUFFIX = ("separate", "training") DEFAULT_IGNORE_DIRS: tuple[str, ...] = () -# Dedicated TQ role launchers are lifecycle utilities, not rollout/training -# examples, so the model/drafter/rollout naming grammar does not apply. -DEFAULT_IGNORE_FILES: tuple[str, ...] = ( - "examples/run_dspark_tq_owner.sh", - "examples/run_dspark_tq_producer.sh", -) +DEFAULT_IGNORE_FILES: tuple[str, ...] = () def _split_tokens(stem: str) -> list[str]: diff --git a/tools/run_dspark_tq_consumer.sh b/tools/run_dspark_tq_consumer.sh deleted file mode 100644 index faeccdaa..00000000 --- a/tools/run_dspark_tq_consumer.sh +++ /dev/null @@ -1,63 +0,0 @@ -#!/usr/bin/env bash -# Copyright 2026 Bytedance Ltd. and/or its affiliates -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -set -euo pipefail -set -x - -# Standalone DSpark Consumer. Start Ray and verl_speco.tq_owner first, then run -# the Producer with the same Ray address, namespace, partition and run ID. -MODEL_PATH=${MODEL_PATH:-/path/to/Qwen3-8B} -DRAFTER_PATH=${DRAFTER_PATH:-/path/to/dspark-drafter} -DRAFT_CKPTS_DIR=${DRAFT_CKPTS_DIR:-/path/to/dspark-tq-checkpoints} -TRAIN_DEVICES=${TRAIN_DEVICES:-0,1,2,3} -TRAIN_GPUS=${TRAIN_GPUS:-4} -BATCH_SIZE_PER_GPU=${BATCH_SIZE_PER_GPU:-2} -MAX_STEPS=${MAX_STEPS:-1000} -DSPARK_NUM_TARGET_LAYERS=${DSPARK_NUM_TARGET_LAYERS:-5} -RAY_ADDRESS=${RAY_ADDRESS:-127.0.0.1:6379} -TQ_NAMESPACE=${TQ_NAMESPACE:-speco-drafter} -TQ_PARTITION_ID=${TQ_PARTITION_ID:-speco_drafter_features} -SPECO_TQ_RUN_ID=${SPECO_TQ_RUN_ID:-dspark-standalone-run} -PYTHON_BIN=${PYTHON_BIN:-python3} - -CUDA_VISIBLE_DEVICES=${TRAIN_DEVICES} PYTHONUNBUFFERED=1 \ -exec "${PYTHON_BIN}" -m verl_speco.draft_train_launcher \ - speco.draft_training.nproc_per_node=${TRAIN_GPUS} \ - speco.draft_training.nnodes=1 \ - actor_rollout_ref.model.path=${MODEL_PATH} \ - actor_rollout_ref.actor.strategy=fsdp2 \ - actor_rollout_ref.rollout.drafter.enable=True \ - actor_rollout_ref.rollout.drafter.enable_drafter_training=True \ - actor_rollout_ref.rollout.drafter.model_path=${DRAFTER_PATH} \ - actor_rollout_ref.rollout.drafter.checkpoint_path=${DRAFT_CKPTS_DIR} \ - actor_rollout_ref.rollout.drafter.speculative_algorithm=DSPARK \ - actor_rollout_ref.rollout.drafter.training.mode=offline \ - actor_rollout_ref.rollout.drafter.training.feature_store.type=tq \ - actor_rollout_ref.rollout.drafter.training.feature_store.path=null \ - actor_rollout_ref.rollout.drafter.training.feature_store.shuffle=False \ - actor_rollout_ref.rollout.drafter.training.feature_store.repeat=False \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.enable=True \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.address=${RAY_ADDRESS} \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.namespace=${TQ_NAMESPACE} \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.partition_id=${TQ_PARTITION_ID} \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.run_id=${SPECO_TQ_RUN_ID} \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.drop_last=True \ - actor_rollout_ref.rollout.drafter.training.batch_size_per_gpu=${BATCH_SIZE_PER_GPU} \ - actor_rollout_ref.rollout.drafter.training.max_steps=${MAX_STEPS} \ - actor_rollout_ref.rollout.drafter.training.dspark_num_target_layers=${DSPARK_NUM_TARGET_LAYERS} \ - actor_rollout_ref.rollout.drafter.training.save_interval_steps=100 \ - actor_rollout_ref.rollout.drafter.training.lr=1e-5 \ - actor_rollout_ref.rollout.drafter.training.lr_warmup_steps=50 \ - actor_rollout_ref.rollout.drafter.training.warmup_style=cosine \ - "$@" diff --git a/tools/run_dspark_tq_consumer_test.sh b/tools/run_dspark_tq_consumer_test.sh deleted file mode 100644 index 370d0be8..00000000 --- a/tools/run_dspark_tq_consumer_test.sh +++ /dev/null @@ -1,84 +0,0 @@ -#!/usr/bin/env bash -# Exercise the real standalone DSpark Consumer with delayed synthetic TQ data. -set -euo pipefail - -: "${MODEL_PATH:?Set MODEL_PATH to the target model directory}" -: "${DRAFTER_PATH:?Set DRAFTER_PATH to the DSpark drafter directory}" - -PYTHON_BIN=${PYTHON_BIN:-python3} -RAY_ADDRESS=${RAY_ADDRESS:-127.0.0.1:6379} -TQ_NAMESPACE=${TQ_NAMESPACE:-speco-drafter} -TQ_PARTITION_ID=${TQ_PARTITION_ID:-speco_drafter_features} -SPECO_TQ_RUN_ID=${SPECO_TQ_RUN_ID:-dspark-consumer-test-$$} -TRAIN_DEVICES=${TRAIN_DEVICES:-0} -TRAIN_GPUS=${TRAIN_GPUS:-1} -BATCH_SIZE_PER_GPU=${BATCH_SIZE_PER_GPU:-1} -NUM_BATCHES=${NUM_BATCHES:-3} -SEQUENCE_LENGTH=${SEQUENCE_LENGTH:-64} -DSPARK_NUM_TARGET_LAYERS=${DSPARK_NUM_TARGET_LAYERS:-5} -INITIAL_DELAY_SECONDS=${INITIAL_DELAY_SECONDS:-5} -BATCH_INTERVAL_SECONDS=${BATCH_INTERVAL_SECONDS:-5} -DRAFT_CKPTS_DIR=${DRAFT_CKPTS_DIR:-/tmp/speco-dspark-tq-consumer-test-${SPECO_TQ_RUN_ID}} - -owner_pid="" -producer_pid="" - -cleanup() { - if [[ -n "${producer_pid}" ]] && kill -0 "${producer_pid}" 2>/dev/null; then - kill "${producer_pid}" 2>/dev/null || true - wait "${producer_pid}" 2>/dev/null || true - fi - if [[ -n "${owner_pid}" ]] && kill -0 "${owner_pid}" 2>/dev/null; then - kill "${owner_pid}" 2>/dev/null || true - wait "${owner_pid}" 2>/dev/null || true - fi -} -trap cleanup EXIT INT TERM - -echo "[1/4] Starting TQ owner run_id=${SPECO_TQ_RUN_ID}" -RAY_ADDRESS="${RAY_ADDRESS}" TQ_NAMESPACE="${TQ_NAMESPACE}" \ -TQ_PARTITION_ID="${TQ_PARTITION_ID}" SPECO_TQ_RUN_ID="${SPECO_TQ_RUN_ID}" \ - bash tools/run_dspark_tq_owner.sh & -owner_pid=$! - -echo "[2/4] Starting delayed synthetic Producer" -"${PYTHON_BIN}" tools/tq_delayed_test_producer.py \ - --model-path "${MODEL_PATH}" \ - --ray-address "${RAY_ADDRESS}" \ - --namespace "${TQ_NAMESPACE}" \ - --partition-id "${TQ_PARTITION_ID}" \ - --run-id "${SPECO_TQ_RUN_ID}" \ - --world-size "${TRAIN_GPUS}" \ - --batch-size-per-gpu "${BATCH_SIZE_PER_GPU}" \ - --num-batches "${NUM_BATCHES}" \ - --sequence-length "${SEQUENCE_LENGTH}" \ - --num-target-layers "${DSPARK_NUM_TARGET_LAYERS}" \ - --initial-delay "${INITIAL_DELAY_SECONDS}" \ - --batch-interval "${BATCH_INTERVAL_SECONDS}" & -producer_pid=$! - -echo "[3/4] Running the real standalone Consumer; it should wait between batches" -MODEL_PATH="${MODEL_PATH}" \ -DRAFTER_PATH="${DRAFTER_PATH}" \ -DRAFT_CKPTS_DIR="${DRAFT_CKPTS_DIR}" \ -TRAIN_DEVICES="${TRAIN_DEVICES}" \ -TRAIN_GPUS="${TRAIN_GPUS}" \ -BATCH_SIZE_PER_GPU="${BATCH_SIZE_PER_GPU}" \ -MAX_STEPS="$((NUM_BATCHES + 1))" \ -DSPARK_NUM_TARGET_LAYERS="${DSPARK_NUM_TARGET_LAYERS}" \ -RAY_ADDRESS="${RAY_ADDRESS}" \ -TQ_NAMESPACE="${TQ_NAMESPACE}" \ -TQ_PARTITION_ID="${TQ_PARTITION_ID}" \ -SPECO_TQ_RUN_ID="${SPECO_TQ_RUN_ID}" \ - bash tools/run_dspark_tq_consumer.sh "$@" - -wait "${producer_pid}" -producer_pid="" - -echo "[4/4] Consumer observed EOS and exited; stopping this test's TQ owner" -kill "${owner_pid}" 2>/dev/null || true -wait "${owner_pid}" 2>/dev/null || true -owner_pid="" -trap - EXIT INT TERM - -echo "DSPARK_TQ_CONSUMER_TEST_OK run_id=${SPECO_TQ_RUN_ID} batches=${NUM_BATCHES}" diff --git a/tools/run_dspark_tq_e2e_test.sh b/tools/run_dspark_tq_e2e_test.sh deleted file mode 100644 index eef3c026..00000000 --- a/tools/run_dspark_tq_e2e_test.sh +++ /dev/null @@ -1,170 +0,0 @@ -#!/usr/bin/env bash -# Copyright 2026 Bytedance Ltd. and/or its affiliates -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -# Run the real standalone TQ Producer and Consumer in one end-to-end test. -set -euo pipefail - -script_dir=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd) -repo_root=$(cd -- "${script_dir}/.." && pwd) -cd "${repo_root}" - -: "${RAY_ADDRESS:?Set RAY_ADDRESS to the running Ray head}" -: "${MODEL_PATH:?Set MODEL_PATH to the target model directory}" -: "${DRAFTER_PATH:?Set DRAFTER_PATH to the DSpark drafter directory}" -: "${PRODUCER_INPUT_PATH:?Set PRODUCER_INPUT_PATH to prompt/response JSONL}" -: "${TARGET_MODEL_REVISION:?Set TARGET_MODEL_REVISION to a revision or checksum}" -: "${TOKENIZER_FINGERPRINT:?Set TOKENIZER_FINGERPRINT to a verified fingerprint}" -: "${TARGET_LAYER_IDS:?Set TARGET_LAYER_IDS as a Hydra list, for example '[2,8,14,20,26]'}" -: "${VLLM_ENDPOINTS:?Set VLLM_ENDPOINTS as a Hydra list, for example '[http://node0:8000/v1]'}" - -PYTHON_BIN=${PYTHON_BIN:-python3} -TOKENIZER_PATH=${TOKENIZER_PATH:-${MODEL_PATH}} -TARGET_MODEL_PATH=${TARGET_MODEL_PATH:-${MODEL_PATH}} -TQ_NAMESPACE=${TQ_NAMESPACE:-speco-drafter} -TQ_PARTITION_ID=${TQ_PARTITION_ID:-speco_drafter_features} -SPECO_TQ_RUN_ID=${SPECO_TQ_RUN_ID:-dspark-e2e-$(date +%Y%m%d-%H%M%S)-$$} -TRAIN_DEVICES=${TRAIN_DEVICES:-0} -TRAIN_GPUS=${TRAIN_GPUS:-1} -BATCH_SIZE_PER_GPU=${BATCH_SIZE_PER_GPU:-1} -DSPARK_NUM_TARGET_LAYERS=${DSPARK_NUM_TARGET_LAYERS:-5} -DRAFT_CKPTS_DIR=${DRAFT_CKPTS_DIR:-/tmp/speco-dspark-tq-e2e-${SPECO_TQ_RUN_ID}} -E2E_TIMEOUT_SECONDS=${E2E_TIMEOUT_SECONDS:-1800} - -if [[ "${TQ_PARTITION_ID}" != "speco_drafter_features" ]]; then - echo "TQ_PARTITION_ID must be speco_drafter_features for protocol v1" >&2 - exit 2 -fi -if [[ ! -f "${PRODUCER_INPUT_PATH}" ]]; then - echo "Producer input does not exist: ${PRODUCER_INPUT_PATH}" >&2 - exit 2 -fi -input_samples=$(awk 'NF { count++ } END { print count + 0 }' "${PRODUCER_INPUT_PATH}") -global_batch_size=$((TRAIN_GPUS * BATCH_SIZE_PER_GPU)) -if (( input_samples < global_batch_size )); then - echo "Producer input has ${input_samples} non-empty records, but one Consumer global batch needs ${global_batch_size}" >&2 - exit 2 -fi - -work_dir=$(mktemp -d "${TMPDIR:-/tmp}/speco-tq-e2e.XXXXXX") -owner_pid="" -consumer_pid="" -producer_pid="" - -cleanup() { - local pid - for pid in "${producer_pid}" "${consumer_pid}" "${owner_pid}"; do - if [[ -n "${pid}" ]] && kill -0 "${pid}" 2>/dev/null; then - kill "${pid}" 2>/dev/null || true - wait "${pid}" 2>/dev/null || true - fi - done - rm -rf -- "${work_dir}" -} -trap cleanup EXIT INT TERM - -echo "[1/4] Starting TQ owner run_id=${SPECO_TQ_RUN_ID}" -RAY_ADDRESS="${RAY_ADDRESS}" \ -TQ_NAMESPACE="${TQ_NAMESPACE}" \ -TQ_PARTITION_ID="${TQ_PARTITION_ID}" \ -SPECO_TQ_RUN_ID="${SPECO_TQ_RUN_ID}" \ - bash tools/run_dspark_tq_owner.sh & -owner_pid=$! - -echo "[2/4] Starting real DSpark Consumer" -( - set +e - MODEL_PATH="${MODEL_PATH}" \ - DRAFTER_PATH="${DRAFTER_PATH}" \ - DRAFT_CKPTS_DIR="${DRAFT_CKPTS_DIR}" \ - TRAIN_DEVICES="${TRAIN_DEVICES}" \ - TRAIN_GPUS="${TRAIN_GPUS}" \ - BATCH_SIZE_PER_GPU="${BATCH_SIZE_PER_GPU}" \ - MAX_STEPS=0 \ - DSPARK_NUM_TARGET_LAYERS="${DSPARK_NUM_TARGET_LAYERS}" \ - RAY_ADDRESS="${RAY_ADDRESS}" \ - TQ_NAMESPACE="${TQ_NAMESPACE}" \ - TQ_PARTITION_ID="${TQ_PARTITION_ID}" \ - SPECO_TQ_RUN_ID="${SPECO_TQ_RUN_ID}" \ - bash tools/run_dspark_tq_consumer.sh "$@" - echo "$?" > "${work_dir}/consumer.status" -) & -consumer_pid=$! - -echo "[3/4] Starting real vLLM-backed Producer" -( - set +e - PYTHON_BIN="${PYTHON_BIN}" \ - RAY_ADDRESS="${RAY_ADDRESS}" \ - TQ_NAMESPACE="${TQ_NAMESPACE}" \ - TQ_PARTITION_ID="${TQ_PARTITION_ID}" \ - SPECO_TQ_RUN_ID="${SPECO_TQ_RUN_ID}" \ - PRODUCER_INPUT_PATH="${PRODUCER_INPUT_PATH}" \ - TARGET_MODEL_PATH="${TARGET_MODEL_PATH}" \ - TARGET_MODEL_REVISION="${TARGET_MODEL_REVISION}" \ - TOKENIZER_PATH="${TOKENIZER_PATH}" \ - TOKENIZER_FINGERPRINT="${TOKENIZER_FINGERPRINT}" \ - TARGET_LAYER_IDS="${TARGET_LAYER_IDS}" \ - VLLM_ENDPOINTS="${VLLM_ENDPOINTS}" \ - bash examples/run_dspark_tq_producer.sh - echo "$?" > "${work_dir}/producer.status" -) & -producer_pid=$! - -started_at=${SECONDS} -while [[ ! -f "${work_dir}/producer.status" || ! -f "${work_dir}/consumer.status" ]]; do - if ! kill -0 "${owner_pid}" 2>/dev/null; then - echo "TQ owner exited before the end-to-end test completed" >&2 - exit 1 - fi - if (( SECONDS - started_at >= E2E_TIMEOUT_SECONDS )); then - echo "E2E timed out after ${E2E_TIMEOUT_SECONDS} seconds" >&2 - exit 124 - fi - if [[ -f "${work_dir}/producer.status" ]]; then - producer_status=$(<"${work_dir}/producer.status") - if [[ "${producer_status}" -ne 0 ]]; then - echo "Producer failed with exit code ${producer_status}" >&2 - exit "${producer_status}" - fi - fi - if [[ -f "${work_dir}/consumer.status" ]]; then - consumer_status=$(<"${work_dir}/consumer.status") - if [[ "${consumer_status}" -ne 0 ]]; then - echo "Consumer failed with exit code ${consumer_status}" >&2 - exit "${consumer_status}" - fi - fi - sleep 1 -done - -producer_status=$(<"${work_dir}/producer.status") -consumer_status=$(<"${work_dir}/consumer.status") -if [[ "${producer_status}" -ne 0 || "${consumer_status}" -ne 0 ]]; then - echo "E2E failed: producer=${producer_status} consumer=${consumer_status}" >&2 - exit 1 -fi - -wait "${producer_pid}" -producer_pid="" -wait "${consumer_pid}" -consumer_pid="" - -echo "[4/4] Producer published EOS and Consumer drained the run; stopping owner" -kill "${owner_pid}" 2>/dev/null || true -wait "${owner_pid}" 2>/dev/null || true -owner_pid="" -trap - EXIT INT TERM -rm -rf -- "${work_dir}" - -echo "DSPARK_TQ_E2E_TEST_OK run_id=${SPECO_TQ_RUN_ID} checkpoints=${DRAFT_CKPTS_DIR}" diff --git a/tools/run_dspark_tq_owner.sh b/tools/run_dspark_tq_owner.sh deleted file mode 100644 index c152407b..00000000 --- a/tools/run_dspark_tq_owner.sh +++ /dev/null @@ -1,29 +0,0 @@ -#!/usr/bin/env bash -# Copyright 2026 Bytedance Ltd. and/or its affiliates -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -set -euo pipefail - -: "${RAY_ADDRESS:?Set RAY_ADDRESS to the running Ray head, for example 10.0.0.1:6379}" -: "${SPECO_TQ_RUN_ID:?Set SPECO_TQ_RUN_ID to a unique pipeline run id}" -TQ_NAMESPACE=${TQ_NAMESPACE:-speco-drafter} -TQ_PARTITION_ID=${TQ_PARTITION_ID:-speco_drafter_features} -PYTHON_BIN=${PYTHON_BIN:-python3} - -exec "${PYTHON_BIN}" -m verl_speco.tq_owner \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.address="${RAY_ADDRESS}" \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.namespace="${TQ_NAMESPACE}" \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.partition_id="${TQ_PARTITION_ID}" \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.run_id="${SPECO_TQ_RUN_ID}" \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.backend.storage_backend=SimpleStorage \ - "$@" diff --git a/tools/run_qwen3-4b_drafter_dspark_mooncake.sh b/tools/run_qwen3-4b_drafter_dspark_mooncake.sh deleted file mode 100644 index 10398a51..00000000 --- a/tools/run_qwen3-4b_drafter_dspark_mooncake.sh +++ /dev/null @@ -1,72 +0,0 @@ -set -euo pipefail -set -x - -# Run each stage in a separate shell: RUN_STAGE=master, vllm, then train. -# On Ascend replace CUDA_VISIBLE_DEVICES below with ASCEND_RT_VISIBLE_DEVICES. -RUN_STAGE=${RUN_STAGE:-train} -MODEL_PATH=${MODEL_PATH:-/path/to/Qwen3-4B} -DATA_PATH=${DATA_PATH:-/path/to/token_replay.jsonl} -DRAFT_CKPTS_DIR=${DRAFT_CKPTS_DIR:-/path/to/draft_checkpoints} -VLLM_DEVICES=${VLLM_DEVICES:-0,1} -TRAIN_DEVICES=${TRAIN_DEVICES:-2,3,4,5} -VLLM_TP=${VLLM_TP:-2} -TRAIN_GPUS=${TRAIN_GPUS:-4} - -export MOONCAKE_MASTER_SERVER=${MOONCAKE_MASTER_SERVER:-127.0.0.1:50051} -export MOONCAKE_METADATA_SERVER=${MOONCAKE_METADATA_SERVER:-http://127.0.0.1:8090/metadata} -export MOONCAKE_PROTOCOL=${MOONCAKE_PROTOCOL:-tcp} -export MOONCAKE_GLOBAL_SEGMENT_SIZE=${MOONCAKE_GLOBAL_SEGMENT_SIZE:-17179869184} -export MOONCAKE_LOCAL_BUFFER_SIZE=${MOONCAKE_LOCAL_BUFFER_SIZE:-2147483648} - -if [ "${RUN_STAGE}" = "master" ]; then - exec mooncake_master \ - --enable_http_metadata_server=true \ - --http_metadata_server_host=0.0.0.0 \ - --http_metadata_server_port=8090 -fi - -if [ "${RUN_STAGE}" = "vllm" ]; then - CUDA_VISIBLE_DEVICES=${VLLM_DEVICES} exec vllm serve "${MODEL_PATH}" \ - --host 0.0.0.0 --port 8000 \ - --tensor-parallel-size "${VLLM_TP}" \ - --gpu-memory-utilization 0.85 \ - --speculative-config '{"method":"extract_hidden_states","num_speculative_tokens":1,"draft_model_config":{"hf_config":{"eagle_aux_hidden_state_layer_ids":[1,9,17,25,33,36]}}}' \ - --kv-transfer-config '{"kv_connector":"SpeCoMooncakeHiddenStatesConnector","kv_connector_module_path":"verl_speco.integration.mooncake_hidden_states_connector","kv_role":"kv_producer"}' \ - --no-enable-chunked-prefill -fi - -if [ "${RUN_STAGE}" != "train" ]; then - echo "RUN_STAGE must be master, vllm, or train" >&2 - exit 2 -fi - -CUDA_VISIBLE_DEVICES=${TRAIN_DEVICES} PYTHONUNBUFFERED=1 \ -python3 -m verl_speco.draft_train_launcher \ - speco.draft_training.nproc_per_node=${TRAIN_GPUS} \ - speco.draft_training.nnodes=1 \ - actor_rollout_ref.model.path=${MODEL_PATH} \ - actor_rollout_ref.actor.strategy=fsdp2 \ - actor_rollout_ref.rollout.drafter.enable=True \ - actor_rollout_ref.rollout.drafter.enable_drafter_training=True \ - actor_rollout_ref.rollout.drafter.checkpoint_path=${DRAFT_CKPTS_DIR} \ - actor_rollout_ref.rollout.drafter.speculative_algorithm=DSPARK \ - actor_rollout_ref.rollout.drafter.training.mode=offline \ - actor_rollout_ref.rollout.drafter.training.feature_store.type=jsonl_token_replay \ - actor_rollout_ref.rollout.drafter.training.feature_store.path=${DATA_PATH} \ - actor_rollout_ref.rollout.drafter.training.feature_store.shuffle=True \ - actor_rollout_ref.rollout.drafter.training.feature_store.repeat=True \ - actor_rollout_ref.rollout.drafter.training.target_feature_replay.backend=vllm_mooncake \ - actor_rollout_ref.rollout.drafter.training.target_feature_replay.vllm_endpoint=http://127.0.0.1:8000/v1 \ - actor_rollout_ref.rollout.drafter.training.target_feature_replay.on_generate=delete \ - actor_rollout_ref.rollout.drafter.training.target_feature_pipeline.enabled=True \ - actor_rollout_ref.rollout.drafter.training.target_feature_pipeline.concurrency=16 \ - actor_rollout_ref.rollout.drafter.training.target_feature_pipeline.transfer_concurrency=8 \ - actor_rollout_ref.rollout.drafter.training.target_feature_pipeline.producer_prefetch_depth=4 \ - actor_rollout_ref.rollout.drafter.training.target_feature_pipeline.prefetch_depth=2 \ - actor_rollout_ref.rollout.drafter.training.batch_size_per_gpu=4 \ - actor_rollout_ref.rollout.drafter.training.max_steps=1000 \ - actor_rollout_ref.rollout.drafter.training.save_interval_steps=100 \ - actor_rollout_ref.rollout.drafter.training.lr=1e-5 \ - actor_rollout_ref.rollout.drafter.training.lr_warmup_steps=50 \ - actor_rollout_ref.rollout.drafter.training.warmup_style=cosine \ - "$@" diff --git a/examples/run_qwen3-8b_drafter_hidden_state_vllm.sh b/tools/run_qwen3-8b_drafter_hidden_state_vllm.sh similarity index 50% rename from examples/run_qwen3-8b_drafter_hidden_state_vllm.sh rename to tools/run_qwen3-8b_drafter_hidden_state_vllm.sh index c318d306..239e4b18 100644 --- a/examples/run_qwen3-8b_drafter_hidden_state_vllm.sh +++ b/tools/run_qwen3-8b_drafter_hidden_state_vllm.sh @@ -19,20 +19,20 @@ set -x # Run this script in its own terminal before starting the training script. # # Ascend example: -# DEVICE_ENV=ASCEND_RT_VISIBLE_DEVICES VLLM_DEVICES_0=0,1 VLLM_DEVICES_1=2,3 VLLM_TP=2 \ -# bash examples/run_qwen3-8b_drafter_hidden_state_vllm.sh +# Set DEVICE_ENV=ASCEND_RT_VISIBLE_DEVICES, VLLM_DEVICES and VLLM_TP below, +# then run: bash tools/run_qwen3-8b_drafter_hidden_state_vllm.sh # CUDA example: -# DEVICE_ENV=CUDA_VISIBLE_DEVICES VLLM_DEVICES_0=0,1 VLLM_DEVICES_1=2,3 VLLM_TP=2 \ -# bash examples/run_qwen3-8b_drafter_hidden_state_vllm.sh +# Set DEVICE_ENV=CUDA_VISIBLE_DEVICES, VLLM_DEVICES and VLLM_TP below, +# then run: bash tools/run_qwen3-8b_drafter_hidden_state_vllm.sh MODEL_PATH=${MODEL_PATH:-/path/to/Qwen3-8B} DEVICE_ENV=${DEVICE_ENV:-ASCEND_RT_VISIBLE_DEVICES} -VLLM_DEVICES_0=${VLLM_DEVICES_0:-0} -VLLM_DEVICES_1=${VLLM_DEVICES_1:-1} +# Devices assigned to target-model vLLM. The script splits this list into +# consecutive groups of VLLM_TP devices and starts one service per group. +VLLM_DEVICES=${VLLM_DEVICES:-0,1,2,3,4,5} VLLM_TP=${VLLM_TP:-1} VLLM_HOST=${VLLM_HOST:-127.0.0.1} -VLLM_PORT_0=${VLLM_PORT_0:-8000} -VLLM_PORT_1=${VLLM_PORT_1:-8001} +VLLM_BASE_PORT=${VLLM_BASE_PORT:-8000} VLLM_GPU_MEMORY_UTILIZATION=${VLLM_GPU_MEMORY_UTILIZATION:-0.8} VLLM_MAX_NUM_SEQS=${VLLM_MAX_NUM_SEQS:-256} # Auxiliary training layers followed by the target model's final hidden-state @@ -41,7 +41,32 @@ VLLM_MAX_NUM_SEQS=${VLLM_MAX_NUM_SEQS:-256} VLLM_HIDDEN_STATE_LAYER_IDS=${VLLM_HIDDEN_STATE_LAYER_IDS:-'[1,9,17,25,33,36]'} HIDDEN_STATES_DIR=${HIDDEN_STATES_DIR:-/tmp/speco-vllm-hidden-states} -mkdir -p "${HIDDEN_STATES_DIR}/service-0" "${HIDDEN_STATES_DIR}/service-1" +if ! [[ "${VLLM_TP}" =~ ^[1-9][0-9]*$ ]]; then + echo "VLLM_TP must be a positive integer, got: ${VLLM_TP}" >&2 + exit 2 +fi +if ! [[ "${VLLM_BASE_PORT}" =~ ^[0-9]+$ ]]; then + echo "VLLM_BASE_PORT must be an integer, got: ${VLLM_BASE_PORT}" >&2 + exit 2 +fi + +visible_devices=${VLLM_DEVICES} + +IFS=',' read -r -a DEVICE_IDS <<< "${visible_devices}" +for index in "${!DEVICE_IDS[@]}"; do + DEVICE_IDS[index]=${DEVICE_IDS[index]//[[:space:]]/} + if [[ -z "${DEVICE_IDS[index]}" ]]; then + echo "Visible device list contains an empty item: ${visible_devices}" >&2 + exit 2 + fi +done + +device_count=${#DEVICE_IDS[@]} +if (( device_count % VLLM_TP != 0 )); then + echo "Visible device count (${device_count}) must be divisible by VLLM_TP (${VLLM_TP}): ${visible_devices}" >&2 + exit 2 +fi +service_count=$((device_count / VLLM_TP)) SPECULATIVE_CONFIG=$(printf '{"method":"extract_hidden_states","num_speculative_tokens":1,"draft_model_config":{"hf_config":{"eagle_aux_hidden_state_layer_ids":%s}}}' "${VLLM_HIDDEN_STATE_LAYER_IDS}") @@ -65,24 +90,40 @@ start_vllm() { STARTED_PID=$! } -PID_0="" -PID_1="" +PIDS=() +ENDPOINTS=() cleanup() { - [[ -n "${PID_0}" ]] && kill "${PID_0}" 2>/dev/null || true - [[ -n "${PID_1}" ]] && kill "${PID_1}" 2>/dev/null || true - [[ -n "${PID_0}" ]] && wait "${PID_0}" 2>/dev/null || true - [[ -n "${PID_1}" ]] && wait "${PID_1}" 2>/dev/null || true + if (( ${#PIDS[@]} > 0 )); then + kill "${PIDS[@]}" 2>/dev/null || true + wait "${PIDS[@]}" 2>/dev/null || true + fi } trap cleanup EXIT INT TERM -start_vllm "${VLLM_DEVICES_0}" "${VLLM_PORT_0}" "${HIDDEN_STATES_DIR}/service-0" "$@" -PID_0=${STARTED_PID} -start_vllm "${VLLM_DEVICES_1}" "${VLLM_PORT_1}" "${HIDDEN_STATES_DIR}/service-1" "$@" -PID_1=${STARTED_PID} +for ((service_index = 0; service_index < service_count; service_index++)); do + first_device=$((service_index * VLLM_TP)) + service_devices=${DEVICE_IDS[first_device]} + for ((tp_index = 1; tp_index < VLLM_TP; tp_index++)); do + service_devices+=",${DEVICE_IDS[first_device + tp_index]}" + done + + service_port=$((VLLM_BASE_PORT + service_index)) + service_hidden_states_dir="${HIDDEN_STATES_DIR}/service-${service_index}" + mkdir -p "${service_hidden_states_dir}" + start_vllm \ + "${service_devices}" \ + "${service_port}" \ + "${service_hidden_states_dir}" \ + "$@" + PIDS+=("${STARTED_PID}") + ENDPOINTS+=("http://${VLLM_HOST}:${service_port}/v1") +done -echo "VLLM_SERVICES_STARTED pid_0=${PID_0} endpoint_0=http://${VLLM_HOST}:${VLLM_PORT_0}/v1 pid_1=${PID_1} endpoint_1=http://${VLLM_HOST}:${VLLM_PORT_1}/v1" +endpoint_list=$(IFS=,; echo "[${ENDPOINTS[*]}]") +echo "VLLM_SERVICES_STARTED count=${service_count} tp=${VLLM_TP} devices=${visible_devices} endpoints=${endpoint_list} pids=${PIDS[*]}" +echo "Use this for standalone training: SPECO_VLLM_ENDPOINTS='${endpoint_list}'" set +e -wait -n "${PID_0}" "${PID_1}" +wait -n "${PIDS[@]}" status=$? set -e exit "${status}" diff --git a/tools/run_tq_connection_smoke.sh b/tools/run_tq_connection_smoke.sh deleted file mode 100644 index e7173b21..00000000 --- a/tools/run_tq_connection_smoke.sh +++ /dev/null @@ -1,53 +0,0 @@ -#!/usr/bin/env bash -# Copyright 2026 Bytedance Ltd. and/or its affiliates -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -set -euo pipefail - -# A Ray head must already be running. This script starts only the two smoke -# roles: an owner that also publishes two synthetic samples, and a client that -# reads, validates, clears, and observes EOS. -PYTHON_BIN=${PYTHON_BIN:-python3} -RAY_ADDRESS=${RAY_ADDRESS:-127.0.0.1:6379} -TQ_NAMESPACE=${TQ_NAMESPACE:-speco-drafter-smoke} -SPECO_TQ_RUN_ID=${SPECO_TQ_RUN_ID:-tq-smoke-$$} -SMOKE_TIMEOUT_SECONDS=${SMOKE_TIMEOUT_SECONDS:-60} - -owner_pid="" - -cleanup() { - if [[ -n "${owner_pid}" ]] && kill -0 "${owner_pid}" 2>/dev/null; then - kill "${owner_pid}" 2>/dev/null || true - wait "${owner_pid}" 2>/dev/null || true - fi -} -trap cleanup EXIT INT TERM - -"${PYTHON_BIN}" tools/tq_connection_smoke.py owner \ - --ray-address "${RAY_ADDRESS}" \ - --namespace "${TQ_NAMESPACE}" \ - --run-id "${SPECO_TQ_RUN_ID}" \ - --timeout "${SMOKE_TIMEOUT_SECONDS}" & -owner_pid=$! - -"${PYTHON_BIN}" tools/tq_connection_smoke.py client \ - --ray-address "${RAY_ADDRESS}" \ - --namespace "${TQ_NAMESPACE}" \ - --run-id "${SPECO_TQ_RUN_ID}" \ - --timeout "${SMOKE_TIMEOUT_SECONDS}" - -wait "${owner_pid}" -owner_pid="" -trap - EXIT INT TERM - -echo "TQ_CONNECTION_SMOKE_OK run_id=${SPECO_TQ_RUN_ID}" diff --git a/tools/tq_connection_smoke.py b/tools/tq_connection_smoke.py deleted file mode 100644 index 40dbdee0..00000000 --- a/tools/tq_connection_smoke.py +++ /dev/null @@ -1,175 +0,0 @@ -# Copyright 2026 Bytedance Ltd. and/or its affiliates -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -"""Real two-process smoke test for Ray + TransferQueue 0.1.10. - -Start a Ray head first, then run ``owner`` and ``client`` in separate shells. -The owner publishes one protocol-valid sample; the client list/get/decodes and -clears it, publishes a done marker, and exits without killing the owner. -""" - -from __future__ import annotations - -import argparse -import time - -import torch - -from verl_speco.integration.transferqueue_bridge import ( - close_transfer_queue_owner, - configure_transfer_queue, - connect_ray_cluster, - list_samples, - put_sample, - start_transfer_queue_owner, -) -from verl_speco.trainer.feature_store import DraftFeatureSample -from verl_speco.trainer.tq_feature_store import TQFeatureStore -from verl_speco.trainer.tq_sample_source import TQFeatureDataLoader -from verl_speco.transport.drafter_sample_protocol import ( - PROTOCOL_SCHEMA_VERSION, - SampleMetadata, - encode_sample, - make_eos_record, - make_ready_tag, - make_sample_key, -) - - -def _config(args) -> dict: - return { - "enable": True, - "package_version": "0.1.10", - "ray": {"address": args.ray_address, "namespace": args.namespace}, - "partition_id": "speco_drafter_features", - "run_id": args.run_id, - "schema_version": PROTOCOL_SCHEMA_VERSION, - "controller": {"polling_mode": True}, - "backend": { - "storage_backend": "SimpleStorage", - "SimpleStorage": { - "total_storage_size": 32, - "num_data_storage_units": 1, - }, - }, - } - - -def _record(run_id: str, sequence_no: int): - meta = SampleMetadata( - schema_version=PROTOCOL_SCHEMA_VERSION, - run_id=run_id, - sample_id=f"smoke-{sequence_no:04d}", - sequence_no=sequence_no, - ) - sample = DraftFeatureSample( - algorithm="DSPARK", - input_ids=torch.tensor([1, 2, 3]) + sequence_no, - loss_mask=torch.tensor([0.0, 1.0, 1.0]), - position_ids=torch.tensor([0, 1, 2]), - hidden_states=( - torch.arange(12, dtype=torch.float32).reshape(3, 4) + sequence_no - ), - ) - return meta, sample - - -def _wait_for(predicate, timeout: float, description: str): - deadline = time.monotonic() + timeout - while time.monotonic() < deadline: - value = predicate() - if value: - return value - time.sleep(0.2) - raise TimeoutError(f"Timed out waiting for {description}") - - -def run_owner(args) -> None: - config = _config(args) - configure_transfer_queue(config) - connect_ray_cluster(args.ray_address, args.namespace) - start_transfer_queue_owner(config) - records = [_record(args.run_id, sequence_no) for sequence_no in range(2)] - keys = [make_sample_key(meta) for meta, _ in records] - try: - owner_ready_key = f"control:v{PROTOCOL_SCHEMA_VERSION}:{args.run_id}:owner-ready" - put_sample( - owner_ready_key, - {"marker": torch.tensor([1], dtype=torch.uint8)}, - tag={ - "record_type": "control", - "status": "owner_ready", - "schema_version": PROTOCOL_SCHEMA_VERSION, - "run_id": args.run_id, - }, - ) - for (meta, sample), key in zip(records, keys, strict=True): - put_sample(key, encode_sample(sample, meta), tag=make_ready_tag(meta)) - eos_key, eos_fields, eos_tag = make_eos_record(args.run_id, len(records)) - put_sample(eos_key, eos_fields, tag=eos_tag) - print(f"OWNER_READY keys={keys}", flush=True) - _wait_for( - lambda: all(key not in list_samples() for key in keys), - args.timeout, - "client to clear the smoke sample keys", - ) - print("OWNER_OBSERVED_SAMPLES_CLEARED", flush=True) - finally: - close_transfer_queue_owner() - print("OWNER_CLOSED", flush=True) - - -def run_client(args) -> None: - config = _config(args) - store = TQFeatureStore.from_config(config) - loader = TQFeatureDataLoader(store, batch_size=2, rank=0, world_size=1) - try: - iterator = iter(loader) - batch = next(iterator) - restored = batch.local_samples - assert restored[0].hidden_states.tolist() == torch.arange(12).reshape(3, 4).tolist() - assert restored[1].hidden_states.tolist() == ( - torch.arange(12).reshape(3, 4) + 1 - ).tolist() - loader.clear_completed_batch(batch.global_keys) - try: - next(iterator) - except StopIteration: - pass - else: - raise AssertionError("TQ Consumer did not stop after draining EOS") - print( - f"CLIENT_OK samples={len(restored)} shape={tuple(restored[0].hidden_states.shape)}", - flush=True, - ) - finally: - store.close_local() - print("CLIENT_CLOSED_LOCAL_ONLY", flush=True) - - -def main() -> None: - parser = argparse.ArgumentParser() - parser.add_argument("role", choices=("owner", "client")) - parser.add_argument("--ray-address", required=True) - parser.add_argument("--namespace", default="speco-drafter-smoke") - parser.add_argument("--run-id", default="tq-smoke") - parser.add_argument("--timeout", type=float, default=60.0) - args = parser.parse_args() - if args.role == "owner": - run_owner(args) - else: - run_client(args) - - -if __name__ == "__main__": - main() diff --git a/tools/tq_delayed_test_producer.py b/tools/tq_delayed_test_producer.py deleted file mode 100644 index af9e3af1..00000000 --- a/tools/tq_delayed_test_producer.py +++ /dev/null @@ -1,222 +0,0 @@ -# Copyright 2026 Bytedance Ltd. and/or its affiliates -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -"""Publish protocol-valid synthetic DSpark features in delayed batches. - -This is an integration-test Producer for the standalone TQ Consumer. It does -not run vLLM: tensor shapes are derived from the target model config so the -normal DSpark preprocessing and training path can consume the samples. -""" - -from __future__ import annotations - -import argparse -import json -import time -from pathlib import Path -from typing import Any - -import torch - -from verl_speco.integration.transferqueue_bridge import ( - close_transfer_queue_client, - configure_transfer_queue, - connect_ray_cluster, - connect_transfer_queue_client, - list_samples, - put_sample, -) -from verl_speco.trainer.feature_store import DraftFeatureSample -from verl_speco.transport.drafter_sample_protocol import ( - PROTOCOL_SCHEMA_VERSION, - SampleMetadata, - encode_sample, - make_eos_record, - make_ready_tag, - make_sample_key, -) - - -def _tq_config(args: argparse.Namespace) -> dict[str, Any]: - return { - "enable": True, - "package_version": "0.1.10", - "ray": {"address": args.ray_address, "namespace": args.namespace}, - "partition_id": args.partition_id, - "run_id": args.run_id, - "schema_version": PROTOCOL_SCHEMA_VERSION, - "controller": {"polling_mode": True}, - "backend": { - "storage_backend": "SimpleStorage", - "SimpleStorage": { - "total_storage_size": 100000, - "num_data_storage_units": 8, - }, - }, - } - - -def _model_dimensions(model_path: str) -> tuple[int, int, int]: - config_path = Path(model_path) / "config.json" - with config_path.open("r", encoding="utf-8") as handle: - config = json.load(handle) - text_config = config.get("text_config") or config - hidden_size = int(text_config["hidden_size"]) - vocab_size = int(text_config["vocab_size"]) - num_hidden_layers = int(text_config.get("num_hidden_layers", 1)) - return hidden_size, vocab_size, num_hidden_layers - - -def _wait_for_owner(args: argparse.Namespace) -> None: - owner_ready_key = f"control:v{PROTOCOL_SCHEMA_VERSION}:{args.run_id}:owner-ready" - deadline = time.monotonic() + args.timeout - while time.monotonic() < deadline: - tag = list_samples().get(owner_ready_key) - if isinstance(tag, dict) and tag.get("status") == "owner_ready": - return - time.sleep(0.2) - raise TimeoutError(f"Timed out waiting for TQ owner key {owner_ready_key!r}") - - -def _target_layer_ids(num_target_layers: int, num_hidden_layers: int) -> list[int]: - if num_target_layers <= 0: - raise ValueError("num_target_layers must be positive") - if num_hidden_layers < 4: - raise ValueError("the DSpark target model must have at least four hidden layers") - if num_target_layers == 1: - return [num_hidden_layers // 2] - start = 1 - end = num_hidden_layers - 3 - span = end - start - return [ - int(round(start + (index * span) / (num_target_layers - 1))) - for index in range(num_target_layers) - ] - - -def _sample( - args: argparse.Namespace, - *, - sequence_no: int, - hidden_size: int, - vocab_size: int, - target_layer_ids: list[int], -) -> tuple[str, dict[str, torch.Tensor], dict[str, Any]]: - generator = torch.Generator(device="cpu") - generator.manual_seed(args.seed + sequence_no) - # The default standalone DSpark configuration enables L1 distillation. - # Its wire layout contains N auxiliary target layers followed by the - # target model's final hidden state, all concatenated on the last axis. - hidden_dim = hidden_size * (len(target_layer_ids) + 1) - input_ids = torch.randint( - low=0, - high=vocab_size, - size=(args.sequence_length,), - generator=generator, - dtype=torch.long, - ) - loss_mask = torch.ones(args.sequence_length, dtype=torch.float32) - loss_mask[: max(1, args.sequence_length // 4)] = 0 - hidden_states = torch.randn( - args.sequence_length, - hidden_dim, - generator=generator, - dtype=torch.bfloat16, - ) - sample_id = f"consumer-test-{sequence_no:08d}" - metadata = SampleMetadata( - schema_version=PROTOCOL_SCHEMA_VERSION, - run_id=args.run_id, - sample_id=sample_id, - sequence_no=sequence_no, - ) - sample = DraftFeatureSample( - algorithm="DSPARK", - input_ids=input_ids, - loss_mask=loss_mask, - position_ids=torch.arange(args.sequence_length, dtype=torch.long), - hidden_states=hidden_states, - metadata={"hidden_states_layout": "dflash_aux_plus_last"}, - ) - key = make_sample_key(metadata) - return key, encode_sample(sample, metadata), make_ready_tag(metadata) - - -def run(args: argparse.Namespace) -> None: - hidden_size, vocab_size, num_hidden_layers = _model_dimensions(args.model_path) - target_layer_ids = _target_layer_ids(args.num_target_layers, num_hidden_layers) - samples_per_batch = args.world_size * args.batch_size_per_gpu - total_samples = args.num_batches * samples_per_batch - config = _tq_config(args) - configure_transfer_queue(config) - connect_ray_cluster(args.ray_address, args.namespace) - connect_transfer_queue_client() - try: - _wait_for_owner(args) - print( - "PRODUCER_CONNECTED " - f"batches={args.num_batches} samples_per_batch={samples_per_batch} " - f"sequence_length={args.sequence_length} hidden_shape=" - f"({args.sequence_length}, {hidden_size * (args.num_target_layers + 1)})", - flush=True, - ) - time.sleep(args.initial_delay) - sequence_no = 0 - for batch_index in range(args.num_batches): - keys = [] - for _ in range(samples_per_batch): - key, fields, tag = _sample( - args, - sequence_no=sequence_no, - hidden_size=hidden_size, - vocab_size=vocab_size, - target_layer_ids=target_layer_ids, - ) - put_sample(key, fields, tag=tag) - keys.append(key) - sequence_no += 1 - print( - f"PRODUCER_BATCH_READY batch={batch_index + 1}/{args.num_batches} " - f"samples={len(keys)} sequence_no=[{sequence_no - len(keys)},{sequence_no})", - flush=True, - ) - if batch_index + 1 < args.num_batches: - time.sleep(args.batch_interval) - eos_key, eos_fields, eos_tag = make_eos_record(args.run_id, total_samples) - put_sample(eos_key, eos_fields, tag=eos_tag) - print(f"PRODUCER_EOS total_samples={total_samples} key={eos_key}", flush=True) - finally: - close_transfer_queue_client() - print("PRODUCER_CLOSED_LOCAL", flush=True) - - -def main() -> None: - parser = argparse.ArgumentParser(description=__doc__) - parser.add_argument("--model-path", required=True) - parser.add_argument("--ray-address", required=True) - parser.add_argument("--namespace", default="speco-drafter") - parser.add_argument("--partition-id", default="speco_drafter_features") - parser.add_argument("--run-id", required=True) - parser.add_argument("--world-size", type=int, default=1) - parser.add_argument("--batch-size-per-gpu", type=int, default=1) - parser.add_argument("--num-batches", type=int, default=3) - parser.add_argument("--sequence-length", type=int, default=64) - parser.add_argument("--num-target-layers", type=int, default=5) - parser.add_argument("--initial-delay", type=float, default=5.0) - parser.add_argument("--batch-interval", type=float, default=5.0) - parser.add_argument("--timeout", type=float, default=120.0) - parser.add_argument("--seed", type=int, default=2026) - args = parser.parse_args() - if args.world_size <= 0 or args.batch_size_per_gpu <= 0: - parser.error("world-size and batch-size-per-gpu must be positive") - if args.num_batches <= 0 or args.sequence_length <= 0: - parser.error("num-batches and sequence-length must be positive") - run(args) - - -if __name__ == "__main__": - main() diff --git a/tools/wait_for_vllm_endpoints.py b/tools/wait_for_vllm_endpoints.py new file mode 100644 index 00000000..a7a3aac4 --- /dev/null +++ b/tools/wait_for_vllm_endpoints.py @@ -0,0 +1,92 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Wait until every OpenAI-compatible vLLM endpoint is ready.""" + +from __future__ import annotations + +import argparse +import time +from urllib.error import HTTPError, URLError +from urllib.request import urlopen + + +def parse_endpoint_list(value: str) -> list[str]: + """Parse the Hydra-style ``[url0,url1]`` used by the training launcher.""" + + raw = value.strip() + if not (raw.startswith("[") and raw.endswith("]")): + raise ValueError("endpoints must use [url0,url1] syntax") + endpoints = [ + item.strip().strip("'\"").rstrip("/") + for item in raw[1:-1].split(",") + if item.strip() + ] + if not endpoints: + raise ValueError("endpoints must contain at least one URL") + return endpoints + + +def wait_for_endpoints( + endpoints: list[str], + *, + timeout_seconds: float, + poll_interval_seconds: float, + request_timeout_seconds: float, +) -> None: + deadline = time.monotonic() + timeout_seconds + pending = set(endpoints) + while pending: + for endpoint in list(pending): + try: + with urlopen( + f"{endpoint}/models", timeout=request_timeout_seconds + ) as response: + if 200 <= int(response.status) < 300: + print(f"EXTERNAL_VLLM_READY endpoint={endpoint}", flush=True) + pending.remove(endpoint) + except (HTTPError, OSError, URLError): + pass + if not pending: + return + if time.monotonic() >= deadline: + raise TimeoutError( + "external hidden-state vLLM is not ready at: " + + ", ".join(sorted(pending)) + ) + time.sleep(poll_interval_seconds) + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--endpoints", required=True) + parser.add_argument("--timeout-seconds", type=float, default=120.0) + parser.add_argument("--poll-interval-seconds", type=float, default=1.0) + parser.add_argument("--request-timeout-seconds", type=float, default=2.0) + args = parser.parse_args() + + try: + endpoints = parse_endpoint_list(args.endpoints) + wait_for_endpoints( + endpoints, + timeout_seconds=args.timeout_seconds, + poll_interval_seconds=args.poll_interval_seconds, + request_timeout_seconds=args.request_timeout_seconds, + ) + except (TimeoutError, ValueError) as exc: + parser.error(str(exc)) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/verl_speco/backends/dflash_trainer_backend.py b/verl_speco/backends/dflash_trainer_backend.py index aa0f0026..cec76210 100644 --- a/verl_speco/backends/dflash_trainer_backend.py +++ b/verl_speco/backends/dflash_trainer_backend.py @@ -573,13 +573,9 @@ def forward( correct_3d = correct.view(bsz, n_blocks, self.block_size) pred_valid_3d = binary_weights[:, :, 1:].bool() pred_correct_3d = correct_3d[:, :, 1:] & pred_valid_3d - simulated_accept_length_sum = ( - pred_correct_3d.float().cumprod(dim=-1).sum() - ) + simulated_accept_length_sum = pred_correct_3d.float().cumprod(dim=-1).sum() simulated_accept_block_count = pred_valid_3d.any(dim=-1).float().sum() - correct_per_position = ( - correct_3d.float().sum(dim=(0, 1)) - ) + correct_per_position = correct_3d.float().sum(dim=(0, 1)) loss_per_position = loss_sum_per_position / count_per_pos acc_per_position = correct_per_position / count_per_pos masked_rows = (binary_eval_mask <= 0.5).sum().to(dtype=torch.float32) diff --git a/verl_speco/draft_train_launcher.py b/verl_speco/draft_train_launcher.py index 087e8a5a..1cde487a 100644 --- a/verl_speco/draft_train_launcher.py +++ b/verl_speco/draft_train_launcher.py @@ -57,7 +57,9 @@ "speco.draft_training.standalone", "actor_rollout_ref.rollout.drafter.training.standalone", ) -_FEATURE_STORE_TYPE_KEY = "actor_rollout_ref.rollout.drafter.training.feature_store.type" +_FEATURE_STORE_TYPE_KEY = ( + "actor_rollout_ref.rollout.drafter.training.feature_store.type" +) _TQ_ENABLE_KEY = "actor_rollout_ref.rollout.drafter.training.transfer_queue.enable" _TQ_RAY_ADDRESS_KEY = ( "actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.address" @@ -190,7 +192,9 @@ def validate_tq_launch_config(overrides: list[str]) -> None: ) run_id = str(_find_override(overrides, (_TQ_RUN_ID_KEY,)) or "").strip() if not run_id or run_id.lower() in {"null", "none"}: - raise ValueError("feature_store.type=tq requires training.transfer_queue.run_id") + raise ValueError( + "feature_store.type=tq requires training.transfer_queue.run_id" + ) def build_torch_distributed_command( diff --git a/verl_speco/inspect_jsonl_samples.py b/verl_speco/inspect_jsonl_samples.py index 217d638d..4af85183 100644 --- a/verl_speco/inspect_jsonl_samples.py +++ b/verl_speco/inspect_jsonl_samples.py @@ -59,7 +59,7 @@ def main() -> int: args = parser.parse_args() path = Path(args.path) - summaries = [] + summaries: list[dict[str, Any]] = [] key_counts: Counter[str] = Counter() schema_counts: Counter[str] = Counter() issues_by_schema: dict[str, list[str]] = defaultdict(list) @@ -230,7 +230,7 @@ def _shape(value: Any) -> list[Any]: if not isinstance(value, list): return ["scalar", type(value).__name__] shape = [] - current = value + current: Any = value while isinstance(current, list): shape.append(len(current)) current = current[0] if current else None diff --git a/verl_speco/integration/mooncake_hidden_states_connector.py b/verl_speco/integration/mooncake_hidden_states_connector.py index 4dce80ec..9d17ffec 100644 --- a/verl_speco/integration/mooncake_hidden_states_connector.py +++ b/verl_speco/integration/mooncake_hidden_states_connector.py @@ -276,9 +276,7 @@ def build_connector_meta( self._active_requests[request.req_id] = request self._request_blocks[request.req_id] = blocks self._response_metadata[request.req_id] = { - "mooncake_key": ( - f"{self._key_prefix}_{_safe_key(request.req_id)}" - ), + "mooncake_key": (f"{self._key_prefix}_{_safe_key(request.req_id)}"), "input_ids_list": token_ids, "tensor_shapes": { "hidden_states": ( diff --git a/verl_speco/integration/oldlogprob_runtime.py b/verl_speco/integration/oldlogprob_runtime.py index 0130a3b3..2d3680c5 100644 --- a/verl_speco/integration/oldlogprob_runtime.py +++ b/verl_speco/integration/oldlogprob_runtime.py @@ -214,7 +214,9 @@ def _oldlogprob_global_step(micro_batch: Any) -> int: try: from verl.utils import tensordict_utils as tu - value = tu.get_non_tensor_data(data=micro_batch, key=OLD_LOGPROB_GLOBAL_STEP_KEY, default=0) + value = tu.get_non_tensor_data( + data=micro_batch, key=OLD_LOGPROB_GLOBAL_STEP_KEY, default=0 + ) except Exception: # noqa: BLE001 try: value = micro_batch.get(OLD_LOGPROB_GLOBAL_STEP_KEY, 0) @@ -539,7 +541,10 @@ def _put_oldlogprob_hidden_refs( # resolves "speco:" keys via TQ instead of ray.get. Falls back to # ray.put otherwise (unchanged behavior). if _speco_tq_enabled_for_oldlogprob(): - from verl_speco.integration.transferqueue_bridge import make_sample_key, put_sample + from verl_speco.integration.transferqueue_bridge import ( + make_sample_key, + put_sample, + ) tq_key = make_sample_key( _oldlogprob_global_step(micro_batch), @@ -549,7 +554,10 @@ def _put_oldlogprob_hidden_refs( put_sample( tq_key, {"hidden": hidden_chunk}, - tag={"global_step": _oldlogprob_global_step(micro_batch), "owner": int(owner)}, + tag={ + "global_step": _oldlogprob_global_step(micro_batch), + "owner": int(owner), + }, ) chunk_ref = tq_key else: diff --git a/verl_speco/integration/sglang_runtime.py b/verl_speco/integration/sglang_runtime.py index 8bb5ebfb..7bdfce1c 100644 --- a/verl_speco/integration/sglang_runtime.py +++ b/verl_speco/integration/sglang_runtime.py @@ -2030,17 +2030,24 @@ async def generate( make_sample_key, put_sample, ) + configure_transfer_queue(training_cfg) if is_transfer_queue_enabled(): - tq_key = make_sample_key(collection_global_steps, self.replica_rank, request_id) + tq_key = make_sample_key( + collection_global_steps, self.replica_rank, request_id + ) tq_payload = {"hidden_states": hidden_states.unsqueeze(0).cpu()} # P2: also offload the other large tensors so they bypass # the driver too. The drafter consumer restores them from # this same TQ payload. if target_logprobs is not None: - tq_payload["target_logprobs"] = target_logprobs.unsqueeze(0).cpu() + tq_payload["target_logprobs"] = target_logprobs.unsqueeze( + 0 + ).cpu() if torch.is_tensor(hidden_raw_target_logprobs): - tq_payload["hidden_raw_target_logprobs"] = hidden_raw_target_logprobs.unsqueeze(0).cpu() + tq_payload["hidden_raw_target_logprobs"] = ( + hidden_raw_target_logprobs.unsqueeze(0).cpu() + ) if torch.is_tensor(hidden_raw_target_logprobs_positions): tq_payload["hidden_raw_target_logprobs_positions"] = ( hidden_raw_target_logprobs_positions.unsqueeze(0).cpu() @@ -2048,7 +2055,10 @@ async def generate( put_sample( tq_key, tq_payload, - tag={"global_step": collection_global_steps, "replica_rank": self.replica_rank}, + tag={ + "global_step": collection_global_steps, + "replica_rank": self.replica_rank, + }, ) drafter_sample["hidden_states_tq_key"] = tq_key drafter_sample["hidden_states"] = None @@ -2057,7 +2067,9 @@ async def generate( if torch.is_tensor(hidden_raw_target_logprobs): drafter_sample["hidden_raw_target_logprobs"] = None if torch.is_tensor(hidden_raw_target_logprobs_positions): - drafter_sample["hidden_raw_target_logprobs_positions"] = None + drafter_sample["hidden_raw_target_logprobs_positions"] = ( + None + ) else: self._speco_log_missing_hidden_states_once( collection_global_steps=collection_global_steps, diff --git a/verl_speco/integration/task_runner.py b/verl_speco/integration/task_runner.py index 2ab07c9b..e153b170 100644 --- a/verl_speco/integration/task_runner.py +++ b/verl_speco/integration/task_runner.py @@ -314,7 +314,10 @@ def _run_with_speco_trainer(self, config): # TaskRunner owns shutdown so a failed fit cannot leak the named # controller/storage. No-op when transfer_queue.enable=false or the # package is not installed. - from verl_speco.integration.transferqueue_bridge import close_transfer_queue, init_transfer_queue + from verl_speco.integration.transferqueue_bridge import ( + close_transfer_queue, + init_transfer_queue, + ) transfer_queue_started = init_transfer_queue(config) try: diff --git a/verl_speco/models/eagle/llama_eagle.py b/verl_speco/models/eagle/llama_eagle.py index 2dcfed0b..d079a10a 100644 --- a/verl_speco/models/eagle/llama_eagle.py +++ b/verl_speco/models/eagle/llama_eagle.py @@ -1326,17 +1326,11 @@ def forward(self, x): down_proj_slices = self.down_proj.weight.split(slice, dim=1) gate_proj = torch.cat( - [ - F.linear(x, gate_proj_slices[i]) - for i in range(pretraining_tp) - ], + [F.linear(x, gate_proj_slices[i]) for i in range(pretraining_tp)], dim=-1, ) up_proj = torch.cat( - [ - F.linear(x, up_proj_slices[i]) - for i in range(pretraining_tp) - ], + [F.linear(x, up_proj_slices[i]) for i in range(pretraining_tp)], dim=-1, ) diff --git a/verl_speco/producer/input_reader.py b/verl_speco/producer/input_reader.py index c68c3245..882506dc 100644 --- a/verl_speco/producer/input_reader.py +++ b/verl_speco/producer/input_reader.py @@ -500,9 +500,7 @@ def _tokenize_chat_response_with_explicit_boundary( "Tokenizer chat template did not preserve the assistant response marker" ) prompt_text, suffix_text = rendered.split(marker, 1) - explicit_prompt_ids = _token_ids( - tokenizer(prompt_text, add_special_tokens=False) - ) + explicit_prompt_ids = _token_ids(tokenizer(prompt_text, add_special_tokens=False)) response_ids = _token_ids(tokenizer(response, add_special_tokens=False)) suffix_ids = _token_ids(tokenizer(suffix_text, add_special_tokens=False)) return explicit_prompt_ids, [*explicit_prompt_ids, *response_ids, *suffix_ids] @@ -541,10 +539,7 @@ def _build_tokenized_request( else full_ids[:feature_end] ) max_sequence_length = int(_config_value(config, "max_sequence_length", 0) or 0) - if ( - max_sequence_length > 0 - and len(request_prompt_token_ids) > max_sequence_length - ): + if max_sequence_length > 0 and len(request_prompt_token_ids) > max_sequence_length: raise ValueError( f"Producer sample {sample_id!r} requires a vLLM prefill of " f"{len(request_prompt_token_ids)} tokens after selecting its training " diff --git a/verl_speco/producer/vllm_feature_client.py b/verl_speco/producer/vllm_feature_client.py index 45c55ffb..96519100 100644 --- a/verl_speco/producer/vllm_feature_client.py +++ b/verl_speco/producer/vllm_feature_client.py @@ -19,6 +19,7 @@ import errno import importlib import inspect +import logging import os import time from dataclasses import dataclass @@ -26,6 +27,12 @@ from typing import Any, Mapping, Sequence +logger = logging.getLogger(__name__) + +_DEFAULT_MAX_RETRIES = 3 +_RETRY_BACKOFF_BASE_SECONDS = 2 + + @dataclass(frozen=True) class VllmEndpoint: base_url: str @@ -218,8 +225,29 @@ async def _request(self, request: Any, *, generate: bool) -> RawVllmFeature: state.inflight += 1 try: async with self._global_semaphore, state.semaphore: + response = await self._request_with_retries( + state, + request, + generate=generate, + ) + raw = await asyncio.to_thread(load_hidden_state_result, response) + state.requests += 1 + return raw + finally: + state.inflight = max(state.inflight - 1, 0) + + async def _request_with_retries( + self, + state: _EndpointState, + request: Any, + *, + generate: bool, + ) -> VllmResponse: + total_attempts = _DEFAULT_MAX_RETRIES + 1 + for attempt in range(1, total_attempts + 1): + try: if generate: - response = await request_generate( + return await request_generate( state.endpoint, state.client, list(request.prompt_token_ids), @@ -227,19 +255,43 @@ async def _request(self, request: Any, *, generate: bool) -> RawVllmFeature: max_tokens=int(request.max_tokens), timeout=self.request_timeout, ) - else: - response = await request_prefill( - state.endpoint, - state.client, - list(request.prompt_token_ids), - model=self.model, - timeout=self.request_timeout, + return await request_prefill( + state.endpoint, + state.client, + list(request.prompt_token_ids), + model=self.model, + timeout=self.request_timeout, + ) + except ValueError: + # Response validation failures are deterministic protocol/data + # errors, equivalent to speculators' InvalidResponseError. + raise + except Exception as exc: + if attempt >= total_attempts: + logger.error( + "vLLM request failed after %s attempts endpoint=%s " + "sample_id=%s error=%s", + total_attempts, + state.endpoint.base_url, + getattr(request, "sample_id", None), + exc, ) - raw = await asyncio.to_thread(load_hidden_state_result, response) - state.requests += 1 - return raw - finally: - state.inflight = max(state.inflight - 1, 0) + raise + backoff = _RETRY_BACKOFF_BASE_SECONDS**attempt + logger.warning( + "vLLM request aborted attempt=%s/%s endpoint=%s " + "sample_id=%s error=%s; retrying in %ss", + attempt, + total_attempts, + state.endpoint.base_url, + getattr(request, "sample_id", None), + exc, + backoff, + ) + await asyncio.sleep(backoff) + raise RuntimeError( + "unreachable: vLLM request retry loop exhausted without returning" + ) async def close(self) -> None: states, self._states = self._states, [] diff --git a/verl_speco/standalone_tq_producer.py b/verl_speco/standalone_tq_producer.py index e06fbae1..3097de64 100644 --- a/verl_speco/standalone_tq_producer.py +++ b/verl_speco/standalone_tq_producer.py @@ -334,9 +334,7 @@ async def request_worker() -> None: raw = await pool.prefill(request) stats.pending_bytes += int(raw.byte_size) try: - sample = feature_from_vllm_payload( - raw, request, feature_contract - ) + sample = feature_from_vllm_payload(raw, request, feature_contract) except HiddenStateAlignmentError as exc: stats.dropped_count += 1 stats.pending_bytes = max( diff --git a/verl_speco/standalone_tq_training_launcher.py b/verl_speco/standalone_tq_training_launcher.py index 2c1907ac..4c9b6fa1 100644 --- a/verl_speco/standalone_tq_training_launcher.py +++ b/verl_speco/standalone_tq_training_launcher.py @@ -235,9 +235,7 @@ def _parse_vllm_endpoints(env: Mapping[str, str]) -> tuple[str, ...]: if item.strip() ) else: - endpoint = str( - env.get("SPECO_VLLM_ENDPOINT", _DEFAULT_VLLM_ENDPOINT) - ).strip() + endpoint = str(env.get("SPECO_VLLM_ENDPOINT", _DEFAULT_VLLM_ENDPOINT)).strip() endpoints = (endpoint.rstrip("/"),) if endpoint else () if not endpoints: raise ValueError("At least one hidden-state vLLM endpoint is required") diff --git a/verl_speco/trainer/base_trainer.py b/verl_speco/trainer/base_trainer.py index f21363d0..de7e0f63 100644 --- a/verl_speco/trainer/base_trainer.py +++ b/verl_speco/trainer/base_trainer.py @@ -733,7 +733,9 @@ def get_training_metrics(self) -> dict[str, float]: if key.startswith(count_prefix) and key.removeprefix(count_prefix).isdigit() ) if not positions: - positions = list(range(int(self._block_drafter_config_value("block_size", 16)))) + positions = list( + range(int(self._block_drafter_config_value("block_size", 16))) + ) for pos in positions: count_key = f"{prefix}/count_per_position/{pos}" count = sums.get(count_key, 0.0) diff --git a/verl_speco/trainer/draft_training_loop.py b/verl_speco/trainer/draft_training_loop.py index b861f10d..a1b50698 100644 --- a/verl_speco/trainer/draft_training_loop.py +++ b/verl_speco/trainer/draft_training_loop.py @@ -91,7 +91,9 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: "actor_rollout_ref.rollout.drafter.training.feature_store.path is required" ) if feature_store_type == "tq" and training_mode != "offline": - raise ValueError("feature_store.type=tq requires standalone training.mode=offline") + raise ValueError( + "feature_store.type=tq requires standalone training.mode=offline" + ) if feature_store_type in replay_feature_store_types and training_mode != "offline": raise ValueError( f"feature_store.type={feature_store_type} is supported only by " @@ -281,9 +283,7 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: transfer_concurrency=int( max( math.ceil( - int( - pipeline_cfg.get("transfer_concurrency", 8) or 8 - ) + int(pipeline_cfg.get("transfer_concurrency", 8) or 8) / world_size ), 1, @@ -480,9 +480,13 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: logger.info("[standalone rank=%s] cleaning trainer resources", rank) await trainer.cleanup_training(clear_data=True) if dist.is_initialized(): - logger.info("[standalone rank=%s] entering final process-group barrier", rank) + logger.info( + "[standalone rank=%s] entering final process-group barrier", rank + ) dist.barrier() - logger.info("[standalone rank=%s] final process-group barrier complete", rank) + logger.info( + "[standalone rank=%s] final process-group barrier complete", rank + ) dist.destroy_process_group() logger.info("[standalone rank=%s] cleanup complete", rank) @@ -1137,7 +1141,9 @@ def _clear_tq_batch_across_ranks( dist.all_reduce(failed, op=dist.ReduceOp.MAX) if bool(failed.item()): if local_error is not None: - raise RuntimeError("rank 0 failed to clear a completed TQ batch") from local_error + raise RuntimeError( + "rank 0 failed to clear a completed TQ batch" + ) from local_error raise RuntimeError("rank 0 failed to clear a completed TQ batch") @@ -1153,7 +1159,9 @@ def _connect_tq_store_across_ranks(store, *, rank: int, device: torch.device) -> if connected: return if local_error is not None: - raise RuntimeError(f"TQ Consumer failed to connect on rank={rank}") from local_error + raise RuntimeError( + f"TQ Consumer failed to connect on rank={rank}" + ) from local_error raise RuntimeError( f"TQ Consumer failed to connect on another rank; rank={rank} is stopping" ) diff --git a/verl_speco/trainer/feature_store.py b/verl_speco/trainer/feature_store.py index 36c5ccb0..cf1ef0df 100644 --- a/verl_speco/trainer/feature_store.py +++ b/verl_speco/trainer/feature_store.py @@ -843,9 +843,7 @@ def _conversation_to_input_ids_and_loss_mask( raise ValueError( "jsonl_token_replay conversations produced no assistant tokens" ) - loss_mask = [0.0] * len(prompt_ids) + [1.0] * ( - len(full_ids) - len(prompt_ids) - ) + loss_mask = [0.0] * len(prompt_ids) + [1.0] * (len(full_ids) - len(prompt_ids)) return ( torch.tensor(full_ids, dtype=torch.long).reshape(-1), torch.tensor(loss_mask, dtype=torch.float32).reshape(-1), @@ -969,16 +967,16 @@ def read(self, key: str) -> DraftFeatureSample: raise ValueError( f"Invalid vllm_safetensors key {key!r}; expected sample index 0" ) - entries = { - str(entry.get("path")): entry for entry in self._load_manifest() - } + entries = {str(entry.get("path")): entry for entry in self._load_manifest()} entry = entries.get(file_name) if entry is None: raise KeyError(f"Missing vllm_safetensors manifest entry for {file_name}") tensors = load_file(str(self.path / file_name), device="cpu") manifest_sample = dict(entry.get("sample") or {}) payload: dict[str, Any] = { - "schema_version": int(manifest_sample.get("schema_version", SCHEMA_VERSION)), + "schema_version": int( + manifest_sample.get("schema_version", SCHEMA_VERSION) + ), "algorithm": manifest_sample.get("algorithm", "EAGLE3"), "input_ids": tensors["input_ids"], "loss_mask": tensors["loss_mask"], @@ -1026,11 +1024,11 @@ def _sample_to_safetensors( "position_ids", ): value = payload.get(optional_key) - if torch.is_tensor(value): + if value is not None and torch.is_tensor(value): tensors[optional_key] = value.contiguous() metadata = payload.get("metadata") or {} hidden_positions = metadata.get("hidden_positions") - if torch.is_tensor(hidden_positions): + if hidden_positions is not None and torch.is_tensor(hidden_positions): tensors["metadata.hidden_positions"] = hidden_positions.long().contiguous() return tensors @@ -1057,7 +1055,9 @@ def build_feature_store_from_config( ) if store_type == "tq": if not read_only: - raise ValueError("feature_store.type=tq is a read-only Consumer data source") + raise ValueError( + "feature_store.type=tq is a read-only Consumer data source" + ) from verl_speco.trainer.tq_feature_store import TQFeatureStore tq_cfg = transfer_queue_cfg @@ -1089,8 +1089,7 @@ def build_feature_store_from_config( tokenizer_path=feature_store_cfg.get("tokenizer_path"), trust_remote_code=bool(feature_store_cfg.get("trust_remote_code", False)), train_on=str( - feature_store_cfg.get("train_on", "last_assistant") - or "last_assistant" + feature_store_cfg.get("train_on", "last_assistant") or "last_assistant" ), ) else: @@ -1198,9 +1197,7 @@ def _load_jsonl_offset(path: Path, offset: int) -> dict[str, Any]: raise ValueError(f"JSONL offset {offset} in {path} points to an empty line") payload = json.loads(line) if not isinstance(payload, dict): - raise TypeError( - f"JSONL offset {offset} in {path} must contain a JSON object" - ) + raise TypeError(f"JSONL offset {offset} in {path} must contain a JSON object") return payload @@ -1238,9 +1235,7 @@ def _normalize_conversation_role(value: Any) -> str: def _conversation_item_to_message(item: Any) -> dict[str, str]: if not isinstance(item, dict): - raise TypeError( - "jsonl_token_replay conversations entries must be JSON objects" - ) + raise TypeError("jsonl_token_replay conversations entries must be JSON objects") role = _normalize_conversation_role(item.get("role", item.get("from"))) content = item.get("content", item.get("value", "")) if role not in {"system", "user", "assistant"}: diff --git a/verl_speco/trainer/mooncake_transfer.py b/verl_speco/trainer/mooncake_transfer.py index 8f494df1..0f4a6c51 100644 --- a/verl_speco/trainer/mooncake_transfer.py +++ b/verl_speco/trainer/mooncake_transfer.py @@ -102,9 +102,7 @@ def value(name: str, default: Any) -> Any: os.getenv("MOONCAKE_LOCAL_BUFFER_SIZE", "1GB"), ) ), - protocol=str( - value("protocol", os.getenv("MOONCAKE_PROTOCOL", "tcp")) - ), + protocol=str(value("protocol", os.getenv("MOONCAKE_PROTOCOL", "tcp"))), device_name=str( value("device_name", os.getenv("MOONCAKE_DEVICE_NAME", "")) ), diff --git a/verl_speco/trainer/speco_ray_trainer.py b/verl_speco/trainer/speco_ray_trainer.py index bf298385..e60682a7 100644 --- a/verl_speco/trainer/speco_ray_trainer.py +++ b/verl_speco/trainer/speco_ray_trainer.py @@ -2228,7 +2228,9 @@ def compute_old_log_prob_without_collection(): tu.assign_non_tensor_data(batch_td, OLD_LOGPROB_HIDDEN_OBJECT_REF_KEY, True) # Stamp the global step so the actor-worker producer can build # step-unique TransferQueue keys for old-logprob hidden chunks (P1). - tu.assign_non_tensor_data(batch_td, OLD_LOGPROB_GLOBAL_STEP_KEY, self.global_steps) + tu.assign_non_tensor_data( + batch_td, OLD_LOGPROB_GLOBAL_STEP_KEY, self.global_steps + ) self._speco_last_oldlogprob_prepare_elapsed_sec = ( time.perf_counter() - prepare_started diff --git a/verl_speco/trainer/standalone_checkpoint.py b/verl_speco/trainer/standalone_checkpoint.py index 04fbc8f3..3bd1a156 100644 --- a/verl_speco/trainer/standalone_checkpoint.py +++ b/verl_speco/trainer/standalone_checkpoint.py @@ -35,7 +35,9 @@ def _ensure_dict_child(config: dict[str, Any], key: str) -> dict[str, Any]: def _source_drafter_model_path(trainer: Any) -> str | None: - model_path = getattr(getattr(getattr(trainer, "config", None), "rollout", None), "drafter", None) + model_path = getattr( + getattr(getattr(trainer, "config", None), "rollout", None), "drafter", None + ) model_path = getattr(model_path, "model_path", None) if not model_path: return None @@ -88,13 +90,17 @@ def _target_runtime_model_type(target_config: dict[str, Any] | None) -> str | No return str(target_config["model_type"]) -def _fill_if_missing(dst: dict[str, Any], src: dict[str, Any], keys: tuple[str, ...]) -> None: +def _fill_if_missing( + dst: dict[str, Any], src: dict[str, Any], keys: tuple[str, ...] +) -> None: for key in keys: if key in src and key not in dst: dst[key] = deepcopy(src[key]) -def _copy_if_present(dst: dict[str, Any], src: dict[str, Any] | None, keys: tuple[str, ...]) -> None: +def _copy_if_present( + dst: dict[str, Any], src: dict[str, Any] | None, keys: tuple[str, ...] +) -> None: if src is None: return for key in keys: @@ -107,7 +113,9 @@ def _drop_keys(config: dict[str, Any], keys: tuple[str, ...]) -> None: config.pop(key, None) -def _target_rope_parameters(target_config: dict[str, Any] | None) -> dict[str, Any] | None: +def _target_rope_parameters( + target_config: dict[str, Any] | None, +) -> dict[str, Any] | None: if not isinstance(target_config, dict): return None rope_parameters = target_config.get("rope_parameters") @@ -240,22 +248,31 @@ def rewrite_standalone_runtime_config( try: completed_future.result() except Exception as exc: # noqa: BLE001 - logger.warning("Skip standalone runtime config rewrite because checkpoint save failed: %s", exc) + logger.warning( + "Skip standalone runtime config rewrite because checkpoint save failed: %s", + exc, + ) return None config_path = os.path.join(checkpoint_path, "config.json") if not os.path.exists(config_path): - logger.warning("Cannot rewrite standalone runtime config: missing %s", config_path) + logger.warning( + "Cannot rewrite standalone runtime config: missing %s", config_path + ) return None try: with open(config_path, "r", encoding="utf-8") as f: training_config = json.load(f) except (OSError, json.JSONDecodeError) as exc: - logger.warning("Cannot rewrite standalone runtime config %s: %s", config_path, exc) + logger.warning( + "Cannot rewrite standalone runtime config %s: %s", config_path, exc + ) return None if not isinstance(training_config, dict): - logger.warning("Cannot rewrite standalone runtime config %s: expected object", config_path) + logger.warning( + "Cannot rewrite standalone runtime config %s: expected object", config_path + ) return None training_config_path = os.path.join(checkpoint_path, "speco_training_config.json") @@ -263,7 +280,11 @@ def rewrite_standalone_runtime_config( with open(training_config_path, "w", encoding="utf-8") as f: json.dump(training_config, f, indent=2, sort_keys=True) except OSError as exc: - logger.warning("Failed to write standalone training config copy %s: %s", training_config_path, exc) + logger.warning( + "Failed to write standalone training config copy %s: %s", + training_config_path, + exc, + ) source_model_path = _source_drafter_model_path(trainer) source_runtime_config = _load_source_drafter_config(trainer) @@ -290,7 +311,13 @@ def rewrite_standalone_runtime_config( runtime_config["speco_training_model_type"] = backend_type target_runtime_keys = ("head_dim", "rope_theta", "max_position_embeddings") - common_alias_keys = ("head_dim", "rope_theta", "target_layer_ids", "mask_token_id", "num_context_layers") + common_alias_keys = ( + "head_dim", + "rope_theta", + "target_layer_ids", + "mask_token_id", + "num_context_layers", + ) _fill_if_missing(runtime_config, training_config, common_alias_keys) _copy_if_present(runtime_config, target_runtime_config, target_runtime_keys) @@ -322,7 +349,9 @@ def rewrite_standalone_runtime_config( "mask_token_id", ), ) - _copy_if_present(dspark_config, target_runtime_config, ("head_dim", "rope_theta")) + _copy_if_present( + dspark_config, target_runtime_config, ("head_dim", "rope_theta") + ) else: dspark_config = {} @@ -340,8 +369,15 @@ def rewrite_standalone_runtime_config( training_config.get("projector_type", "domino") or "domino" ) - source_has_rope_theta = isinstance(source_runtime_config, dict) and "rope_theta" in source_runtime_config - if backend_type == "dspark" and str(target_model_type or "").lower().startswith("qwen3") and not source_has_rope_theta: + source_has_rope_theta = ( + isinstance(source_runtime_config, dict) + and "rope_theta" in source_runtime_config + ) + if ( + backend_type == "dspark" + and str(target_model_type or "").lower().startswith("qwen3") + and not source_has_rope_theta + ): _drop_keys(runtime_config, ("rope_theta",)) _drop_keys(dflash_config, ("rope_theta",)) _drop_keys(dspark_config, ("rope_theta",)) @@ -355,18 +391,28 @@ def rewrite_standalone_runtime_config( or dflash_config.get("target_layer_ids") or dspark_config.get("target_layer_ids") ) - if target_layer_ids is not None and "eagle_aux_hidden_state_layer_ids" not in runtime_config: + if ( + target_layer_ids is not None + and "eagle_aux_hidden_state_layer_ids" not in runtime_config + ): try: - runtime_config["eagle_aux_hidden_state_layer_ids"] = [int(layer_id) + 1 for layer_id in target_layer_ids] + runtime_config["eagle_aux_hidden_state_layer_ids"] = [ + int(layer_id) + 1 for layer_id in target_layer_ids + ] except (TypeError, ValueError): - logger.warning("Invalid target_layer_ids in standalone exported config: %r", target_layer_ids) + logger.warning( + "Invalid target_layer_ids in standalone exported config: %r", + target_layer_ids, + ) try: with open(config_path, "w", encoding="utf-8") as f: json.dump(runtime_config, f, indent=2, sort_keys=True) f.write("\n") except OSError as exc: - logger.warning("Failed to write standalone runtime config %s: %s", config_path, exc) + logger.warning( + "Failed to write standalone runtime config %s: %s", config_path, exc + ) return None return source_model_path diff --git a/verl_speco/trainer/target_feature_pipeline.py b/verl_speco/trainer/target_feature_pipeline.py index 916616f5..fbf9f6c0 100644 --- a/verl_speco/trainer/target_feature_pipeline.py +++ b/verl_speco/trainer/target_feature_pipeline.py @@ -109,7 +109,9 @@ def submit_next() -> bool: ] else: futures = [ - self._request_executor.submit(self.replayer.materialize, [sample]) + self._request_executor.submit( + self.replayer.materialize, [sample] + ) for sample in samples ] pending.append((started, futures)) diff --git a/verl_speco/trainer/target_feature_replay.py b/verl_speco/trainer/target_feature_replay.py index 3f020d5b..6ed068bb 100644 --- a/verl_speco/trainer/target_feature_replay.py +++ b/verl_speco/trainer/target_feature_replay.py @@ -388,6 +388,10 @@ def feature_from_vllm_payload( raise TypeError("vLLM hidden-states payload must be a mapping") token_ids = values.get("token_ids") hidden = values.get("hidden_states") + if token_ids is None or hidden is None: + raise ValueError( + "vLLM hidden-states payload must contain token_ids and hidden_states" + ) if not torch.is_tensor(token_ids) or not torch.is_tensor(hidden): raise ValueError( "vLLM hidden-states payload must contain token_ids and hidden_states" @@ -1306,6 +1310,8 @@ def _request_vllm_response(self, prompt_ids: list[int]) -> Any: attempted_endpoints: set[int] = set() for attempt in range(self.vllm_max_retries + 1): state = self._acquire_vllm_endpoint(attempted_endpoints) + if state.client is None: + raise RuntimeError(f"vLLM endpoint {state.url} has no client") attempt_started = time.perf_counter() try: response = state.client.completions.create( @@ -1355,6 +1361,8 @@ def _request_vllm_hidden_states(self, prompt_ids: list[int]) -> dict[str, Any]: attempted_endpoints: set[int] = set() for attempt in range(self.vllm_max_retries + 1): state = self._acquire_vllm_endpoint(attempted_endpoints) + if state.client is None: + raise RuntimeError(f"vLLM endpoint {state.url} has no client") try: attempt_started = time.perf_counter() if log_request: @@ -1390,7 +1398,7 @@ def _request_vllm_hidden_states(self, prompt_ids: list[int]) -> dict[str, Any]: hidden_states = payload.get("hidden_states") hidden_shape = ( tuple(hidden_states.shape) - if torch.is_tensor(hidden_states) + if hidden_states is not None and torch.is_tensor(hidden_states) else None ) logger.info( diff --git a/verl_speco/trainer/tq_feature_store.py b/verl_speco/trainer/tq_feature_store.py index 4447a173..7e25af9a 100644 --- a/verl_speco/trainer/tq_feature_store.py +++ b/verl_speco/trainer/tq_feature_store.py @@ -189,7 +189,9 @@ def close(self) -> None: def _require_connected(self) -> None: if not self._connected: - raise RuntimeError("TQFeatureStore.connect() must be called before data access") + raise RuntimeError( + "TQFeatureStore.connect() must be called before data access" + ) def _plain_dict(value: Any) -> dict[str, Any]: @@ -206,10 +208,7 @@ def _plain_dict(value: Any) -> dict[str, Any]: except ImportError: # pragma: no cover - the project depends on OmegaConf pass if isinstance(value, Mapping): - return { - str(key): _plain_value(item) - for key, item in value.items() - } + return {str(key): _plain_value(item) for key, item in value.items()} raise TypeError(f"Expected a mapping configuration, got {type(value)!r}") diff --git a/verl_speco/trainer/tq_sample_source.py b/verl_speco/trainer/tq_sample_source.py index a676c7b5..3ec0ce50 100644 --- a/verl_speco/trainer/tq_sample_source.py +++ b/verl_speco/trainer/tq_sample_source.py @@ -151,10 +151,15 @@ def __iter__(self) -> Iterator[TQLocalBatch]: return if command.get("kind") != "batch": raise RuntimeError(f"Unsupported TQ loader command: {command!r}") - assignments = command.get("assignments") - if not isinstance(assignments, list) or len(assignments) != self.world_size: + wire_assignments = command.get("assignments") + if ( + not isinstance(wire_assignments, list) + or len(wire_assignments) != self.world_size + ): raise RuntimeError("TQ loader received malformed rank assignments") - local_entries = [_entry_from_wire(item) for item in assignments[self.rank]] + local_entries = [ + _entry_from_wire(item) for item in wire_assignments[self.rank] + ] samples = self.store.get_many(local_entries) global_keys = ( [str(key) for key in command.get("global_keys", [])] @@ -171,7 +176,9 @@ def clear_completed_batch(self, global_keys: Sequence[str] | None) -> None: if self.rank != 0: return if not global_keys: - raise ValueError("rank 0 requires global_keys to clear a completed TQ batch") + raise ValueError( + "rank 0 requires global_keys to clear a completed TQ batch" + ) self.store.clear_many(global_keys) def _broadcast_command(self, command: dict[str, Any] | None) -> dict[str, Any]: diff --git a/verl_speco/transport/drafter_sample_protocol.py b/verl_speco/transport/drafter_sample_protocol.py index 5ffa1ac4..2bb5d355 100644 --- a/verl_speco/transport/drafter_sample_protocol.py +++ b/verl_speco/transport/drafter_sample_protocol.py @@ -72,7 +72,9 @@ def from_dict(cls, payload: Mapping[str, Any]) -> "SampleMetadata": sequence_no=int(payload["sequence_no"]), ) except KeyError as exc: - raise ValueError(f"ready tag missing required field {exc.args[0]!r}") from exc + raise ValueError( + f"ready tag missing required field {exc.args[0]!r}" + ) from exc meta.validate() return meta @@ -125,9 +127,9 @@ def is_ready_sample_tag( run_id: str, schema_version: int = PROTOCOL_SCHEMA_VERSION, ) -> bool: - return parse_ready_tag( - tag, run_id=run_id, schema_version=schema_version - ) is not None + return ( + parse_ready_tag(tag, run_id=run_id, schema_version=schema_version) is not None + ) def encode_sample( @@ -178,7 +180,9 @@ def encode_sample( manifest["hidden_states_kind"] = "list" manifest["hidden_states_fields"] = hidden_fields else: - raise TypeError("DraftFeatureSample.hidden_states must be a tensor or tensor list") + raise TypeError( + "DraftFeatureSample.hidden_states must be a tensor or tensor list" + ) manifest["present_fields"].append("hidden_states") manifest["metadata"] = _encode_metadata_tree( payload.get("metadata", {}), fields, path="metadata" @@ -207,7 +211,9 @@ def decode_sample( raise ValueError(f"TQ sample {key!r} has an invalid or unexpected ready tag") expected_key = make_sample_key(meta) if key != expected_key: - raise ValueError(f"TQ sample key mismatch: got {key!r}, expected {expected_key!r}") + raise ValueError( + f"TQ sample key mismatch: got {key!r}, expected {expected_key!r}" + ) if _MANIFEST_FIELD not in fields: raise ValueError(f"TQ sample {key!r} missing field {_MANIFEST_FIELD!r}") manifest = _tensor_to_json(fields[_MANIFEST_FIELD], name=_MANIFEST_FIELD) @@ -269,7 +275,9 @@ def _encode_metadata_tree( encoded: dict[str, Any] = {} for key, item in value.items(): if not isinstance(key, str): - raise TypeError(f"DraftFeatureSample metadata key at {path} must be str") + raise TypeError( + f"DraftFeatureSample metadata key at {path} must be str" + ) encoded[key] = _encode_metadata_tree(item, fields, path=f"{path}.{key}") return {"__tq_mapping__": encoded} if isinstance(value, (list, tuple)): @@ -331,7 +339,9 @@ def _decode_hidden_states( def _json_to_tensor(payload: Mapping[str, Any]) -> torch.Tensor: - raw = json.dumps(dict(payload), sort_keys=True, separators=(",", ":")).encode("utf-8") + raw = json.dumps(dict(payload), sort_keys=True, separators=(",", ":")).encode( + "utf-8" + ) return torch.tensor(list(raw), dtype=torch.uint8) @@ -350,7 +360,7 @@ def _tensor_to_json(value: Any, *, name: str) -> dict[str, Any]: def _require_tensor(fields: Mapping[str, Any], name: str) -> torch.Tensor: value = fields.get(name) - if not torch.is_tensor(value): + if value is None or not torch.is_tensor(value): raise TypeError(f"TQ field {name!r} must be a torch.Tensor") return value.detach().cpu().contiguous() diff --git a/verl_speco/vllm_hidden_states_generate.py b/verl_speco/vllm_hidden_states_generate.py index d7e78e10..16ab4375 100644 --- a/verl_speco/vllm_hidden_states_generate.py +++ b/verl_speco/vllm_hidden_states_generate.py @@ -75,9 +75,7 @@ def generate_vllm_safetensors_features(config) -> dict[str, Any]: replay_config = deepcopy(config) with open_dict(replay_config): - target_feature_replay = ( - replay_config.actor_rollout_ref.rollout.drafter.training.target_feature_replay - ) + target_feature_replay = replay_config.actor_rollout_ref.rollout.drafter.training.target_feature_replay target_feature_replay.backend = "vllm_file" input_store = build_feature_store_from_config(input_cfg, read_only=True) diff --git a/verl_speco/workers/speco_worker.py b/verl_speco/workers/speco_worker.py index d30ec7ed..ddfb6716 100644 --- a/verl_speco/workers/speco_worker.py +++ b/verl_speco/workers/speco_worker.py @@ -425,7 +425,9 @@ def __init__( # collect_rollout_features can branch without re-reading config. from verl_speco.integration.transferqueue_bridge import configure_transfer_queue - self._speco_tq_enabled = configure_transfer_queue(self.config.rollout.drafter.training) + self._speco_tq_enabled = configure_transfer_queue( + self.config.rollout.drafter.training + ) def _ensure_process_group_initialized(self): if not dist.is_initialized(): @@ -937,7 +939,7 @@ def collect_rollout_features(self, samples: list[dict]): # Per-step cross-sample cache for chunk fetches. Reset every collect # call so keys (which carry global_step) never stale and the cache # cannot grow unbounded. Shared across all samples in this step. - self._tq_chunk_cache = {} + self._tq_chunk_cache: dict[str, Any] = {} for sample in samples: if not sample: continue From 6f6291811463a2e7503b5fcd8243d170fbbfa5a7 Mon Sep 17 00:00:00 2001 From: xxyyrr598 <1615681145@qq.com> Date: Thu, 27 Aug 2026 15:18:03 +0800 Subject: [PATCH 43/50] chore: stop tracking dev-only docs to align with upstream --- ...sync_vllm_mooncake_dspark_training_plan.md | 2902 ----------------- docs/standalone_tq_consumer_implementation.md | 713 ---- ...standalone_tq_foundation_implementation.md | 1030 ------ docs/standalone_tq_producer.md | 212 -- ...standalone_vllm_tq_dspark_training_plan.md | 1130 ------- docs/transferqueue_integration_plan.md | 145 - 6 files changed, 6132 deletions(-) delete mode 100644 docs/async_vllm_mooncake_dspark_training_plan.md delete mode 100644 docs/standalone_tq_consumer_implementation.md delete mode 100644 docs/standalone_tq_foundation_implementation.md delete mode 100644 docs/standalone_tq_producer.md delete mode 100644 docs/standalone_vllm_tq_dspark_training_plan.md delete mode 100644 docs/transferqueue_integration_plan.md diff --git a/docs/async_vllm_mooncake_dspark_training_plan.md b/docs/async_vllm_mooncake_dspark_training_plan.md deleted file mode 100644 index 525da0fb..00000000 --- a/docs/async_vllm_mooncake_dspark_training_plan.md +++ /dev/null @@ -1,2902 +0,0 @@ -# verl-SpeCo 异步 vLLM → TransferQueue(Mooncake 后端)→ DSpark 流式训练方案 - -Last updated: 08/21/2026 - -## 1. 文档范围 - -本文只讨论当前 `verl-SpeCo-ls` 项目的独立草稿模型训练入口: - -```text -examples/run_qwen3-8b_drafter_separate_training.sh - → python -m verl_speco.standalone_tq_training_launcher - ├─→ TransferQueue owner - ├─→ vLLM hidden-state producer - └─→ TransferQueue consumer - → python -m verl_speco.draft_train_launcher - → torch.distributed.run - → python -m verl_speco.draft_train - → run_standalone_draft_training() -``` - -目标是训练 `mode=train/offline` 的草稿模型,不是 RL 与草稿模型一起训练,也不是先生成全量 hidden states 再长期保存到磁盘。 - -输入既可以是已有 response 的 replay 文件,也可以是 verl prompt-only Parquet。新流水线需要: - -1. Producer 读取 prompt;缺少 response 时由 target vLLM 生成; -2. Producer 并行请求一个或多个 vLLM endpoint 做 prefill; -3. hidden states 写入 TransferQueue(简称 TQ),TQ 的数据后端使用 Mooncake; -4. DSpark 的各训练 rank 从 TQ/Mooncake 并行读取自己的 local batch; -5. 所有 rank 完成同一个 optimizer step 后,清理这一批 TQ 数据; -6. 不使用自研 Stream Coordinator;第一版复用 PR #48 已验证的 TQ KV 传输模式,由 TQ 保存 payload 和 tags,rank 0 通过 `kv_list` 发现 ready key 并组成 global batch。 - -本文不照搬 AngelSpec 的进程组织。实现依据是旁边只读参考仓库 `verl-SpeCo` 的 PR #48 分支,尤其是: - -```text -verl_speco/integration/transferqueue_bridge.py -verl_speco/integration/sglang_runtime.py -verl_speco/integration/oldlogprob_runtime.py -verl_speco/workers/speco_worker.py -verl_speco/integration/task_runner.py -``` - -参考仓库只用于理解和复用设计,不修改其代码。真正实现仍放在本项目 `verl-SpeCo-ls`,服务于 standalone drafter training。 - -建议按下面顺序阅读: - -1. 第一部分先介绍 verl/SpeCo 中的进程、`TokenOutput`、`DataProto`、WorkerGroup、replica 和 owner rank; -2. 再完整解释 PR #48 的 SGLang TQ 路径; -3. 再解释 PR #48 的 old-logprob TQ 路径; -4. 再解释 hidden 恢复后如何进入 online buffer 和 optimizer step; -5. 第二部分才讨论如何把相同 TQ 传输能力改造成独立 Producer/DSpark Consumer。 - -## 第一部分:PR #48 原始 TQ 流程 - -这一部分只解释只读参考仓库 `../verl-SpeCo` 的 PR #48。这里的 Producer、Consumer、Ray driver、数据结构和生命周期全部是 PR #48 当前代码的行为,不是本文为 standalone 设计的行为。 - -### 1A. 阅读 PR #48 前必须知道的项目对象 - -#### SGLang server - -SGLang server 是 rollout 推理进程。它接收 prompt,执行 target model 推理并生成 response。SpeCo 的 SGLang patch 还会在推理过程中收集指定层 hidden states。 - -它不是 drafter trainer,也不执行 optimizer step。在 PR #48 中它是 hidden-state Producer。 - -#### TokenOutput - -`TokenOutput` 是一次 rollout request 的返回对象。核心字段可以概括为: - -```python -TokenOutput( - token_ids=list[int], - log_probs=..., - routed_experts=..., - extra_fields={ - "global_steps": int, - "drafter_sample": dict | None, - }, -) -``` - -`token_ids` 是生成结果;`extra_fields` 是 SpeCo 添加的旁路字段。`drafter_sample` 不影响正常 response 返回,它用于把草稿训练所需信息从 rollout 侧带回训练控制层。 - -#### DataProto 和 non_tensor_batch - -verl 将一批 rollout 结果整理成 `DataProto`。它通常分为: - -```python -DataProto( - batch=TensorDict(...), - non_tensor_batch={...}, - meta_info={...}, -) -``` - -- `batch`:规则的批量 tensor,例如 prompts、responses、attention mask; -- `non_tensor_batch`:不能直接组成规则 dense batch 的 Python 对象或 object array; -- `meta_info`:批次级配置和指标。 - -每个 request 的 `TokenOutput.extra_fields["drafter_sample"]` 最终会被 rollout/agent-loop 聚合到: - -```python -gen_batch_output.non_tensor_batch["drafter_sample"] -``` - -所以 SGLang 侧写的是单 request `TokenOutput`,driver 侧拿到的是批量 `DataProto`。 - -#### RayPPOTrainer driver - -`SpecoRayPPOTrainer` 所在进程是控制进程,简称 driver。它负责按训练 step 依次调用 rollout、old-logprob、actor update,以及向各 WorkerGroup 发 RPC。 - -driver 不执行 drafter 模型 forward。PR #48 之前,它会接触包含 hidden tensor 的 `drafter_sample`;PR #48 之后,它只处理中小 tensor、标量 metadata 和 TQ key。 - -#### WorkerGroup - -WorkerGroup 是 verl 对一组 Ray worker 的调用封装。代码: - -```python -self.drafter_wg.collect_rollout_features(buckets) -``` - -不是普通本地函数调用,而是根据注册的 dispatch rule,将参数分发到多张卡上的 `SpecoWorker.collect_rollout_features()`。 - -#### Rollout replica - -rollout replica 是一组共同承载一个 rollout 模型副本的进程/GPU。一个任务可能存在多个 data-parallel rollout replica。`replica_rank` 标识当前 sample 是哪个 rollout replica 生成的: - -```text -replica_rank = 0, 1, 2, ... -``` - -#### Drafter training replica、DP rank 和 SP rank - -drafter 训练也可能按 data parallel 和 sequence parallel 组织: - -```text -drafter replica / DP rank 0 - ├─ SP rank 0 - └─ SP rank 1 - -drafter replica / DP rank 1 - ├─ SP rank 0 - └─ SP rank 1 -``` - -同一个 drafter DP replica 内的 SP ranks 共同执行一个模型副本的训练。`replica_rank` 用来把 rollout replica 产生的数据路由到对应 drafter DP replica。 - -#### Owner rank - -`collect_rollout_features` 注册了: - -```python -@register( - dispatch_mode=make_nd_compute_dispatch_fn( - mesh_name="drafter_owner_route" - ) -) -``` - -每个 drafter DP replica 的 `SP rank 0` 被标记为 collect leader/owner。driver 传入的是按 replica 分好的 bucket,dispatch 层负责把 bucket 发到对应训练组。一个 replica 内可能有多个 rank 需要共同训练,但只有指定 leader 负责汇总 RPC 返回。 - -这里的 owner route 是 PR #48 为什么不能简单让任意 worker 随机拿 sample 的原因:sample 必须进入与当前 drafter device mesh 一致的训练组。 - -## 2. PR #48 改造前的 online 特征流程 - -PR #48 改造的是 SpeCo 的 online drafter feature transport。hidden states 有两条主要来源。 - -### 2.1 SGLang rollout hidden 路径 - -改造前: - -```text -SGLang server - → 生成 drafter_sample,其中直接包含 hidden_states CPU tensor - → TokenOutput.extra_fields - → RayPPOTrainer driver 收集 drafter_sample - → driver 按 drafter replica/owner 分桶 - → Ray dispatch / object store - → SpecoWorker.collect_rollout_features(samples) - → _store_rollout_sample() - → online drafter buffer/train -``` - -此时 `drafter_sample` 类似: - -```python -{ - "input_ids": int64[1, total_len], - "prompts": int64[1, prompt_len], - "responses": int64[1, response_len], - "hidden_states": bf16[1, hidden_rows, hidden_dim], - "hidden_positions": int64[1, hidden_rows], - "target_logprobs": tensor | None, - "global_step": 42, - "replica_rank": 1, -} -``` - -问题是整个字典经过 driver,而 `hidden_states` 是其中最大的字段。driver 的 host memory 和 Ray object store 都会承载这些 tensor。 - -### 2.2 old-logprob hook hidden 路径 - -另一条路径在 actor old-logprob forward 中捕获 hidden states。改造前,大 chunk 通过: - -```python -chunk_ref = ray.put(hidden_chunk) -``` - -sample 不直接带 tensor,而是带: - -```python -{ - "hidden_states_ref_chunks": [ - { - "ref": ray_object_ref, - "start": 0, - "length": 512, - }, - ], -} -``` - -drafter worker 再 `ray.get(ref)`,根据 `start/length` 切出每条 sample 所需行。 - -### 2.3 PR #48 要改变的边界 - -PR #48 没有改变: - -- rollout 什么时候产生 sample; -- driver 如何触发 drafter worker; -- drafter worker 如何调用 `_store_rollout_sample()`; -- drafter model 的训练逻辑; -- drafter 权重发布。 - -它只改变大 tensor 的跨进程介质: - -```text -改造前:Producer → Ray driver/object store → Consumer -改造后:Producer → TQ storage → Consumer - key 仍走原 Ray 控制路径 -``` - -## 3. PR #48 改造后的完整 TQ 流程 - -### 3.0 总览 - -PR #48 的目标不是让 TQ 自己产生训练 batch,而是把原来经过 Ray driver/Ray object store 的大 hidden tensor 攁到 TQ。原来的控制路径继续存在,只是控制路径上从“大 tensor”变成“小 key”。 - -整体结构是: - -```text - 原 Ray 控制路径 - drafter_sample / chunk ref -Producer ───────────────────── key ───────────────────▶ Consumer - │ │ - │ kv_put(large tensor) │ kv_batch_get(key) - ▼ ▼ -TransferQueue storage ─────────────────────────────────────┘ -``` - -因此 PR #48 同时保留两条通道: - -```text -控制通道:Producer → Ray driver → drafter worker -数据通道:Producer → TQ storage → drafter worker -``` - -控制通道负责告诉 Consumer “本次应该处理哪个样本、对应哪个 TQ key”;数据通道负责传输 hidden states 等大 tensor。 - -#### 3.0.1 配置放在哪里 - -PR #48 在 drafter training 配置下增加: - -```yaml -actor_rollout_ref: - rollout: - drafter: - training: - transfer_queue: - enable: false - backend: - storage_backend: SimpleStorage - SimpleStorage: - total_storage_size: 100000 - num_data_storage_units: 8 -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/config/speco_base.yaml -``` - -这里的 `transfer_queue` 是 SpeCo drafter 自己的配置,不是 standalone `feature_store.type`,也不是简单把上游 verl 的 `transfer_queue.enable` 打开。 - -#### 3.0.2 TaskRunner 创建整套 TQ - -RL 任务启动时,`SpecoTaskRunner.run()` 在创建 workers 之前调用: - -```python -from verl_speco.integration.transferqueue_bridge import ( - close_transfer_queue, - init_transfer_queue, -) - -transfer_queue_started = init_transfer_queue(config) -try: - trainer.init_workers() - trainer.fit() -finally: - if transfer_queue_started: - close_transfer_queue() -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/integration/task_runner.py:319 -``` - -`init_transfer_queue(config)` 内部读取: - -```python -config.actor_rollout_ref.rollout.drafter.training.transfer_queue -``` - -然后执行: - -```python -tq.init(_to_plain_dict(tq_cfg)) -``` - -并记录: - -```python -_state["initialized"] = True -_state["owner"] = True -``` - -这里的 owner 是“创建/拥有 TQ 生命周期的进程”。只有 owner 在任务结束时执行 `tq.close()`。 - -关键顺序是: - -```text -SpecoTaskRunner -→ tq.init(完整配置) -→ trainer.init_workers() -→ Ray workers 启动 -``` - -也就是说,PR #48 假设 TQ Controller/storage 已经由 TaskRunner 在 Ray 集群环境中建立,后启动的 worker 只需要连接。 - -#### 3.0.3 每个 Producer/Consumer 进程怎么连接 TQ - -bridge 中的 `_ensure_initialized()` 是进程级懒初始化: - -```python -def _ensure_initialized(): - if _state["initialized"]: - return - - with _state_lock: - if _state["initialized"]: - return - - tq.init() - _state["initialized"] = True -``` - -注意这里是: - -```python -tq.init() -``` - -不是: - -```python -tq.init(config) -``` - -无参初始化的含义是连接 TaskRunner 已创建的同一套 TQ。它依赖 PR #48 所处的 Ray 运行环境完成服务发现。 - -因此 PR #48 不是“每个 worker 各创建一套 TQ”,而是: - -```text -TaskRunner:tq.init(config),创建一次 -SGLang producer:tq.init(),连接 -actor producer:tq.init(),连接 -drafter consumer:tq.init(),连接 -``` - -#### 3.0.4 SGLang Producer 怎么写 hidden states - -SGLang 已经完成 rollout,并组装出 `drafter_sample` 后,PR #48 执行: - -```python -configure_transfer_queue(training_cfg) - -if is_transfer_queue_enabled(): - tq_key = make_sample_key( - collection_global_steps, - self.replica_rank, - request_id, - ) - - tq_payload = { - "hidden_states": hidden_states.unsqueeze(0).cpu(), - } - - if target_logprobs is not None: - tq_payload["target_logprobs"] = ( - target_logprobs.unsqueeze(0).cpu() - ) - - put_sample( - tq_key, - tq_payload, - tag={ - "global_step": collection_global_steps, - "replica_rank": self.replica_rank, - }, - ) - - drafter_sample["hidden_states_tq_key"] = tq_key - drafter_sample["hidden_states"] = None -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/integration/sglang_runtime.py:2020 -``` - -这里发生了两条不同的数据流: - -```text -大 tensor:SGLang → TQ -小字典/key:SGLang → 原 Ray/driver 路径 → drafter worker -``` - -写进 TQ 后将: - -```python -drafter_sample["hidden_states"] = None -``` - -是为了避免 hidden states 继续经过 driver/Ray object store。driver 仍收到 sample,但大 tensor 已替换成: - -```python -drafter_sample["hidden_states_tq_key"] -``` - -#### 3.0.5 `put_sample()` 实际怎么写 - -bridge 中: - -```python -def put_sample(key, tensor_dict, *, tag=None): - payload = { - k: v - for k, v in tensor_dict.items() - if torch.is_tensor(v) - } - - _ensure_initialized() - - tq.kv_put( - key=key, - partition_id="speco_drafter_features", - fields=payload, - tag=tag or {}, - ) -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/integration/transferqueue_bridge.py:206 -``` - -这里可以明确看到: - -- PR #48 使用 TQ 高层 KV API; -- 一个 key 对应一个 sample; -- `fields` 是 tensor 字典; -- `tag` 是小 metadata; -- partition 当前写死为 `speco_drafter_features`; -- 写入前 tensor 已 `.cpu()`; -- 写入失败直接抛异常,不静默回退。 - -key 的生成代码是: - -```python -def make_sample_key(global_step, replica_rank, request_id): - return f"speco:{global_step}:{replica_rank}:{request_id}" -``` - -这个 key 对 RL rollout 是合理的,因为它用 step、rollout replica 和 request ID 标识一次在线采样。 - -#### 3.0.5.1 PR #48 中“数据”和“元数据”分别长什么样 - -PR #48 一条 SGLang 样本在写 TQ 之前,`drafter_sample` 同时包含大 tensor 和控制信息。简化后类似: - -```python -drafter_sample = { - # 普通训练输入,仍走原 sample/Ray 控制路径 - "input_ids": int64[1, prompt_len + response_len], - "prompts": int64[1, prompt_len], - "responses": int64[1, response_len], - - # 大 tensor,开启 TQ 后从这个字典移除 - "hidden_states": bf16[1, hidden_rows, hidden_dim], - "target_logprobs": fp32[1, rows, topk_or_vocab] | None, - - # hidden 与 token 对齐所需的小字段 - "hidden_positions": int64[1, hidden_rows] | None, - "hidden_position_start": int, - "hidden_position_end": int, - "hidden_window_start": int, - "hidden_window_end": int, - - # 控制信息 - "global_step": int, - "replica_rank": int, -} -``` - -执行 `put_sample()` 时,并不是把整个 `drafter_sample` 放进 TQ。PR #48 只抽取占用大的 tensor: - -```python -tq_payload = { - "hidden_states": bf16[1, hidden_rows, hidden_dim], - "target_logprobs": fp32[1, rows, ...], # 可选 - "hidden_raw_target_logprobs": ..., # 可选 - "hidden_raw_target_logprobs_positions": ..., # 可选 -} -``` - -这就是 TQ 的 data payload。它被传给: - -```python -tq.kv_put(fields=tq_payload) -``` - -另外还有 TQ tag: - -```python -tag = { - "global_step": 42, - "replica_rank": 1, -} -``` - -tag 是 TQ 侧轻量 metadata,用于描述/检索对象,不承载 hidden tensor。 - -写入完成后,仍经 Ray 传递的轻量 `drafter_sample` 变成: - -```python -drafter_sample = { - "input_ids": int64[1, total_len], - "prompts": int64[1, prompt_len], - "responses": int64[1, response_len], - - "hidden_states": None, - "target_logprobs": None, - "hidden_states_tq_key": "speco:42:1:req-007", - - "hidden_positions": int64[1, hidden_rows] | None, - "hidden_position_start": 128, - "hidden_position_end": 640, - "global_step": 42, - "replica_rank": 1, -} -``` - -因此 PR #48 实际存在三类对象: - -| 对象 | 内容 | 传输路径 | 作用 | -|---|---|---|---| -| TQ fields/payload | `hidden_states` 等大 tensor | Producer → TQ storage → Consumer | 避免 driver 搬运大 tensor | -| TQ tag | `global_step`、`replica_rank` | TQ control/index metadata | 描述该 key | -| 轻量 `drafter_sample` | tokens、位置、标量 metadata、`hidden_states_tq_key` | 原 Ray driver 路径 | 告诉 Consumer 取哪个 key,以及如何解释取回的 tensor | - -代码实现解耦的关键不是“所有内容都进 TQ”,而是: - -```python -drafter_sample["hidden_states_tq_key"] = tq_key -drafter_sample["hidden_states"] = None -``` - -第一行给 Consumer 留下寻址信息;第二行阻止大 tensor 继续沿旧路径传输。 - -#### 3.0.5.2 Consumer 如何把两部分重新合成一个训练样本 - -Consumer 最初拿到的是轻量 sample: - -```python -sample["hidden_states"] is None -sample["hidden_states_tq_key"] == "speco:42:1:req-007" -``` - -它执行: - -```python -payload = get_sample(sample["hidden_states_tq_key"]) -sample["hidden_states"] = payload["hidden_states"] -``` - -合并后: - -```python -sample = { - "input_ids": ..., - "prompts": ..., - "responses": ..., - "hidden_positions": ..., - "hidden_states": bf16[1, hidden_rows, hidden_dim], - "hidden_states_tq_key": "speco:42:1:req-007", - ... -} -``` - -后面的 `_store_rollout_sample()` 看到的结构与关闭 TQ 时基本一致,所以训练主体不需要增加 TQ 分支。TQ bridge 只改变“大 tensor 从哪里恢复”,不改变 drafter trainer 的输入语义。 - -#### 3.0.6 old-logprob Producer 怎么写 chunk - -PR #48 的另一条 Producer 路径来自 actor old-logprob hidden hook。原来是: - -```python -chunk_ref = ray.put(hidden_chunk) -``` - -开启 TQ 后改成: - -```python -tq_key = make_sample_key( - global_step, - owner, - f"chunk{len(chunk_refs)}", -) - -put_sample( - tq_key, - {"hidden": hidden_chunk}, - tag={ - "global_step": global_step, - "owner": owner, - }, -) - -chunk_ref = tq_key -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/integration/oldlogprob_runtime.py:533 -``` - -后面的 driver 不需要知道 ref 是 Ray ObjectRef 还是 TQ string key,它只把 ref 当不透明 token 继续传递。 - -#### 3.0.7 Consumer 怎么根据 key 读取 - -drafter worker 收到原来的 sample 小字典后: - -```python -tq_key = sample.get("hidden_states_tq_key") - -if tq_key is not None and self._speco_tq_enabled: - payload = get_sample(tq_key) - - for field in ( - "hidden_states", - "target_logprobs", - "hidden_raw_target_logprobs", - "hidden_raw_target_logprobs_positions", - ): - if payload.get(field) is not None: - sample[field] = payload[field] - - if sample.get("hidden_states") is None: - raise RuntimeError( - "TQ key exists but hidden_states payload is missing" - ) -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/workers/speco_worker.py:848 -``` - -恢复 tensor 后,后面的逻辑仍使用原 `sample`/`batch`,drafter trainer 不需要知道 tensor 来自 Ray 还是 TQ。 - -#### 3.0.8 `get_sample()` 实际怎么读 - -```python -def get_sample(key): - _ensure_initialized() - - result = tq.kv_batch_get( - keys=[key], - partition_id="speco_drafter_features", - ) - - value = _extract_value(result, key) - return _tensordict_to_dict(value) -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/integration/transferqueue_bridge.py:237 -``` - -`_extract_value()` 兼容三种返回形态: - -```python -if isinstance(result, dict): - return result.get(key) -if isinstance(result, (list, tuple)): - return result[0] -return result -``` - -这是因为不同 TQ 版本/后端返回包装可能不同。 - -#### 3.0.9 为什么需要 `_densify_tq_tensor()` - -PR #48 后续修复发现,TQ 把单样本 tensor 放进 TensorDict 后,`kv_batch_get` 可能返回 NestedTensor,并额外带 batch 维。旧代码要执行: - -```python -tensor[start:start + length] -``` - -但 NestedTensor 不支持在 jagged dim 直接 slice。因此加入: - -```python -def _densify_tq_tensor(tensor): - if tensor.is_nested: - parts = [ - part - for part in tensor.unbind() - if part.numel() > 0 - ] - tensor = torch.cat(parts, dim=0) - - if tensor.dim() == 3: - tensor = tensor.squeeze(0) - elif tensor.dim() == 1: - tensor = tensor.unsqueeze(0) - - return tensor.contiguous() -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/workers/speco_worker.py:72 -``` - -standalone Consumer 同样必须做这个转换,不能假定 `kv_batch_get` 返回普通 dense `[seq, hidden]`。 - -#### 3.0.10 为什么需要 per-step cache - -old-logprob 路径里,一个 owner hidden chunk 可能被约 16 个 sample 共同引用。如果每个 sample 都: - -```python -get_sample(same_tq_key) -``` - -就会重复传输同一个数百 MB chunk。PR #48 在每次 `collect_rollout_features()` 开始时创建: - -```python -self._tq_chunk_cache = {} -``` - -解析 ref 时: - -```python -cache_key = ref if isinstance(ref, str) else id(ref) - -if cache_key not in cache: - cache[cache_key] = _resolve_tq_or_ray_ref(ref) - -tensor = cache[cache_key] -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/workers/speco_worker.py:98 -../verl-SpeCo/verl_speco/workers/speco_worker.py:854 -``` - -独立训练如果每个 sample 都是独立 TQ key,主要使用 `kv_batch_get(keys=[...])` 批量读取,不一定需要跨 sample chunk cache;但 prefetch 重试或共享 packed object 时仍应保留 key cache。 - -#### 3.0.11 PR #48 什么时候删除数据 - -PR #48 没有在 `get_sample()` 后删除。其注释明确说明:同一 drafter replica 的多个 TP/SP rank 可能读取同一 key,第一次读取后立刻删除会让剩余 rank 失败。 - -当前策略是任务结束时由 owner: - -```python -tq.close() -``` - -统一结束 TQ 生命周期。也就是说,PR #48 当前没有实现精细的逐 step `kv_clear`。 - -#### 3.0.12 PR #48 的完整时序 - -```text -SpecoTaskRunner - → tq.init(config) - → 启动 Ray workers - -SGLang/actor Producer process - → configure_transfer_queue() - → 第一次 put 时 tq.init() - → kv_put(key, tensor fields, tag) - → 把 key 塞回原 sample/ref - -Ray driver - → 只中转小 sample/key - -drafter worker Consumer process - → 第一次 get 时 tq.init() - → kv_batch_get([key]) - → 解包 TensorDict/NestedTensor - → 恢复 sample["hidden_states"] - → 原 drafter collect/train 逻辑 - -任务结束 - → TaskRunner owner tq.close() -``` - -### 3.1 已经实现的可复用能力 - -PR #48 新增 `verl_speco/integration/transferqueue_bridge.py`,锁定思路是把 TQ 当成独立传输库,不修改上游 verl。它提供: - -```python -configure_transfer_queue(training_cfg) -init_transfer_queue(config) -make_sample_key(global_step, replica_rank, request_id) -put_sample(key, tensor_dict, tag=...) -get_sample(key) -close_transfer_queue() -``` - -实际写入调用是: - -```python -tq.kv_put( - key=key, - partition_id="speco_drafter_features", - fields=payload, - tag=tag, -) -``` - -实际读取调用是: - -```python -result = tq.kv_batch_get( - keys=[key], - partition_id="speco_drafter_features", -) -``` - -另外,PR #48 已经处理了多项 standalone 方案也需要的问题: - -1. TQ 返回值可能是 direct value、`{key: value}` 或 list,需要统一解包; -2. TQ 返回的 tensor 可能是 NestedTensor,必须通过 `unbind + cat` 恢复为生产端写入的 dense tensor; -3. 多个样本引用同一 hidden chunk 时,需要 per-step cache,避免重复 `kv_batch_get` 同一个大对象; -4. TQ key 存在但 payload 缺失时 fail loud,不能静默丢样本; -5. `enable=false` 时保留原传输路径。 - -这些逻辑应直接作为本项目 TQ adapter 的参考。 - -### 3.2 PR #48 的数据流 - -PR #48 优化的是 RL online 路径: - -```text -SGLang/actor worker - → kv_put(hidden states) - → 把 hidden_states_tq_key 塞进原 drafter_sample - → 原 Ray driver 继续传递小 sample/key - → drafter worker collect_rollout_features() - → kv_batch_get(key) -``` - -它没有让 consumer 自己从 TQ 发现下一批 key;key 仍沿原来的 Ray 控制路径到达 drafter worker。 - -### 3.3 PR #48 没有提供的 standalone 能力 - -PR #48 当前没有实现: - -- 从预生成 response 文件读取数据的独立 Producer; -- Producer 并行请求外部 vLLM endpoint; -- standalone DSpark trainer 主动发现 ready key; -- global batch 到各 torchrun rank 的分片; -- 每个 optimizer step 后精确 `kv_clear`; -- EOS; -- standalone 无 Ray 的 TQ bootstrap; -- MooncakeStore 的实际运行验证。 - -PR #48 当前配置是: - -```yaml -transfer_queue: - enable: false - backend: - storage_backend: SimpleStorage - SimpleStorage: - total_storage_size: 100000 - num_data_storage_units: 8 -``` - -并且 bridge 注释针对 `TransferQueue==0.1.7`。旁边当前 verl 主线已使用 `0.1.8` 文案并包含 `MooncakeStore` 配置。当前机器没有安装 `transfer_queue` 包,因此正式实现前必须锁定版本并实机验证 API 签名,不能把 0.1.7 和 0.1.8 混用。 - -### 3.4 standalone 方案对 PR #48 的扩展 - -不再自行实现 `publish_ready/claim/lease/ack` HTTP 服务。第一版在 PR #48 KV 模式上补四个操作: - -```python -tq.kv_batch_put(...) # Producer 批量写 -tq.kv_list(...) # rank 0 列出 key + tag -tq.kv_batch_get(...) # 各 rank 并行读 -tq.kv_clear(...) # optimizer step 成功后删 -``` - -第一版不依赖 TQ `Sampler/StreamingDataLoader/get_meta`,因为 PR #48 并未使用或验证这些接口。等 KV 流程稳定后再升级为 TQ StreamingDataLoader。 - -### 3.5 PR #48 与 standalone 独立训练逐项映射 - -| PR #48 online RL | standalone drafter training | -|---|---| -| `SpecoTaskRunner` 调用 `tq.init(config)` | 新增独立 TQ owner/bootstrap 进程,或在确认安全后由 Producer owner 调用 `tq.init(config)` | -| SGLang/actor worker 是 Producer | `feature_producer.py` 是独立 Producer | -| rollout 过程中已经得到 hidden states | Producer 读取预生成 response,再请求外部 vLLM prefill | -| `make_sample_key(global_step, replica_rank, request_id)` | `make_replay_sample_key(dataset,row,tokens,target fingerprint)` | -| `put_sample()` / `kv_put()` | 继续复用同一写入模式,可扩展为 `kv_batch_put()` | -| key 通过 Ray driver/sample 传给 consumer | 没有 driver;rank 0 用 `kv_list()` 主动发现 ready keys | -| drafter worker `get_sample(key)` | 每个 torchrun rank `kv_batch_get(local_keys)` | -| `collect_rollout_features()` 恢复 sample | 转换为 `DraftFeatureSample` 后调用现有 `prepare_training_batch_from_samples()` | -| 多 TP/SP rank 可能读同一 key,所以不立即删除 | data-parallel rank 读取互不重叠的 key;全 rank step 成功后统一 `kv_clear(global_keys)` | -| TaskRunner 结束时 `tq.close()` | 每 step clear;输入 drain 完成后 owner 最后 `tq.close()` | - -standalone 需要新增的控制流是: - -```text -Producer DSpark rank 0 其他 ranks - │ │ │ - │ kv_put(sample key, fields, tag) │ │ - ├────────────────────────────────────▶│ │ - │ │ kv_list READY keys │ - │ │ │ - │ │ broadcast selected_keys ──▶│ - │ │ │ - │ │ kv_batch_get(local keys) │ kv_batch_get(local keys) - │ │ │ - │ ├──── DSpark synchronized step ────┤ - │ │ │ - │ │ kv_clear(global keys) │ -``` - -这个映射中,TQ 同时承担: - -- 大 tensor 存储/传输; -- key、tag 和 partition 的轻量索引。 - -但第一版 global batch 的选择仍由单个 DSpark job 的 rank 0 完成。这样最接近 PR #48 的 KV API,避免在同一次改造中再引入未经该 PR 验证的 Sampler/StreamingDataLoader。 - -### 3.6 PR #48 的 SGLang 路径:逐函数、逐对象完整流程 - -下面从一次生成请求开始,不省略中间层。 - -#### 阶段 1:SGLang完成生成并收集 hidden states - -执行进程:SGLang rollout server。 - -输入是一次 request 对应的 prompt 和生成配置。生成结束时,代码已经持有: - -```python -prompt_tensor: int64[prompt_len] -response_tensor: int64[response_len] -hidden_states: bf16[hidden_rows, hidden_dim] -hidden_positions: int64[hidden_rows] | None -target_logprobs: tensor | None -request_id: str -collection_global_steps: int -self.replica_rank: int -``` - -这些变量的语义: - -- `prompt_tensor`:输入 prompt token IDs; -- `response_tensor`:SGLang生成的 response token IDs; -- `hidden_states`:target model 指定层在部分 token positions 上的输出; -- `hidden_positions`:每个 hidden row 对应完整 `prompt+response` 序列中的哪个 token position; -- `target_logprobs`:可选的目标概率监督; -- `request_id`:当前 rollout request 标识; -- `replica_rank`:执行该 request 的 rollout replica。 - -SGLang 先构造完整 sample: - -```python -drafter_sample = { - "input_ids": torch.cat( - [prompt_tensor, response_tensor], dim=0 - ).unsqueeze(0), - "prompts": prompt_tensor.unsqueeze(0), - "responses": response_tensor.unsqueeze(0), - "hidden_states": hidden_states.unsqueeze(0).cpu(), - "hidden_positions": hidden_positions.unsqueeze(0).cpu(), - "target_logprobs": ( - target_logprobs.unsqueeze(0).cpu() - if target_logprobs is not None - else None - ), - "global_step": collection_global_steps, - "replica_rank": self.replica_rank, - # 还有 hidden window/alignment metadata -} -``` - -前导维 `1` 表示这是一个单样本 batch。`.cpu()` 表示跨进程传输前把大 tensor 放到 CPU 内存。 - -#### 阶段 2:PR #48 将大 fields 写入 TQ - -同一个 SGLang进程执行: - -```python -tq_key = make_sample_key( - collection_global_steps, - self.replica_rank, - request_id, -) - -tq_payload = { - "hidden_states": hidden_states.unsqueeze(0).cpu(), -} - -put_sample( - tq_key, - tq_payload, - tag={ - "global_step": collection_global_steps, - "replica_rank": self.replica_rank, - }, -) -``` - -调用展开后是: - -```python -tq.init() # 当前进程第一次使用时 -tq.kv_put( - key=tq_key, - partition_id="speco_drafter_features", - fields=tq_payload, - tag=tag, -) -``` - -效果是 TQ 中增加一行: - -```text -partition = speco_drafter_features -key = speco:42:1:req-007 -fields = {hidden_states: bf16[1, H, D], ...} -tag = {global_step: 42, replica_rank: 1} -``` - -`kv_put` 返回后,SGLang侧将旧 sample 改成: - -```python -drafter_sample["hidden_states_tq_key"] = tq_key -drafter_sample["hidden_states"] = None -``` - -此后该 Python 字典不再携带 hidden tensor,只携带定位它的 key。 - -#### 阶段 3:把 drafter_sample 放进 TokenOutput.extra_fields - -SGLang返回: - -```python -TokenOutput( - token_ids=token_ids, - log_probs=log_probs, - routed_experts=routed_experts, - extra_fields={ - "global_steps": collection_global_steps, - "drafter_sample": drafter_sample, - }, -) -``` - -此时 `TokenOutput` 中有两类输出: - -- 正常 rollout 输出:`token_ids/log_probs`; -- SpeCo 训练旁路输出:`extra_fields.drafter_sample`。 - -TQ payload 不在 `TokenOutput` 中,只有 `hidden_states_tq_key` 在其中。 - -#### 阶段 4:多个 TokenOutput 聚合成 gen_batch_output - -rollout/agent-loop 层将多个 request 的输出合并为批量 `DataProto`: - -```python -gen_batch_output.non_tensor_batch["drafter_sample"] -``` - -可能是 object array: - -```python -array([ - {"hidden_states_tq_key": "speco:42:0:req-A", ...}, - {"hidden_states_tq_key": "speco:42:1:req-B", ...}, -], dtype=object) -``` - -之所以进入 `non_tensor_batch`,是因为每条 sample 的 hidden window、Python metadata 和可选字段不一定具有统一 dense shape。 - -#### 阶段 5:driver 从 DataProto 取出 drafter samples - -`generate_sequences_with_speco()` 包装原 rollout 调用: - -```python -gen_batch_output = original_generate_sequences(...) -collected = self._speco_collect_generation_samples(gen_batch_output) -``` - -`_speco_collect_generation_samples()` 调用: - -```python -samples = pop_drafter_samples(gen_batch_output) -``` - -`pop_drafter_samples()` 实际执行: - -```python -non_tensor_batch = gen_batch_output.non_tensor_batch -samples_array = non_tensor_batch.pop("drafter_sample", None) -samples = normalize_drafter_samples(samples_array) -``` - -这里 `pop` 有两个作用: - -1. 取得 SpeCo drafter side-channel samples; -2. 从正常 PPO 的 `gen_batch_output` 中移除该旁路字段,避免后续 PPO batch 继续携带它。 - -`normalize_drafter_samples()` 将 dict、object array 或 list 统一成: - -```python -samples: list[dict] -``` - -#### 阶段 6:driver 按 replica_rank 分桶 - -假设有两个 rollout/drafter replicas,收到: - -```python -samples = [ - {"replica_rank": 1, "hidden_states_tq_key": "k1", ...}, - {"replica_rank": 0, "hidden_states_tq_key": "k2", ...}, - {"replica_rank": 1, "hidden_states_tq_key": "k3", ...}, -] -``` - -执行: - -```python -buckets = bucket_drafter_samples_by_replica( - samples, - num_replicas=2, -) -``` - -结果: - -```python -buckets = [ - [sample_k2], # bucket 0 - [sample_k1, sample_k3], # bucket 1 -] -``` - -分桶依据只有: - -```python -owner_rank = int(sample["replica_rank"]) -buckets[owner_rank].append(sample) -``` - -这一步没有读取 TQ,也没有处理 hidden tensor;只对小字典做路由。 - -#### 阶段 7:driver 通过 WorkerGroup RPC 分发 buckets - -driver 调用: - -```python -self._speco_set_drafter_global_step() -self._speco_collect_rollout_features_rpc( - "rollout", - buckets, -) -``` - -RPC 内部调用: - -```python -self.drafter_wg.collect_rollout_features(buckets) -``` - -因为 worker 方法注册了 `drafter_owner_route` dispatch,WorkerGroup 将 `buckets[0]` 发给 drafter DP replica 0,将 `buckets[1]` 发给 drafter DP replica 1。一个训练 replica 内的 SP ranks 根据 mesh dispatch 规则参与对应调用。 - -这里传输的对象仍是: - -```python -list[dict] -``` - -其中包含 tokens、position metadata 和 TQ key,不包含被置空的 hidden states。 - -#### 阶段 8:SpecoWorker 根据 key 从 TQ 恢复 tensor - -目标 worker 执行: - -```python -def collect_rollout_features(self, samples): - for sample in samples: - tq_key = sample.get("hidden_states_tq_key") - payload = get_sample(tq_key) - sample["hidden_states"] = payload["hidden_states"] -``` - -`get_sample()` 展开为: - -```python -tq.init() # 此 Consumer 进程第一次使用时 -result = tq.kv_batch_get( - keys=[tq_key], - partition_id="speco_drafter_features", -) -payload = _extract_value(result, tq_key) -payload = _tensordict_to_dict(payload) -``` - -现在 `sample` 再次包含: - -```python -{ - "input_ids": ..., - "hidden_positions": ..., - "hidden_states": bf16[1, hidden_rows, hidden_dim], - "hidden_states_tq_key": "...", -} -``` - -这与关闭 TQ 时 worker 收到的逻辑内容一致。 - -#### 阶段 9:worker 构造 DrafterBaseTrainer 所需 batch dict - -worker 先保留 token fields: - -```python -batch = { - "input_ids": sample["input_ids"], - "prompts": sample["prompts"], - "responses": sample["responses"], -} -``` - -再复制 hidden alignment metadata,例如: - -```python -batch["hidden_positions"] -batch["hidden_position_start"] -batch["hidden_position_end"] -batch["hidden_states_layout"] -batch["global_step"] -``` - -hidden tensor 单独作为参数: - -```python -self._store_rollout_sample( - batch=batch, - hidden_states=hidden, - target_logprobs=target_logprobs, -) -``` - -#### 阶段 10:样本进入在线 buffer 或落盘 - -`_store_rollout_sample()` 根据 training mode 分支: - -```python -if mode == "collect_only": - self._write_rollout_feature_sample( - batch, - hidden_states, - target_logprobs, - ) -else: - self.trainer.collect_online_data( - batch, - hidden_states, - target_logprobs, - ) -``` - -`collect_only` 会转换成 `DraftFeatureSample` 并写 `TorchShardFeatureStore`。online 模式则进入 `DrafterBaseTrainer.collect_online_data()`。 - -`collect_online_data()` 做: - -1. 将 `input_ids/hidden_states/positions/logprobs` 规范化到 CPU; -2. 按 batch 维拆成逐样本; -3. 根据 `hidden_positions` 校验 hidden row 与 token position; -4. 截取可训练窗口; -5. 构造内部 training item; -6. 保存到当前 step 的 `collected_data`,或在启用 data buffer 时保存到跨 step buffer。 - -因此 TQ get 完成不代表立即 optimizer step。它先恢复 online training sample,再进入现有数据准备逻辑。 - -#### 阶段 11:driver 在 actor update 周期触发 drafter 训练 - -driver 包装了 `update_actor()`: - -```python -should_train_drafter = ( - self._speco_should_attempt_drafter_train_this_step() -) - -actor_output = original_update_actor(...) - -if should_train_drafter: - drafter_trained, metrics = self._speco_train_drafter() -``` - -`_speco_train_drafter()` 再向 WorkerGroup 发: - -```python -self.drafter_wg.train_drafter() -``` - -每个 `SpecoWorker.train_drafter()`: - -1. 检查是否属于 drafter training group; -2. 检查 `training_interval_steps`; -3. 激活 drafter training model; -4. 循环 `train_steps_per_trigger` 次; -5. 每次调用 `self.trainer.training_step(global_step)`; -6. 成功时准备需要发布的 drafter state dict; -7. 清理训练期间临时状态。 - -`training_step()` 从刚才的 online `collected_data/DataBuffer` 组成 batch,执行 drafter forward、loss、backward 和 optimizer step。 - -所以 SGLang TQ 路径的最终效果是: - -```text -TQ 只替换 hidden tensor 跨进程传输 -→ sample 收集逻辑不变 -→ online buffer 不变 -→ drafter training trigger 不变 -→ loss/optimizer 不变 -``` - -### 3.7 PR #48 old-logprob 路径的完整差异 - -old-logprob 路径没有 `TokenOutput.extra_fields`。它从 PPO 的 actor old-logprob forward 开始。 - -#### 阶段 1:driver 构造 collect plan - -driver 根据 batch、collect interval 和 drafter owner 数量决定: - -```python -collect_mask: bool[batch] -hidden_positions: list/tensor per sample -owner_rank: int64[batch] -prompt_lens: int64[batch] -response_lens: int64[batch] -``` - -并把 `global_step` 等控制字段放入 old-logprob micro-batch。 - -#### 阶段 2:actor forward hook 选择 hidden rows - -actor worker 在 old-logprob forward 中捕获指定层 hidden states,根据 `collect_mask/position_mask` 只保留需要训练的样本和 token rows。 - -输出可以是 dense selected tensor,也可以是 sparse rows。随后 `_put_oldlogprob_hidden_refs()` 将同一个 owner 的多条 sample rows 拼成一个较大的 `hidden_chunk`。 - -#### 阶段 3:hidden chunk 写入 TQ - -改造前: - -```python -chunk_ref = ray.put(hidden_chunk) -``` - -PR #48: - -```python -tq_key = make_sample_key( - global_step, - owner, - f"chunk{chunk_index}", -) - -put_sample( - tq_key, - {"hidden": hidden_chunk}, - tag={"global_step": global_step, "owner": owner}, -) - -chunk_ref = tq_key -``` - -TQ fields: - -```python -{"hidden": bf16[total_owner_rows, hidden_dim]} -``` - -控制路径中的 chunk metadata: - -```python -chunk_info = { - "sample_indices": [0, 3, 5], - "starts": [0, 128, 384], - "lengths": [128, 256, 96], - "row_indices": [...], - "dtype": "bfloat16", - "shape": [480, hidden_dim], -} -``` - -`starts/lengths` 描述每条 sample 在共享 chunk 中对应的行区间。 - -#### 阶段 4:driver 将 chunk ref 还原为逐样本引用 - -driver 的 `_speco_collect_oldlogprob_features()` 读取: - -```python -chunk_refs = ["speco:42:0:chunk0", ...] -chunk_meta = [chunk_info, ...] -``` - -然后为每个 batch sample 构造: - -```python -sample["hidden_states_ref_chunks"] = [ - { - "ref": "speco:42:0:chunk0", - "chunk_start": 128, - "chunk_length": 256, - "chunk_row_indices": ..., - "dtype": "bfloat16", - "shape": [480, hidden_dim], - } -] -``` - -同时构造该 sample 的: - -```python -input_ids -prompts -responses -hidden_positions -hidden_states_layout -replica_rank=owner -``` - -再按 owner 放入 `buckets[owner]`,通过同一个 `collect_rollout_features()` RPC 发给 drafter worker。 - -#### 阶段 5:Consumer 获取共享 chunk 并切片 - -drafter worker 发现: - -```python -sample.get("hidden_states") is None -sample.get("hidden_states_ref_chunks") is not None -``` - -于是调用 `_resolve_hidden_state_chunks()`。对字符串 ref: - -```python -if ref.startswith("speco:"): - full_chunk = get_sample(ref)["hidden"] - full_chunk = _densify_tq_tensor(full_chunk) -``` - -然后按 sample metadata 取行: - -```python -sample_hidden = full_chunk[ - chunk_start : chunk_start + chunk_length -] -``` - -同一个 chunk 被多个 sample 复用,所以使用: - -```python -self._tq_chunk_cache[ref] = full_chunk -``` - -保证一次 `collect_rollout_features()` 中同一个 TQ key 只 get 一次。 - -得到逐样本 hidden 后,后续 `_store_rollout_sample → collect_online_data → train_drafter` 与 SGLang 路径相同。 - -### 3.8 PR #48 数据生命周期和清理 - -PR #48 的 TQ row 生命周期是: - -```text -TaskRunner tq.init(config) -→ Producer kv_put -→ key 经 Ray 控制路径传递 -→ 一个或多个 drafter TP/SP rank kv_batch_get -→ online drafter 收集/训练继续执行 -→ 整个 trainer.fit() 结束 -→ TaskRunner finally 调用 tq.close() -``` - -当前没有: - -```python -tq.kv_clear(key) -``` - -原因是同一 sample/key 可能被一个 drafter replica 的多个 rank 读取。若第一个 rank get 后删除,其余 rank 可能 get 失败。 - -因此 PR #48 采用任务级生命周期,而不是样本级确认和回收。这简化了并发正确性,但意味着长任务中的 TQ storage 占用可能持续增长;代码注释也把“leader 在 barrier 后精细 clear”留作后续工作。 - -### 3.9 PR #48 开启与关闭时的行为差异 - -`configure_transfer_queue()` 返回: - -```python -enabled_in_config and transfer_queue_importable -``` - -关闭时: - -```text -SGLang drafter_sample 继续内联 hidden_states -old-logprob 继续 ray.put(hidden_chunk) -Consumer 继续 ray.get/ref resolve -``` - -开启时: - -```text -SGLang hidden fields → TQ,sample 只带 key -old-logprob hidden chunk → TQ,ref 变成字符串 key -Consumer 根据 key 类型走 TQ get -``` - -如果配置开启但 `transfer_queue` 包未安装,bridge 会记录 warning,并让 `is_transfer_queue_enabled()` 返回 false,保留旧 Ray 路径。若已经进入 `put_sample/get_sample` 却发生 TQ 错误,则抛异常,不静默丢 hidden data。 - -## 第二部分:基于 PR #48 的 standalone drafter training 适配 - -从这一部分开始才讨论 `verl-SpeCo-ls` 的独立训练。下面的 `kv_list`、rank 0 选 global keys、逐 step `kv_clear`、独立 vLLM Producer 都是需要在本项目新增的逻辑,不是 PR #48 已有逻辑。 - -### 当前 standalone 基线 - -当前独立训练是: - -```text -draft_train_launcher -→ torch.distributed.run -→ 每个 rank 创建 DraftFeatureDataLoader -→ 每个 rank 在训练循环内同步 TargetFeatureReplayer.materialize() -→ vLLM/file hidden payload -→ prepare_training_batch_from_samples() -→ training_step_from_batch() -``` - -新方案要把 `TargetFeatureReplayer.materialize()` 从训练 rank 的同步取数路径移到独立 Producer,同时保持后两步训练接口不变。 - -## 4. TQ metadata 到底记录什么 - -### 4.1 Partition - -一次训练运行使用一个独立 partition: - -```python -partition_id = f"speco:{run_id}:dspark_train" -``` - -partition 用来隔离: - -- 不同训练 run; -- train 和 validation; -- 不同 target checkpoint 生成的 hidden states。 - -不能让两个 target 模型共用同一 partition,否则训练侧可能消费错误的 hidden states。 - -### 4.2 Sample key - -每条输入样本使用稳定 key: - -```python -sample_key = sha256( - dataset_id - + row_id - + prompt_token_ids - + response_token_ids - + tokenizer_fingerprint - + target_model_fingerprint - + target_layer_ids - + hidden_states_layout -).hexdigest() -``` - -稳定 key 用于: - -- vLLM HTTP 请求重试时不生成不同对象; -- Producer 重启后识别相同样本; -- 检查 hidden states 是否属于正确模型和正确层; -- TQ/Mooncake 清理时准确定位对象。 - -### 4.3 Fields 与 READY 约定 - -每个样本包含固定字段: - -```python -{ - "input_ids": int64[seq], - "loss_mask": float32[seq], - "position_ids": int64[seq], - "hidden_states": bf16[seq, aux_hidden_dim or aux_hidden_dim+hidden_size], -} -``` - -这里必须保持当前 `DraftFeatureSample` 的契约:`TargetFeatureReplayer._feature_from_vllm_payload()` 会把多个 aux layers flatten;当 layout 是 `dflash_aux_plus_last` 时,还会把 final hidden 拼到同一个 `hidden_states` tensor 尾部。`DSparkTrainerBackend.preprocess_individual_items()` 再根据 metadata 中的 `hidden_states_layout` 将 final hidden 切出来,生成训练 batch 的 `target_last_hidden_states`。 - -因此 TQ 不需要新增一个独立 `target_last_hidden_states` field。开启 DSpark L1 时,要求: - -```text -metadata.hidden_states_layout = dflash_aux_plus_last -hidden_states.shape[-1] = num_context_layers * hidden_size + hidden_size -``` - -完整 TQ native metadata API 可以做字段级 ready 判定,但 PR #48 当前走的是高层 KV API:一次 `kv_put` 把一个样本的多个 tensor fields 一起写入。因此第一版采用更直接的约定: - -```python -required_fields = [ - "input_ids", - "loss_mask", - "position_ids", - "hidden_states", -] -``` - -Producer 先在内存中验证所有必需字段,再进行一次 `kv_put`。只有 `kv_put` 成功返回,key 才会出现在 `kv_list` 结果中,并带有: - -```python -tag={ - "status": "ready", - "run_id": run_id, - "sample_id": sample_key, -} -``` - -Consumer 只选择 `status=ready` 且 `run_id` 匹配的 key。不要先 put `input_ids`、再单独 put `hidden_states`,否则 Consumer 可能观察到半成品。 - -### 4.4 Tags - -tags 是轻量 metadata,不放大 tensor: - -```python -tags = { - "sample_id": sample_key, - "source_row": row_id, - "seq_len": seq_len, - "payload_bytes": payload_bytes, - "target_model_fp": target_model_fingerprint, - "target_layers": "8,16,24", - "hidden_layout": "dflash_aux_plus_last", - "producer_status": "success", -} -``` - -tags 可用于过滤、监控、背压统计和错误排查。它不能替代 tensor shape/dtype 校验。 - -### 4.5 Run ID,而不是先依赖 task_name - -PR #48 的 `kv_put/kv_batch_get` 路径没有使用 `task_name` 或 Sampler consumption history。第一版按它的已验证接口,在 partition 和 tags 中放 `run_id`: - -```python -partition_id = "speco_drafter_features" -tag = { - "run_id": run_id, - "status": "ready", -} -``` - -不同 run 最好直接使用不同 partition: - -```python -partition_id = f"speco_drafter_features_{run_id}" -``` - -这样 job 重启和清理更简单。`task_name=dspark_train` 留到后续迁移 TQ native metadata/Sampler 时再使用。 - -### 4.6 standalone 中一条样本的完整对象形态 - -standalone 没有 PR #48 的轻量 `drafter_sample → Ray driver` 路径,因此需要让 TQ tag 承担“如何找到和解释 payload”的 metadata 作用。 - -#### Producer 读到的原始记录 - -```python -source_record = { - "dataset_id": "math-train", - "row_id": 12345, - "prompt": "...", - "response": "已经提前生成的 response", -} -``` - -#### Token replay 样本 - -分词和对齐后: - -```python -replay_sample = DraftReplaySample( - input_ids=int64[full_seq], - loss_mask=float32[full_seq], - position_ids=int64[full_seq], - feature_positions=int64[feature_rows], - draft_position_ids=int64[feature_rows], - metadata={ - "dataset_id": "math-train", - "row_id": 12345, - }, -) -``` - -这里的 `full_seq` 是 prompt 与预生成 response 拼接后的长度;`feature_positions` 指明哪些 token 位置最终进入 drafter 监督样本。 - -#### vLLM 返回的原始 hidden payload - -当前文件协议要求 safetensors 至少包含: - -```python -vllm_payload = { - "token_ids": int64[prefill_rows], - "hidden_states": bf16[prefill_rows, returned_layers, hidden_size], -} -``` - -这还不能直接给 DSpark。Producer 应复用当前: - -```python -TargetFeatureReplayer._feature_from_vllm_payload(...) -``` - -完成 token 校验、position 对齐、选层和 flatten。 - -#### Producer 最终得到的 DraftFeatureSample - -```python -feature = DraftFeatureSample( - algorithm="DSpark", - input_ids=int64[feature_rows], - loss_mask=float32[feature_rows], - position_ids=int64[feature_rows], - hidden_states=bf16[feature_rows, feature_hidden_dim], - metadata={ - "hidden_states_layout": "dflash_aux_plus_last", - "target_layer_ids": [8, 16, 24], - "target_model_path": "...", - "target_config_fingerprint": "...", - "feature_start": 128, - "feature_end": 640, - "sequence_length": 512, - }, -) -``` - -若有 3 个 aux layer、target hidden size 为 4096,并包含 final hidden: - -```text -feature_hidden_dim = 3 * 4096 + 4096 = 16384 -hidden_states.shape = [feature_rows, 16384] -``` - -前 `12288` 维是 aux/context hidden,最后 `4096` 维是 DSpark L1 所需 final hidden。现有 DSpark backend 会根据 `dflash_aux_plus_last` 自动拆分。 - -#### 写入 TQ 的 data fields - -第一版建议一个 key 对应一个已经规范化完成的 `DraftFeatureSample`: - -```python -tq_fields = { - "input_ids": feature.input_ids.cpu(), - "loss_mask": feature.loss_mask.cpu(), - "position_ids": feature.position_ids.cpu(), - "hidden_states": feature.hidden_states.cpu(), -} -``` - -这里只放 tensor,因为 PR #48 的 `put_sample()` 会过滤非 tensor: - -```python -payload = { - key: value - for key, value in tensor_dict.items() - if torch.is_tensor(value) -} -``` - -#### 写入 TQ 的 tag metadata - -```python -tq_tag = { - "run_id": "run-20260818-001", - "status": "ready", - "sample_id": sample_key, - "sequence_no": 12345, - "algorithm": "DSpark", - "hidden_states_layout": "dflash_aux_plus_last", - "target_model_fingerprint": "sha256:...", - "target_layer_ids": "8,16,24", - "feature_rows": 512, - "hidden_dim": 16384, - "payload_bytes": 16777216, -} -``` - -tag 中只放 TQ 版本支持序列化的小标量/字符串。列表等复杂对象可以编码为稳定字符串或 JSON。`payload_bytes` 用于背压统计。 - -#### TQ 中逻辑上保存的 row - -```text -partition: speco_drafter_features_run-20260818-001 -key: 86a4...ef2 - -fields: - input_ids → int64[512] - loss_mask → float32[512] - position_ids → int64[512] - hidden_states → bf16[512, 16384] - -tag: - status → ready - sequence_no → 12345 - hidden_layout → dflash_aux_plus_last - target_model_fp → sha256:... -``` - -#### Consumer 恢复出的对象 - -rank 0 通过 `kv_list` 同时取得 key 和 tag;每个 rank 用 local keys 调用 `kv_batch_get` 取得 fields,然后组合: - -```python -feature = DraftFeatureSample( - algorithm=tag["algorithm"], - input_ids=densify(fields["input_ids"]).reshape(-1), - loss_mask=densify(fields["loss_mask"]).reshape(-1), - position_ids=densify(fields["position_ids"]).reshape(-1), - hidden_states=densify(fields["hidden_states"]), - metadata={ - "hidden_states_layout": tag["hidden_states_layout"], - "target_model_fingerprint": tag["target_model_fingerprint"], - }, -) - -feature.validate(strict=True) -``` - -这样传给: - -```python -trainer.prepare_training_batch_from_samples([feature, ...]) -``` - -的数据结构,与现有文件 feature store 读出的 `DraftFeatureSample` 一致。也就是说,TQ 替换的是存取介质和调度方式,不改变 DSpark backend 的样本契约。 - -## 5. 新的整体架构 - -```text - 小 metadata - ┌──────────────────────────┐ - │ TransferQueueController │ - │ KV metadata / key / tags │ - │ partition / storage map │ - └────────────┬─────────────┘ - │ -JSONL/token replay │ - │ │ - ▼ │ -Feature Producer │ - ├─ tokenizer/window │ - ├─ asyncio bounded concurrency │ - ├─ vLLM endpoint pool │ - ├─ validate/pack │ - └─ TQ put ─────────────────────┤ - ▼ - TQ Mooncake backend - hidden-state tensors - │ - ┌───────────────────┼───────────────────┐ - ▼ ▼ ▼ - DSpark rank 0 DSpark rank 1 DSpark rank N - TQ get TQ get TQ get - └───────────────────┼───────────────────┘ - ▼ - synchronized optimizer step - │ - ▼ - TQ clear after success -``` - -大 tensor 的路径是: - -```text -vLLM/Producer memory → TQ Mooncake backend → each training rank -``` - -不会走: - -```text -Mooncake → rank 0 → rank 1/2/3 -``` - -rank 0 最多只广播 `kv_list` 得到的 key 字符串列表;tags 只在选 batch 时由 rank 0 使用。 - -### 5.1 standalone 每一步为什么能实现推理和训练异步 - -#### 步骤 A:Producer 独立推进输入 cursor - -Producer 自己维护: - -```python -reader_cursor = 12346 -``` - -它不等待 Trainer 请求样本。只要 TQ ready bytes 没超过背压上限,就继续读取文件并创建 vLLM task。 - -效果是 Producer 的执行进度与 `optimizer_step` 解耦: - -```text -Producer sequence_no: 1200,1201,1202,... -Trainer optimizer_step: 87 -``` - -两者通过 TQ 中的 ready rows 衔接,不互相直接调用。 - -#### 步骤 B:并发 vLLM task 完成顺序可以乱序 - -例如 Producer 同时提交: - -```text -sequence_no 100 → endpoint 0 -sequence_no 101 → endpoint 1 -sequence_no 102 → endpoint 0 -``` - -完成顺序可能是: - -```text -101 → 100 → 102 -``` - -每个 task 完成后独立执行 `kv_put`,所以慢请求不会阻塞已经完成的请求写入。tag 中的 `sequence_no` 保留原数据顺序。 - -#### 步骤 C:`kv_put` 成功是 READY 可见性的边界 - -Producer 在调用前已经得到完整 `DraftFeatureSample`。一次 `kv_put` 写入该 sample 的全部 tensor fields,并在 tag 中标记 `status=ready`。 - -因此 Consumer 的判断规则是: - -```text -kv_list 能列出该 key -且 tag.run_id 匹配 -且 tag.status == ready -→ 可以尝试 kv_batch_get -``` - -Consumer 仍需对取回 fields 做完整性校验;tag 是调度 metadata,不是正确性证明。 - -#### 步骤 D:rank 0 只负责选 key - -rank 0 执行: - -```python -entries = list_ready_keys() -selected = sorted(entries, key=sequence_no)[:global_batch_size] -``` - -这一步处理的数据只是: - -```python -[ - {"key": "k100", "sequence_no": 100, ...}, - {"key": "k101", "sequence_no": 101, ...}, -] -``` - -不包含 `[seq, hidden_dim]` hidden tensor,所以 rank 0 不成为大数据中转瓶颈。 - -#### 步骤 E:广播保证所有 rank 对同一个 global step 达成一致 - -所有 rank 调用同一次: - -```python -dist.broadcast_object_list(holder, src=0) -``` - -广播结束后,每个 rank 看到完全相同的 global key list。然后按确定性区间切分: - -```text -rank 0: keys[0:per_rank] -rank 1: keys[per_rank:2*per_rank] -... -``` - -这样不会出现 rank 0 训练 batch A、rank 1 训练 batch B 数量不同,或者某个 rank 没进入 backward 的情况。 - -#### 步骤 F:各 rank 直接读取 Mooncake 后端 - -每个 rank 执行: - -```python -tq.kv_batch_get(keys=local_keys, partition_id=partition_id) -``` - -TQ 根据 key 找到 storage backend 中的数据。使用 MooncakeStore 时,大 tensor 数据路径是 Mooncake → 本 rank;rank 0 不读取其他 rank 的 local payload。 - -因此: - -```text -控制面:rank 0 → broadcast small keys -数据面:Mooncake → each rank directly -``` - -#### 步骤 G:恢复现有 DraftFeatureSample 契约 - -每个 rank 将 TQ fields 与 tag 合并、densify、校验,得到 `list[DraftFeatureSample]`。从这里开始继续执行项目现有代码: - -```python -batch = trainer.prepare_training_batch_from_samples( - materialized_samples, - step=optimizer_step, -) - -ok = await trainer.training_step_from_batch( - batch, - optimizer_step, -) -``` - -所以 TQ 不进入 DSpark model/backend 内部,训练数学逻辑不变。 - -#### 步骤 H:全 rank 成功以后才能清理 - -每个 rank 的 `ok` 通过现有 `_all_ranks_true()` 聚合: - -```text -rank 0 ok = true -rank 1 ok = true -rank 2 ok = true -rank 3 ok = true -→ global_ok = true -``` - -只有此时 rank 0 执行: - -```python -tq.kv_clear(keys=global_keys, partition_id=partition_id) -``` - -这样保证被删除的数据已经参与完成的 optimizer step。若任何 rank get/OOM/backward 失败,不执行 clear,便于作业失败后的诊断或恢复。 - -#### 步骤 I:异步重叠如何形成 - -时间线上: - -```text -时间 ─────────────────────────────────────────▶ - -Producer: vLLM(batch N+1) ─ put ─ vLLM(batch N+2) ─ put -Trainer: get(batch N) ─ train(batch N) ─ get/train(batch N+1) -``` - -Producer 和 Trainer 是不同进程,互相不调用;TQ ready rows 是缓冲区。因此 vLLM prefill、网络传输和 DSpark GPU 训练可以重叠。背压只在缓冲区达到容量上限时暂停 Producer。 - -## 6. Producer:读取预生成 response 并并行请求 vLLM - -### 6.1 输入处理 - -Producer 从现有 JSONL/token replay 数据源读取: - -```python -sample = { - "row_id": "12345", - "prompt": "...", - "response": "提前生成好的文本", -} -``` - -构造: - -```python -prompt_ids = tokenizer.encode(sample["prompt"]) -response_ids = tokenizer.encode(sample["response"]) -input_ids = prompt_ids + response_ids -``` - -同时产生: - -```python -loss_mask -position_ids -feature_positions -sample_key -``` - -### 6.2 有界并发 - -不能按样本串行请求: - -```python -for sample in samples: - result = request_vllm(sample) -``` - -改成: - -```python -async def run_producer(samples): - semaphore = asyncio.Semaphore(max_inflight_requests) - - async def run_one(sample): - async with semaphore: - result = await vllm_pool.prefill(sample) - feature = validate_and_pack(sample, result) - await tq_transport.put(feature) - - async with asyncio.TaskGroup() as group: - for sample in samples: - group.create_task(run_one(sample)) -``` - -`max_inflight_requests` 是 Producer 同时未完成的请求数量,不是 vLLM batch size。vLLM 服务端仍会对同时到达的请求做自己的 continuous batching。 - -### 6.3 多 endpoint - -多个 endpoint 例如: - -```yaml -vllm_endpoints: - - http://node0:8000/v1 - - http://node1:8000/v1 - - http://node2:8000/v1 -``` - -调度器维护每个 endpoint 的 inflight 数: - -```python -endpoint = min( - endpoints, - key=lambda item: item.inflight, -) -``` - -请求前 `inflight += 1`,在 `finally` 中 `inflight -= 1`。失败只重试对应样本,不阻塞全部 Producer。 - -### 6.4 当前 vLLM 文件桥接 - -当前客户端协议期望: - -```python -response.kv_transfer_params["hidden_states_path"] -``` - -所以第一阶段仍然是: - -```text -vLLM 写临时 safetensors -→ Producer load_file -→ 校验 token_ids/hidden_states -→ TQ put 到 Mooncake backend -→ TQ put 成功后删除临时文件 -``` - -删除必须发生在 TQ put 成功之后: - -```python -path = request_vllm_hidden_file(sample) -try: - feature = load_and_validate(path) - await tq_transport.put(feature) -finally: - if put_succeeded: - Path(path).unlink(missing_ok=True) -``` - -### 6.5 目标版本:vLLM 直接写 TQ/Mooncake - -目标响应可改成: - -```json -{ - "kv_transfer_params": { - "backend": "transfer_queue", - "partition_id": "speco:run-1:dspark_train", - "sample_key": "abc123" - } -} -``` - -服务端顺序必须是: - -```text -prefill -→ 捕获指定层 hidden states -→ TQ/Mooncake put 完成 -→ 返回 HTTP success 和 sample key -``` - -这需要定制 vLLM exporter;当前 `verl-SpeCo-ls` 中没有服务端 writer 实现。 - -## 7. 按 PR #48 扩展 TQ bridge - -不要重新发明一套 transport。将 PR #48 的 bridge 设计移植到本项目并增加 standalone 所需方法: - -```python -class StandaloneTQTransport: - def put_sample(self, key, tensor_dict, tag): ... - def list_ready_keys(self, run_id): ... - def get_samples(self, keys, fields=None): ... - def clear_samples(self, keys): ... - def put_control(self, key, tag): ... - def close(self): ... -``` - -写入延续 PR #48 的真实形式: - -```python -tq.kv_put( - key=key, - partition_id=partition_id, - fields={ - "input_ids": feature.input_ids.cpu(), - "loss_mask": feature.loss_mask.cpu(), - "position_ids": feature.position_ids.cpu(), - "hidden_states": feature.hidden_states.cpu(), - }, - tag={ - "run_id": run_id, - "status": "ready", - "sequence_no": sequence_no, - "payload_bytes": payload_bytes, - }, -) -``` - -批量读取延续 PR #48 的 `kv_batch_get`: - -```python -result = tq.kv_batch_get( - keys=keys, - partition_id=partition_id, - fields=required_fields, # 0.1.7 是否支持该参数需实机确认 -) -``` - -新增发现和清理: - -```python -items = tq.kv_list(partition_id=partition_id) -tq.kv_clear(keys=keys, partition_id=partition_id) -``` - -这里的 `kv_list/kv_clear` 参数名必须根据锁定的 TQ 版本验证。当前参考仓库只实际调用了 `kv_put/kv_batch_get`,没有为这两个接口提供运行证据。 - -读取结果继续复用 PR #48 的两个适配函数: - -```python -value = _extract_value(result, key) -row = _tensordict_to_dict(value) -row["hidden_states"] = _densify_tq_tensor(row["hidden_states"]) -``` - -## 8. DSpark 多 rank 如何消费 - -### 8.1 第一版:rank 0 用 kv_list 发现 READY keys - -PR #48 中 key 由 Ray driver 传给 drafter worker;standalone 没有这条控制路径,所以 rank 0 需要主动列出 key: - -```python -rank = dist.get_rank() -world_size = dist.get_world_size() -global_batch_size = batch_size_per_gpu * world_size - -if rank == 0: - entries = tq_transport.list_ready_keys(run_id=run_id) - entries.sort(key=lambda x: (x.tag["sequence_no"], x.key)) - selected_keys = [x.key for x in entries[:global_batch_size]] -else: - selected_keys = None - -holder = [selected_keys] -dist.broadcast_object_list(holder, src=0) -selected_keys = holder[0] -``` - -`sequence_no` 是输入文件顺序。它让多个 Producer 并发完成顺序不同的情况下,Trainer 仍能确定性地组成 batch。 - -rank 0 广播的是字符串 key 列表,不是 hidden-state tensor。 - -### 8.2 各 rank 切自己的 keys - -例如 global batch keys: - -```text -[s0, s1, s2, s3, s4, s5, s6, s7] -``` - -world size 为 4、每卡 batch size 为 2: - -```text -rank 0 → [s0, s1] -rank 1 → [s2, s3] -rank 2 → [s4, s5] -rank 3 → [s6, s7] -``` - -代码: - -```python -def shard_keys(keys, rank, world_size): - assert len(keys) % world_size == 0 - per_rank = len(keys) // world_size - start = rank * per_rank - end = start + per_rank - return keys[start:end] -``` - -### 8.3 每个 rank 并行 get - -所有进程执行: - -```python -local_keys = shard_keys( - selected_keys, - rank=rank, - world_size=world_size, -) - -local_payloads = tq_transport.get_samples(local_keys) -``` - -数据路径: - -```text -rank 0 ← Mooncake(s0,s1) -rank 1 ← Mooncake(s2,s3) -rank 2 ← Mooncake(s4,s5) -rank 3 ← Mooncake(s6,s7) -``` - -不是 rank 0 get 全部后再 scatter。 - -### 8.4 转成当前训练格式 - -TQ 返回的数据需要先按 PR #48 的规则解包、densify,再转换成现有 `DraftFeatureSample`: - -```python -def tq_row_to_feature(row, tag): - return DraftFeatureSample( - algorithm="DSpark", - input_ids=row["input_ids"], - loss_mask=row["loss_mask"], - position_ids=row["position_ids"], - hidden_states=row["hidden_states"], - metadata={ - "hidden_states_layout": tag["hidden_states_layout"], - **row.get("metadata", {}), - }, - ) -``` - -`hidden_states_layout` 不能只放在无法恢复的临时 Python 对象里;应随 TQ tag 保存。rank 0 从 `kv_list` 取得 key/tag 后,需要把每个 local key 对应的 tag 一起交给本 rank 的转换逻辑。Consumer 必须把 layout 放回 `DraftFeatureSample.metadata`。DSpark L1 所需的 `target_last_hidden_states` 随后由现有 backend 从拼接的 `hidden_states` 中切出,transport 层不计算 loss,也不拆 layout。 - -## 9. 修改当前训练循环 - -在 `run_standalone_draft_training()` 中增加数据源分支: - -```python -feature_store_type = str(feature_store_cfg.get("type", "torch_shard")) - -if feature_store_type == "transfer_queue": - tq_stream = build_transfer_queue_stream( - config=config, - rank=rank, - world_size=world_size, - ) - store = None - loader = None - feature_replayer = None -else: - store = build_feature_store_from_config( - feature_store_cfg, - read_only=True, - ) - loader = DraftFeatureDataLoader(...) -``` - -流式训练循环: - -```python -while successful_steps < max_steps: - global_keys, materialized_samples = tq_stream.next_local_batch() - - batch = trainer.prepare_training_batch_from_samples( - materialized_samples, - step=optimizer_step, - ) - - has_batch = batch is not None - if not _all_ranks_true(has_batch, trainer.runtime_device): - raise RuntimeError("at least one rank failed to fetch its TQ batch") - - ok = await trainer.training_step_from_batch( - batch, - optimizer_step, - ) - - if not _all_ranks_true(ok, trainer.runtime_device): - raise RuntimeError("DSpark step failed on at least one rank") - - dist.barrier() - if rank == 0: - tq_stream.clear_global_batch(global_keys) - dist.barrier() -``` - -TQ 接在数据输入层,不放进 `DSparkTrainerBackend`。Backend 只负责模型、forward、loss、backward 和 optimizer。 - -## 10. READY key、inflight key 和训练提交 - -### 10.1 Ready - -在本方案中 ready 表示: - -> `dspark_train` 所需的所有 fields 已经成功写入 TQ storage,训练侧可以读取。 - -第一版不是通过 PR #48 尚未验证的 Sampler 自动判定,而是 Producer 完整校验 payload 后一次 `kv_put`,并设置 `tag.status=ready`。rank 0 的 `kv_list` 只选择这些 key。 - -### 10.2 Inflight key - -rank 0 选出一个 global batch 后,在本地保存: - -```python -inflight_global_keys = selected_keys -``` - -其他 step 不应再次选中这批 key。因为 standalone 只有一个同步 DSpark job,第一版由 rank 0 进程内 `set` 排除 inflight key 即可: - -```python -ready = [x for x in listed if x.key not in inflight_keys] -``` - -若以后允许多个独立 Trainer job 同时消费同一 partition,进程内 set 就不够,届时必须使用 TQ Controller/Sampler 的原子消费分配能力。 - -### 10.3 Optimizer committed - -optimizer committed 表示所有 DSpark rank 已经完成: - -```text -forward → backward → gradient synchronization → optimizer.step -``` - -它比 `kv_list` 发现 key、甚至 `kv_batch_get` 取出 tensor 都更晚。 - -第一版推荐简单语义: - -```text -TQ 负责 key/tag 和 tensor 传输 -rank 0 负责单 Trainer job 的 batch 选择和 inflight set -训练失败 → 整个作业 fail-fast -训练成功 → kv_clear payload,并从 inflight set 移除 -恢复 → 从最近 checkpoint + 输入 cursor 重新启动 -``` - -这样不需要重新实现 lease/ack 状态机,也没有夸大 PR #48 当前尚未使用的 Sampler 能力。 - -## 11. 为什么训练完一个 step 才清理 - -不能在 `kv_batch_get()` 后立即 clear: - -```text -get 成功 -→ clear -→ forward OOM -→ 数据已不存在,无法重试 -``` - -正确顺序: - -```text -rank 0..N get -→ 所有 rank 确认 batch 有效 -→ training_step_from_batch -→ _all_ranks_true(ok) -→ rank 0 kv_clear global keys -``` - -当前 `_all_ranks_true()` 已经是项目现有的跨 rank 同步工具,可以继续复用。 - -## 12. 背压 - -背压表示 Producer 生成速度高于 Trainer 消费速度时,Producer 必须暂停继续提交,以免 Mooncake 内存无限增长。 - -建议限制: - -```yaml -max_vllm_inflight_requests: 32 -max_pending_put_bytes: 8589934592 -max_tq_ready_samples: 256 -max_tq_ready_bytes: 68719476736 -``` - -Producer 在 tags 中写: - -```python -{"payload_bytes": payload_bytes} -``` - -周期性通过 TQ list/metadata 统计当前 partition 尚未消费的数据量: - -```python -while ready_bytes >= max_tq_ready_bytes: - await asyncio.sleep(backpressure_poll_interval) -``` - -如果锁定的 TQ 版本提供容量/ready 统计接口,应直接使用,避免全量扫描 keys。 - -## 13. Stable ID、幂等和孤儿数据 - -### 13.1 幂等 - -幂等表示同一操作重复执行,最终逻辑结果仍只有一份。 - -Producer 对同一样本重试时必须使用相同 `sample_key`: - -```python -await tq.put(key="abc123", ...) -await tq.put(key="abc123", ...) -``` - -不能每次生成随机 key: - -```text -abc123-retry-1 -abc123-retry-2 -``` - -否则一个输入可能训练多次并持续占用 Mooncake。 - -### 13.2 孤儿数据 - -孤儿数据表示 tensor 已经写进 storage,但由于 Producer 崩溃或 metadata 更新失败,没有进入正常消费路径。 - -使用 TQ 后,metadata 和 storage 由同一套系统管理,可以减少“Mooncake有对象、自研 Coordinator 没记录”的双系统窗口,但仍要配置 TTL/partition cleanup: - -```text -训练正常结束 → clear partition -训练异常退出 → 下次启动检查旧 partition -超过 TTL → 清理未消费数据 -``` - -## 14. EOS 和 drop-last - -EOS 表示 Producer 已经读完输入,并且所有 vLLM/TQ put 都已完成。 - -TQ 中需要一种结束条件,具体使用 Controller API、特殊 metadata 或单独的运行状态记录取决于锁定版本。不能把“当前暂时没有 ready sample”当成 EOS,因为 Producer 可能仍在请求 vLLM。 - -最后不足一个 global batch 时: - -```python -global_batch_size = batch_size_per_gpu * world_size -``` - -第一版使用 `drop_last=true`,避免某些 rank 有数据、某些 rank 没数据导致分布式训练不同步。 - -结束条件: - -```text -producer_done == true -and ready_samples < global_batch_size -and inflight_requests == 0 -and pending_puts == 0 -``` - -## 15. 双缓冲预取 - -训练 batch N 时,CPU 后台线程预取 batch N+1: - -```python -next_future = executor.submit(tq_stream.next_local_batch) - -current_batch = first_batch -while current_batch is not None: - next_batch = next_future.result() - next_future = executor.submit(tq_stream.next_local_batch) - - train(current_batch) - current_batch = next_batch -``` - -实际顺序应调整为避免等待 future 后才训练。推荐: - -```python -current = tq_stream.next_local_batch() - -while current is not None: - future = executor.submit(tq_stream.next_local_batch) - train_and_clear(current) - current = future.result() -``` - -第一版只预取一个 global batch,避免 Trainer 崩溃时大量数据已被采样但未训练。 - -如果 Mooncake/TQ Python get 是阻塞函数,使用专用 `ThreadPoolExecutor`,不要阻塞 Producer 的 asyncio event loop。 - -## 16. 建议代码结构 - -```text -verl_speco/ - trainer/ - tq_transport.py # TQ client、put/get/meta/clear 封装 - tq_feature_stream.py # kv_list、global keys、rank shard、decode/prefetch - feature_producer.py # JSONL → 并发 vLLM → TQ - draft_training_loop.py # 增加 transfer_queue 数据源分支 - target_feature_replay.py # 复用 vLLM payload 校验/转换逻辑 -``` - -不要新增: - -```text -coordinator.py -coordinator_client.py -``` - -建议抽象: - -```python -class StreamingFeatureSource(Protocol): - def next_local_batch(self) -> tuple[list[str], list[DraftFeatureSample]]: ... - def clear_global_batch(self, keys: list[str]) -> None: ... - def close(self) -> None: ... -``` - -这样训练循环不依赖 TQ 的具体类型。 - -## 17. 配置草案 - -```yaml -actor_rollout_ref: - rollout: - drafter: - training: - mode: offline - backend: dspark - batch_size_per_gpu: 2 - max_steps: 1000 - - feature_store: - type: transfer_queue - partition_id: speco_drafter_features_${run_id} - drop_last: true - prefetch_steps: 1 - - transfer_queue: - # 与 PR #48 的配置层级和 init 方式保持一致。 - enable: true - package_version: 0.1.8 # 最终以实测版本为准 - backend: - storage_backend: MooncakeStore - MooncakeStore: - auto_init: false - metadata_server: localhost:50123 - master_server_address: localhost:50124 - local_hostname: localhost - protocol: tcp - global_segment_size: 4294967296 - local_buffer_size: 1073741824 - device_name: "" - - required_fields: - - input_ids - - loss_mask - - position_ids - - hidden_states - - producer: - input_path: /path/to/generated_responses.jsonl - vllm_endpoints: - - http://node0:8000/v1 - - http://node1:8000/v1 - max_inflight_requests: 32 - max_pending_put_bytes: 8589934592 - max_ready_samples: 256 - max_ready_bytes: 68719476736 -``` - -当前 examples 中的: - -```bash -transfer_queue.enable=False -``` - -属于 verl RL 主入口配置,当前 standalone `draft_train_launcher` 不读取它。不能只改成 `True`;必须实现上述 `feature_store.type=transfer_queue` 分支。 - -## 18. 启动顺序 - -逻辑顺序: - -```text -1. 启动 Mooncake metadata/master 服务; -2. 启动 standalone TQ owner 进程,调用一次 `tq.init(tq_config)` 创建/连接 TQ Controller 和 MooncakeStore backend; -3. 启动一个或多个定制 vLLM server -4. 启动 Feature Producer -5. Producer 和所有 Trainer rank 调用无参 `tq.init()`,连接 owner 创建的同一套 TQ; -6. 启动 verl_speco.draft_train_launcher -7. torchrun 启动所有 DSpark rank -8. 各 rank 连接 TQ -9. rank 0 用 `kv_list` 选择 global keys,各 rank 并行 `kv_batch_get`/train; -10. 输入耗尽后 Producer 发布 done 状态 -11. Trainer drain 完整 global batches 后退出 -12. 清理 partition,停止 Producer、vLLM、TQ、Mooncake -``` - -PR #48 的 `init_transfer_queue()` 在 Ray `SpecoTaskRunner` 中调用 `tq.init(config)`,worker 的无参 `tq.init()`依靠同一个 Ray 集群发现 named Controller;这部分不能原样复制到 standalone torchrun。 - -本项目要求一开始不使用 Ray,因此 Phase 0 必须先证明锁定的 TQ 版本支持独立 owner/controller 进程,以及 Producer/torchrun rank 如何获得连接信息。若 TQ 0.1.7/0.1.8 实际只能通过 Ray named actor 完成发现,那么有两个选择: - -1. 接受仅用 Ray 承载 TQ 控制面的最小方案; -2. 给 TQ 增加或使用其已有的独立 ZMQ/server-info bootstrap。 - -在这项验证完成前,文档不能声称 PR #48 已经提供“无 Ray TQ standalone 启动”。 - -## 19. 故障处理 - -### vLLM 请求失败 - -- 对单个 sample 按稳定 key 重试; -- 指数退避; -- 超过次数记录失败,并根据配置 fail-fast 或跳过; -- 不写不完整 TQ fields。 - -### vLLM 文件读取成功,但 TQ put 失败 - -- 暂时保留临时文件; -- 重试 TQ put; -- put 成功后再删除; -- 不把样本视为 ready。 - -### 某个训练 rank get 失败 - -- 该 rank 报告 `local_ok=false`; -- `_all_ranks_true()` 使全部 rank 得到一致失败结果; -- 第一版整个训练 fail-fast; -- 不 clear global batch。 - -### OOM/optimizer step 失败 - -- 不 clear; -- 所有 rank 一致退出; -- 从最近训练 checkpoint 恢复; -- 根据 TQ 消费提交语义决定是否重放当前 batch。 - -### clear 失败 - -- optimizer 已成功,不能再次训练这批; -- 将 batch keys 写入本地小型 `gc_pending` 日志; -- 后台重试 clear; -- checkpoint 保存最近 committed sample IDs,避免恢复时重复消费。 - -## 20. 观测指标 - -Producer: - -```text -producer/vllm_inflight -producer/vllm_requests_per_sec -producer/vllm_prefill_tokens_per_sec -producer/vllm_p50_latency -producer/vllm_p95_latency -producer/tq_put_bytes_per_sec -producer/tq_put_failures -producer/pending_put_bytes -``` - -TQ/Mooncake: - -```text -tq/ready_samples -tq/ready_bytes -tq/consumed_samples -tq/storage_bytes -tq/clear_failures -mooncake/put_bandwidth -mooncake/get_bandwidth -``` - -Trainer: - -```text -trainer/tq_wait_seconds -trainer/tq_get_seconds -trainer/tq_get_bytes_per_sec -trainer/decode_seconds -trainer/h2d_seconds -trainer/step_seconds -trainer/data_stall_ratio -trainer/successful_steps -``` - -## 21. 实施阶段 - -### Phase 0:锁定依赖和契约 - -- 从 PR #48 的 `TransferQueue==0.1.7` 起验证,同时对比当前 verl 文档使用的 0.1.8; -- 实测 `tq.init(config)`、worker `tq.init()`、`kv_put`、`kv_batch_get`、`kv_list`、`kv_clear`; -- 验证该版本是否支持无 Ray owner/controller 部署以及连接信息传递; -- 先用 PR #48 已配置的 `SimpleStorage` 做最小闭环; -- 再把 backend 换成 `MooncakeStore`,验证 tcp,最后再验证 rdma; -- 固定 fields、tags、partition 和 `run_id/sequence_no/status`; -- 写 fake TQ 单元测试。 - -### Phase 1:文件桥接 + TQ KV 模式 - -- 新增独立 Producer; -- 32 个有界并发 vLLM 请求; -- 读取 vLLM 临时 safetensors; -- TQ put 成功后删除文件; -- standalone trainer 由 rank 0 `kv_list` 并广播 global keys; -- 各 rank 并行 `kv_batch_get`; -- 复用 PR #48 的返回值解包、NestedTensor densify 和 per-step cache; -- optimizer 成功后 `kv_clear`。 - -验收:连续训练 1000 step,临时文件数量、TQ ready bytes 和 Mooncake占用均保持有界。 - -### Phase 2:双缓冲与多 endpoint - -- 增加多 endpoint 最少 inflight 调度; -- 增加一个 global batch 预取; -- 动态背压; -- 注入单 rank get 失败,确认所有 rank 一致退出而非死锁。 - -### Phase 3:vLLM 直接写 TQ/Mooncake - -- 修改外部定制 vLLM exporter; -- 去掉 `hidden_states_path` 临时文件; -- HTTP 响应返回 partition/sample key; -- 验证 HTTP 重试的幂等性。 - -### Phase 4:可选升级到 TQ StreamingDataLoader - -- 在当前保守方案稳定后再引入 RankAwareSampler; -- 让每个 rank 自动取得 local micro-batch; -- 去掉 rank 0 手工 key-list 广播; -- 验证与 torchrun/DSpark 的 global step 对齐。 - -## 22. 最终推荐 - -针对当前 `verl-SpeCo-ls`,推荐的第一版不是 AngelSpec 的完整架构,也不是单独写 Coordinator,而是: - -```text -当前预生成 response 文件 -→ 独立 asyncio Producer -→ 并行访问多个 vLLM endpoint -→ 读取并校验临时 hidden-state 文件 -→ TransferQueue put -→ Mooncake storage backend -→ rank 0 kv_list 获取 READY global keys -→ broadcast key list -→ 各 DSpark rank 并行 kv_batch_get -→ 现有 prepare_training_batch_from_samples() -→ 现有 training_step_from_batch() -→ 全 rank 成功 -→ TQ clear -``` - -这套方案保留当前 standalone DSpark 训练主体,只替换 `DraftFeatureDataLoader + TargetFeatureReplayer.materialize()` 所在的数据输入路径。它真正建立在 PR #48 已实现的 KV transport 之上,而不是假设 PR #48 已经实现了 standalone Sampler/StreamingDataLoader。 - -## 23. 参考 - -- verl TransferQueue: -- TransferQueue: -- Mooncake Store: -- 本地只读参考:`../verl-SpeCo/verl_speco/integration/transferqueue_bridge.py`(PR #48) -- 本地只读参考:`../verl-SpeCo/docs/transferqueue_integration_plan.md`(PR #48) diff --git a/docs/standalone_tq_consumer_implementation.md b/docs/standalone_tq_consumer_implementation.md deleted file mode 100644 index 78075686..00000000 --- a/docs/standalone_tq_consumer_implementation.md +++ /dev/null @@ -1,713 +0,0 @@ -# 独立 DSpark 训练 TQ Consumer 实现说明 - -Last updated: 08/21/2026 - -## 1. 文档范围和当前结论 - -本文只说明当前仓库中已经实现的独立训练 Consumer。这里的 Consumer 是由 `torchrun` 启动的 DSpark 草稿模型训练任务:它持续从 TransferQueue(下文简称 TQ)发现样本,各训练 rank 分别取得自己负责的 Tensor,复用原有 DSpark 训练逻辑完成一次 optimizer step,然后由 rank 0 删除这一整个 global batch 对应的 TQ 记录。 - -当前已完成的能力是: - -1. `feature_store.type=tq` 可以作为独立训练的数据源,不要求磁盘 `path`。 -2. 每个训练 rank 都连接同一个 Ray 集群、同一个 TQ Controller 和同一个 partition。 -3. 只有 rank 0 调用 `kv_list` 发现 ready key,并把 key/tag 分配给各 rank。 -4. key 和 tag 通过 `torch.distributed.broadcast_object_list` 传输;hidden states 等 Tensor 不经过该广播。 -5. 每个 rank 根据分配到的 key,直接调用 TQ `kv_batch_get` 获取自己的 Tensor。 -6. TQ Tensor 被解码成原训练代码已经认识的 `DraftFeatureSample`,然后复用 `DrafterBaseTrainer` 的 batch 构造、DSpark loss、反向传播和 optimizer step。 -7. 只有当所有 rank 都成功完成该 step 后,rank 0 才调用 `kv_clear` 删除整个 global batch。 -8. Producer 发布 EOS 后,如果剩余样本不足一个 global batch,当前第一版会丢弃并清理这部分尾样本,然后正常结束训练迭代。 - -本文不会把尚未实现的 Producer 写成现有能力。Producer 后续需要复用本文第 6 节所述的公共协议,调用 `encode_sample()` 生成 fields,再使用 bridge 写入相同 TQ。 - -## 2. 本次涉及的文件 - -### 2.1 本次新增的 Consumer 核心文件 - -| 文件 | 实现的组件 | 作用 | -|---|---|---| -| `verl_speco/trainer/tq_feature_store.py` | `TQFeatureStore`、`ReadyEntry`、`EosMetadata` | 将公共 TQ bridge 包装成 Consumer 数据访问层,负责连接、发现、批量读取、解码、删除和读取 EOS | -| `verl_speco/trainer/tq_sample_source.py` | `TQFeatureDataLoader`、`TQLocalBatch`、`build_assignments()` | 实现多 rank 流式取数:rank 0 发现样本并分配 key,各 rank 自己从 TQ 取 Tensor | -| `tests/unit/test_tq_consumer.py` | Consumer 单元测试 | 覆盖 store 构建、ready 过滤排序、协议解码、rank 分配、EOS、清理和非 rank 0 行为 | - -### 2.2 本次修改的既有文件 - -| 文件 | 修改内容 | 为什么要改 | -|---|---|---| -| `verl_speco/trainer/feature_store.py` | factory 新增 `type=tq` 分支 | 让既有独立训练入口能够像选择磁盘 feature store 一样选择流式 TQ 数据源 | -| `verl_speco/trainer/draft_training_loop.py` | 接入 TQ store/loader、跨 rank 连接检查、训练成功后清理 | 将流式取数接入原训练循环,同时保留原 DSpark trainer、loss、optimizer、metric 和 checkpoint 逻辑 | -| `verl_speco/draft_train_launcher.py` | 增加 TQ 启动参数的 fail-fast 检查 | 在启动多个 torchrun 子进程前检查 `enable`、Ray address 和 `run_id`,避免各 rank 启动后才失败 | -| `verl_speco/config/speco_base.yaml` | 标注 `feature_store.type=tq` 为无路径流式数据源 | 保留统一 Hydra 配置入口;TQ 的公共配置仍位于 sibling `training.transfer_queue` | -| `tests/unit/test_draft_train_launcher.py` | 增加 TQ 参数检查测试 | 验证必要配置缺失时 launcher 直接拒绝启动 | -| `tests/unit/test_draft_training_loop.py` | 增加连接和 clear 时序测试 | 验证只由 rank 0 清理、clear 失败会报告、连接失败会传播 | - -### 2.3 直接复用的公共基础 - -以下文件不是这次 Consumer 才创造的概念,但 Consumer 直接使用它们: - -| 文件 | 被复用的能力 | -|---|---| -| `verl_speco/integration/transferqueue_bridge.py` | 屏蔽 TQ 0.1.7 API 细节,提供连接、`kv_list`、`kv_batch_get`、`kv_clear` 和本地关闭接口 | -| `verl_speco/transport/drafter_sample_protocol.py` | 定义 key、tag、fields、metadata 格式,以及 `encode_sample()` / `decode_sample()` | -| `verl_speco/trainer/feature_store.py` | 复用 `DraftFeatureSample`,使 TQ 数据进入训练侧后与磁盘 feature sample 类型一致 | -| `verl_speco/trainer/base_trainer.py` 及既有 backend | 复用 `DrafterBaseTrainer.prepare_training_batch_from_samples()` 和 `training_step_from_batch()` 等训练实现 | - -## 3. 运行时角色 - -### 3.1 TQ Owner - -TQ Owner 是单独的普通 Python 进程。它连接指定 Ray 集群,并以带配置的 `tq.init(config)` 创建任务级 named Controller 和 storage actors。Owner 持有全局 TQ 生命周期;Consumer 结束时不能关闭它。 - -Owner 不是训练 rank,也不执行 DSpark 模型。它的主要作用是让 Producer 和 Consumer 能通过同一个 Ray actor registry 找到同一个 TQ Controller。 - -### 3.2 Producer - -Producer 是后续需要实现的独立推理进程。它应并行调用 vLLM hidden-state 接口,构造一条条 `DraftFeatureSample` 和 `SampleMetadata`,再写入 TQ。 - -Producer 与 Consumer 不通过 Ray RPC 互相调用,也不通过 HTTP 直接传 Tensor。二者只需满足: - -- 连接同一个 Ray address; -- 使用同一个 Ray namespace; -- 使用同一个 TQ partition; -- 使用同一个 `run_id` 和协议版本。 - -### 3.3 Consumer launcher - -`python -m verl_speco.draft_train_launcher` 是父进程。它检查命令行 override,构造 `python -m torch.distributed.run ...` 命令,然后启动训练子进程。 - -launcher 自己不连接 TQ、不取样本、也不持有 GPU 模型。 - -### 3.4 Consumer training rank - -`torchrun --nproc_per_node=N` 会启动 N 个训练 OS 进程。每个进程有独立的: - -- global rank; -- local rank; -- GPU; -- `DrafterBaseTrainer`; -- `TQFeatureStore` 和本地 TQ client; -- DSpark 模型分片及 optimizer 状态。 - -这些 rank 共同执行一个分布式草稿模型训练任务。rank 0 额外负责发现和删除 TQ key;但所有 rank 都会取得各自的训练 Tensor,并参加模型 collective、梯度同步和 optimizer step。 - -### 3.5 Ray 和 torch.distributed 的职责不同 - -本方案仍然使用 Ray,但只因为 TQ 0.1.7 通过 Ray named actor 找 Controller。Consumer 不创建用于训练的 Ray actor,训练本身仍由 `torchrun` 和 `torch.distributed` 执行。 - -两种通信分别是: - -- Ray/TQ:Owner、Producer、每个 Consumer rank 连接共享 TQ;大 Tensor 通过 TQ backend 传输。 -- `torch.distributed`:训练 rank 之间广播小型 key/tag 命令、同步成功状态、训练模型 collective。 - -## 4. 共同配置以及“连接同一个 TQ”的实现 - -关键配置位于: - -```yaml -actor_rollout_ref: - rollout: - drafter: - training: - mode: offline - feature_store: - type: tq - path: null - transfer_queue: - enable: true - ray: - address: 127.0.0.1:6379 - namespace: speco-drafter - partition_id: speco_drafter_features - run_id: dspark-standalone-run - schema_version: 1 - poll_interval_seconds: 0.5 - drop_last: true -``` - -这些字段的含义如下: - -| 字段 | 使用者 | 含义 | -|---|---|---| -| `feature_store.type=tq` | Consumer | 选择流式 TQ source,而不是磁盘 shard/replay source | -| `feature_store.path=null` | Consumer | TQ 不从本地路径读文件,因此无需 path | -| `transfer_queue.enable` | Owner、Producer、Consumer | 开启 bridge 的 TQ 路径 | -| `ray.address` | 三端 | 连接同一个 Ray 集群 | -| `ray.namespace` | 三端 | 在同一 actor namespace 查找 named Controller | -| `partition_id` | 三端 | 对同一个 TQ KV 分区执行 put/list/get/clear | -| `run_id` | Producer、Consumer | 在共享 partition 中区分本次训练数据;Consumer 只接收匹配的样本 | -| `schema_version` | Producer、Consumer | 共同使用的数据协议版本 | -| `poll_interval_seconds` | Consumer rank 0 | ready 数量不足时的轮询间隔 | -| `drop_last` | Consumer | 第一版必须为 true;EOS 后不足 global batch 的尾样本被清理 | - -三端并不是通过共享 Python 对象得到这些配置。每个进程都各自读取相同取值,然后执行: - -```python -configure_transfer_queue(config) -connect_ray_cluster(ray_address, ray_namespace) -connect_transfer_queue_client() -``` - -`connect_ray_cluster()` 内部调用 `ray.init(address=..., namespace=...)`。`connect_transfer_queue_client()` 再使用与 Owner 相同的 native 配置调用 `tq.init(config)`;TQ 会优先在当前 Ray namespace 查找 Owner 创建的 named Controller,找到时忽略本次配置并只创建本地 Client。即使 Client 意外先于 Owner 初始化,也会使用同一份 backend/controller 配置,而不会按默认配置创建服务。之后所有 KV 操作都显式携带相同的 `partition_id`。 - -因此,“连接同一个 TQ”实际由三层身份共同决定:同一 Ray 集群、同一 namespace 下的同一 named Controller、同一 `partition_id`。 - -## 5. 一条样本在 TQ 中的实际格式 - -### 5.1 一条 key 对应一个 sample - -本协议没有把一个训练 batch 存成一个 TQ key。一条 key 对应一条独立训练样本。假设: - -```text -run_id = dspark-run-001 -sequence_no = 17 -sample_id = prompt-000017 -``` - -则 key 为: - -```text -drafter:v1:dspark-run-001:000000000017:prompt-000017 -``` - -`sequence_no` 是本次 run 内的样本顺序号,不是 batch 编号,也不是 optimizer step。Consumer 用它稳定排序,之后每次从有序 ready 列表前部取一个 global batch。 - -### 5.2 tag:用于轻量发现和过滤 - -该 key 的 tag 是普通小字典: - -```python -{ - "record_type": "sample", - "status": "ready", - "schema_version": 1, - "run_id": "dspark-run-001", - "sequence_no": 17, - "sample_id": "prompt-000017", - "algorithm": "DSPARK", -} -``` - -tag 存在 TQ 的 KV 元信息中。`kv_list(partition_id=...)` 返回 `key -> tag`,不需要先加载 hidden states。rank 0 正是依靠 tag 筛选当前 run、当前 schema、DSPARK 且状态为 ready 的记录。 - -### 5.3 fields:真正的 Tensor payload - -同一 key 的 fields 是一个 Tensor 字典: - -```python -{ - "input_ids": Tensor[int64, shape=[L]], - "loss_mask": Tensor[float32, shape=[L]], - "position_ids": Tensor[int64, shape=[L]], - "hidden_states": Tensor[dtype, shape=[L, D]], - "metadata_json": Tensor[uint8, shape=[M]], - # 以下是可选字段: - "last_hidden_states": Tensor[..., ...], - "target": Tensor[..., ...], - "target_logprobs": Tensor[..., ...], -} -``` - -这里 `L` 是 feature window 的 token 数,`D` 是目标模型 hidden size,`M` 是 metadata JSON 序列化后的 UTF-8 字节数。 - -`hidden_states` 等大 Tensor 只存于 fields,通过 TQ `kv_batch_get` 传输;不会放进 tag,也不会通过训练 rank 的 object broadcast。 - -### 5.4 metadata_json:内容丰富但仍随 fields 读取 - -TQ fields 只能承载 Tensor,因此结构化 metadata 被编码为 `uint8` Tensor。解码后的字典格式是: - -```python -{ - "schema_version": 1, - "run_id": "dspark-run-001", - "sample_id": "prompt-000017", - "sequence_no": 17, - "algorithm": "DSPARK", - "target_model_id": "/models/Qwen3-8B", - "target_model_revision": "main", - "tokenizer_fingerprint": "...", - "target_layer_ids": [35], - "hidden_states_layout": "token_major", - "hidden_dtype": "bfloat16", - "hidden_shape": [L, D], - "feature_length": L, - "full_sequence_length": 256, - "feature_start": 64, - "feature_end": 64 + L, - "use_logits": False, -} -``` - -字段分工是: - -- tag:只放发现、过滤、排序所需的小字段;`kv_list` 可直接得到。 -- fields:放训练 Tensor 和完整 metadata;只有被某个 rank 选中后才 `kv_batch_get`。 -- key:把 tag 和 fields 重新关联起来,也是清理记录时传给 `kv_clear` 的标识。 - -### 5.5 控制记录 - -控制记录与 sample 放在同一 partition,但通过 tag 的 `record_type=control` 区分。 - -Owner readiness key: - -```text -control:v1::owner-ready -``` - -EOS key: - -```text -control:v1::eos -``` - -EOS tag 包含 `status=eos` 和 `total_samples`。EOS 表示 Producer 不会再为本次 run 增加新样本;它不是一条训练样本。 - -## 6. 公共协议如何把 Producer 输出还原为训练对象 - -Producer 应调用: - -```python -fields = encode_sample(sample, metadata) -key = make_sample_key(metadata) -tag = make_ready_tag(metadata) -put_sample(key, fields, tag=tag) -``` - -`encode_sample()` 会将所有 Tensor detach、转到 CPU、整理为 contiguous,并统一 `input_ids/position_ids` 为 int64、`loss_mask` 为 float32。随后校验 token 长度、hidden shape 和 metadata 一致,再把 metadata JSON 编成 uint8 Tensor。 - -Consumer 的逆过程位于 `TQFeatureStore.get_many()`: - -```python -records = get_samples([entry.key for entry in entries]) -sample = decode_sample( - key=key, - tag=entry.tag, - fields=fields, - expected_config=self.expected_config, -) -``` - -`get_samples()` 最终调用一次 TQ `kv_batch_get(keys=[...], partition_id=...)`。bridge 将 TQ 返回的 batched TensorDict 或 mapping 拆成与请求 key 顺序一致的普通 fields 字典。 - -`decode_sample()` 随后: - -1. 检查必需 fields 是否存在。 -2. 将 `metadata_json` 从 uint8 Tensor 还原为字典和 `SampleMetadata`。 -3. 根据 metadata 重新计算 key,并与实际 key 比较。 -4. 比较 tag 与 metadata 的公共身份字段。 -5. 检查 Consumer 的 expected config。 -6. 将 Tensor detach 到 CPU,统一基础 dtype/shape。 -7. 检查所有主 Tensor 第一维等于 `feature_length`,hidden shape/dtype 与 metadata 一致。 -8. 构造 `DraftFeatureSample`。 - -输出不再是 TQ 专用对象,而是既有训练代码使用的: - -```python -DraftFeatureSample( - input_ids=..., - loss_mask=..., - position_ids=..., - hidden_states=..., - metadata=..., - ..., -) -``` - -这是能够复用原训练逻辑的关键边界:TQ 只负责上游存储和传输,`decode_sample()` 后的数据类型与磁盘 feature store 读取结果一致。 - -## 7. Consumer 从启动到结束的完整执行流程 - -### 阶段 1:launcher 检查配置并启动 torchrun - -执行者是 launcher 父进程。入口是 `verl_speco.draft_train_launcher.main()`。 - -当 override 中出现 `feature_store.type=tq`,`validate_tq_launch_config()` 会要求: - -- `training.transfer_queue.enable=true`; -- `training.transfer_queue.ray.address` 非空; -- `training.transfer_queue.run_id` 非空。 - -检查成功后构造: - -```text -python -m torch.distributed.run - --nnodes=... - --nproc_per_node=... - -m verl_speco.draft_train - <全部 Hydra overrides> -``` - -配置参数是普通子进程命令行参数。此阶段没有 TQ Tensor 传输。 - -### 阶段 2:每个 rank 初始化训练运行时 - -每个 torchrun 子进程进入 `run_standalone_draft_training()`,调用 `_init_distributed()` 得到 `rank/local_rank/world_size`,绑定本 rank GPU,然后构造原有 `DrafterBaseTrainer` 和 DSpark backend。 - -`speculative_algorithm=DSPARK` 决定 backend 和 DSpark 模型训练实现;`feature_store.type=tq` 只改变数据来源,不替换 trainer。 - -### 阶段 3:factory 创建 TQFeatureStore - -训练循环调用: - -```python -store = build_feature_store_from_config( - feature_store_cfg, - read_only=True, - transfer_queue_cfg=training_cfg.get("transfer_queue"), -) -``` - -factory 在 `type=tq` 时不读取 `feature_store.path`,而是把 sibling `training.transfer_queue` 交给 `TQFeatureStore.from_config()`。 - -TQ store 被限定为 `read_only=True`,意思是它是训练 Consumer source。这里的“read only”不表示永不修改 TQ;成功消费后仍可通过明确的 `clear_many()` 删除记录,但不会把它当作通用 feature writer。 - -### 阶段 4:所有 rank 分别连接同一个 TQ - -训练循环调用 `_connect_tq_store_across_ranks()`。每个 rank 都独立执行 `store.connect()`: - -```text -configure_transfer_queue -→ ray.init(address, namespace) -→ tq.init(same native config) 连接 named Controller -→ 本 rank 设置 _connected=True -``` - -之后 `_all_ranks_true()` 使用 `dist.all_reduce(MIN)` 汇总连接结果。只要一个 rank 连接失败,所有 rank 都停止,不允许部分 rank 进入后续 broadcast 或 FSDP collective。 - -这里没有“rank 0 建一个 client 给其他 rank 共用”。TQ client 是进程本地对象,N 个 rank 有 N 个 client,但它们指向同一 Controller/partition。 - -### 阶段 5:创建 TQFeatureDataLoader - -每个 rank 构造自己的 loader,参数包括相同的 `batch_size_per_gpu`、`world_size`、轮询间隔和 drop-last,以及不同的 `rank`。 - -假设: - -```text -world_size = 2 -batch_size_per_gpu = 2 -global_batch_size = 4 -``` - -那么只有 ready 数量至少为 4,rank 0 才发布一个 batch 命令。 - -### 阶段 6:rank 0 发现 ready key - -rank 0 首先检查 `owner_ready()`。Owner 尚未发布 readiness marker 时,rank 0 sleep 后继续轮询,不会让其他 rank 开始取数。 - -Owner ready 后,rank 0 调用 `list_ready()`,其底层是: - -```text -tq.kv_list(partition_id) -→ key -> tag -→ 按 record_type/status/run/schema/algorithm 过滤 -→ 按 (sequence_no, key) 排序 -``` - -此阶段没有读取 fields,因此 hidden states 尚未传到训练进程。 - -### 阶段 7:rank 0 切分 global batch - -若排序后的前四条是 `k0、k1、k2、k3`,`build_assignments()` 产生: - -```python -assignments = [ - [ReadyEntry(k0, tag0), ReadyEntry(k1, tag1)], # rank 0 - [ReadyEntry(k2, tag2), ReadyEntry(k3, tag3)], # rank 1 -] -``` - -每条样本只出现在一个 rank 的 assignment 中,因此各 rank 不会取得同一训练样本。这里采用连续、不重叠的切片。 - -rank 0 随后构造普通 Python 命令字典: - -```python -{ - "kind": "batch", - "global_keys": [k0, k1, k2, k3], - "assignments": [ - [{"key": k0, "tag": tag0}, {"key": k1, "tag": tag1}], - [{"key": k2, "tag": tag2}, {"key": k3, "tag": tag3}], - ], -} -``` - -`global_keys` 只用于 rank 0 在训练完成后一次清理整个 batch;`assignments` 用于每个 rank 知道自己应该 get 哪些 key。 - -### 阶段 8:小型命令通过 torch.distributed 广播 - -各 rank 同时进入: - -```python -dist.broadcast_object_list(payload, src=0) -``` - -rank 0 的 payload 中是上述字典,其他 rank 的初始值是 `None`。PyTorch 会序列化这个普通 Python 对象并广播给所有 rank。 - -这条边界只传输字符串、整数和小字典 tag。`hidden_states`、`input_ids` 等 fields 不在命令中,所以不会经 rank 0 中转,也不会随 broadcast 复制完整 global batch Tensor。 - -### 阶段 9:每个 rank 直接从 TQ 取本地 payload - -每个 rank 从 `assignments[self.rank]` 还原自己的 `ReadyEntry`: - -```python -local_entries = [_entry_from_wire(item) for item in assignments[self.rank]] -samples = self.store.get_many(local_entries) -``` - -在上述例子中: - -- rank 0 调用 `kv_batch_get(keys=[k0, k1], partition_id=...)`; -- rank 1 调用 `kv_batch_get(keys=[k2, k3], partition_id=...)`。 - -大 Tensor 的数据面因此是 TQ storage 到目标训练 rank,不经过训练 rank 0 的 Python 内存。每个 rank 得到两个 CPU `DraftFeatureSample`。 - -loader yield: - -```python -TQLocalBatch( - local_keys=[本 rank 的 key], - local_samples=[本 rank 的 DraftFeatureSample], - global_keys=[完整 global batch key] if rank == 0 else None, -) -``` - -非 rank 0 不保存 `global_keys`,避免多个 rank 都尝试 clear。 - -### 阶段 10:复用已有训练 batch 构造 - -训练循环识别 `TQLocalBatch` 后,只取: - -```python -samples = tq_local_batch.local_samples -``` - -然后调用原有接口: - -```python -batch = trainer.prepare_training_batch_from_samples( - materialized_samples, - step=optimizer_step, -) -``` - -TQ 路径禁止同时开启 `target_feature_pipeline`,因为样本已经包含目标模型 hidden states,不需要训练侧再访问 vLLM materialize 一次。 - -此时 TQ 专用的 key/tag 已不参与 DSpark 数学计算;训练接口看到的是普通 `DraftFeatureSample`,并按原逻辑整理 input ids、hidden states、mask、position ids 和 DSpark 训练所需输入。 - -### 阶段 11:所有 rank 同步 batch 是否可训练 - -每个 rank 判断 `batch is not None`,再通过 `_all_ranks_true()` 做 `all_reduce(MIN)`。 - -只有所有 rank 都成功构造 batch,才能进入训练。如果任一 rank 解码或 batch 构造失败,TQ 路径直接报错,而且这些 key 不会被删除。 - -### 阶段 12:执行原有 DSpark training step - -每个 rank 调用: - -```python -ok = await trainer.training_step_from_batch(batch, optimizer_step) -``` - -该调用复用既有模型 forward、DSpark loss(包括配置开启时的 L1 loss)、backward、梯度同步和 optimizer step。TQ 新代码没有重新实现 loss 或 optimizer。 - -之后再次以 `_all_ranks_true(ok)` 同步。只有所有 rank 都返回成功,才认为这一个 global batch 已经安全消费。 - -### 阶段 13:训练成功后由 rank 0 删除 global batch - -训练循环调用 `_clear_tq_batch_across_ranks()`: - -1. rank 0 使用 `tq_local_batch.global_keys` 调用 `loader.clear_completed_batch()`。 -2. loader 调用 `store.clear_many(global_keys)`。 -3. bridge 最终调用 `tq.kv_clear(keys=[k0,k1,k2,k3], partition_id=...)`。 -4. 所有 rank 通过 `all_reduce(MAX)` 同步 clear 是否失败。 - -删除发生在 optimizer step 全 rank 成功之后。不是“某个 rank get 完就删除”,因为 get 完只代表 Tensor 已读取,不能代表训练 step 已成功。 - -clear 成功后才增加 `successful_steps`,然后复用原有 metrics 和 checkpoint 调度。 - -### 阶段 14:EOS 和尾 batch - -当 ready 样本少于一个 global batch时,rank 0 查询 EOS: - -- 没有 EOS:说明 Producer 以后仍可能写入更多样本,sleep 后继续轮询。 -- 已有 EOS 且 ready 为空:广播 `{"kind": "stop"}`,所有 rank 结束迭代。 -- 已有 EOS 且存在不足一个 global batch 的尾样本:rank 0 先 clear 这些尾 key,再广播 stop。 - -第一版强制 `drop_last=true`,因此不会构造各 rank batch size 不一致的最后一步。 - -### 阶段 15:checkpoint 和退出清理 - -正常 step 完成后仍按原 `save_interval_steps` 保存 checkpoint;循环结束后按 `save_final_checkpoint` 决定是否保存最终 checkpoint。 - -`finally` 中每个 rank 调用 `store.close()`。对 `TQFeatureStore` 而言,这只是: - -```text -关闭本进程 TQ client -→ 如果本进程自行 ray.init,则 ray.shutdown() -``` - -它不会调用全局 `tq.close()`,不会杀死 Owner 创建的 Controller,也不会影响仍在运行的 Producer 或其他 rank。 - -## 8. 控制面和数据面的完整边界 - -| 数据 | 从哪里到哪里 | 传输机制 | 是否经过 rank 0 | -|---|---|---|---| -| 启动配置 | launcher 到 torchrun 子进程 | 命令行 Hydra overrides | 每个 rank 都收到 | -| ready key/tag | TQ Controller 到 rank 0 | `tq.kv_list` | 是,只有 rank 0 list | -| batch assignment | rank 0 到全部 rank | `dist.broadcast_object_list` | 由 rank 0 发出 | -| hidden states 等 fields | TQ storage 到被分配的 rank | `tq.kv_batch_get` | rank 1 的 Tensor 不经过 rank 0 | -| batch 准备/训练成功状态 | 全部 rank 之间 | Tensor `all_reduce` | collective,无单点 payload relay | -| clear 请求 | rank 0 到 TQ | `tq.kv_clear(global_keys)` | 只有 rank 0 发起 | -| 梯度和模型 collective | 训练 rank 之间 | 既有 PyTorch distributed/FSDP 路径 | 与 TQ 无关 | - -## 9. 当前“最简单校验”具体简单在哪里 - -`TQFeatureStore` 构造的 expected config 只固定: - -```python -ExpectedFeatureConfig( - run_id=<当前训练 run_id>, - schema_version=<当前 schema>, -) -``` - -TQ 会保留 Producer 写入的 `SampleMetadata.algorithm`,但不使用它选择训练 backend, -也不额外与启动配置比较。与原有离线 feature-store 训练一致,实际 trainer/backend 只由 -`rollout.drafter.speculative_algorithm` 和既有 backend factory 决定。 - -因此当前不会拿 Consumer 配置额外比较: - -- target model ID/revision; -- tokenizer fingerprint; -- target layer IDs; -- hidden layout; -- hidden dtype 的外部预期值。 - -但这不等于完全不校验。`decode_sample()` 仍然强制检查: - -- 必需 fields 存在; -- key、tag、metadata 三者身份一致; -- schema/run/algorithm 符合 Consumer; -- Tensor 类型正确; -- input/mask/position/hidden 长度一致; -- hidden 实际 shape/dtype 与该样本 metadata 一致; -- feature window 合法。 - -这满足“第一版少做外部模型身份检查”,同时避免把结构损坏或错 run 的数据送入训练。 - -## 10. 失败、删除和重复消费语义 - -当前实现遵循以下规则: - -1. 连接失败:所有 rank 同步停止。 -2. rank 0 list/EOS 失败:rank 0 广播 error 命令,其他 rank 不会永久等待 batch broadcast。 -3. 某 rank get/decode 失败:`_next_batch_across_ranks()` 将失败同步给全部 rank,不进入模型训练 collective。 -4. 某 rank 无法构造 batch:报错,global keys 保留在 TQ。 -5. 某 rank training step 失败:报错,global keys 保留在 TQ。 -6. 全 rank training step 成功:rank 0 clear 整个 global batch。 -7. clear 失败:错误传播到全部 rank,训练停止;不会把该 step 继续当成已正常完成。 -8. 达到 `max_steps`:循环停止;尚未选择的 ready 样本保留在 TQ。 - -第一版尚未实现完整的崩溃恢复协议。尤其是“optimizer step 已成功,但进程在 clear 前崩溃”时,key 仍存在;重新启动 Consumer 可能再次读取它。要实现严格 exactly-once,需要把 checkpoint step、已消费 sequence 或事务状态纳入协议。该能力应作为后续增强,而不是当前已实现能力。 - -## 11. 如何启动和检查 - -正式运行统一使用端到端 launcher;它负责 Ray、TQ Owner、Producer 和 Consumer 的启动与清理: - -```bash -bash examples/run_qwen3-8b_drafter_separate_training.sh -``` - -## 12. 已完成的测试 - -### 12.1 Consumer/factory/协议/launcher 单元测试 - -已执行: - -```text -python -m pytest \ - tests/unit/test_tq_consumer.py \ - tests/unit/test_draft_train_launcher.py \ - tests/unit/test_transferqueue_bridge.py \ - tests/unit/test_drafter_sample_protocol.py \ - tests/unit/test_draft_feature_store.py \ - -q -``` - -结果:`44 passed`。 - -### 12.2 真实 TQ 0.1.7 跨进程 smoke - -早期临时跨进程工具验证过以下行为,当前回归由协议、bridge、Consumer和launcher单元测试承担: - -- Owner 发布 owner-ready; -- 两条 sample 写入 TQ; -- Consumer 经 `TQFeatureStore` 和 `TQFeatureDataLoader` 读到两条样本; -- hidden shape 正确; -- Consumer clear 已完成 batch; -- EOS 后迭代停止; -- Consumer 只关闭本地 client,Owner 仍能继续观察完成标记并正常关闭。 - -实际 smoke 输出包含: - -```text -CLIENT_OK samples=2 shape=(3,4) -CLIENT_CLOSED_LOCAL_ONLY -OWNER_OBSERVED_SAMPLES_CLEARED -OWNER_CLOSED -``` - -### 12.3 当前环境未覆盖的部分 - -完整 `tests/unit/test_draft_training_loop.py` 在当前 Windows 环境无法完整收集,因为上游 `verl/ray` 依赖不齐;新增训练循环测试代码已通过 Python 编译检查,连接/clear helper 也通过针对性单元逻辑验证。真实多 GPU DSpark 训练仍需要在目标 Linux GPU 环境执行集成测试。 - -## 13. 当前限制和后续建议 - -当前第一版有意不实现以下复杂能力: - -1. Producer 本身尚未在本次 Consumer 改动中实现。 -2. TQ 公共协议和 Consumer 已不再写死 `DSPARK`;当前测试 Producer、启动脚本和已验证的 - feature 语义仍是 DSPARK。其他算法若能复用当前公共 dense fields,只需由对应 Producer - 生成正确的 `DraftFeatureSample`;若字段结构不同,则在协议模块增加对应 codec,不需要改 - TQ 的 key/tag 发现、rank 分配和 clear 流程。 -3. 只支持 `drop_last=true`。 -4. 不支持 TQ 与 `target_feature_pipeline.enabled=true` 同时开启。 -5. 不提供严格的 crash exactly-once 或 checkpoint/queue 联合恢复。 -6. rank 0 仍通过 `kv_list` 轮询整个 partition;数据量很大时可考虑 cursor/ready queue 优化。 -7. 当前外部 expected config 校验较简化,后续可把 model revision、tokenizer fingerprint、layer/layout/dtype 预期接入 Hydra 配置。 -8. 尚需在真实多机、多 GPU、Mooncake backend 环境验证吞吐、背压、Owner 生命周期和网络故障行为。 - -建议下一阶段优先完成 Producer,并严格复用 `drafter_sample_protocol.py`,不要在 Producer 另造一套 key/tag/fields 格式。完成 Producer 后,首先跑 world size 1 的端到端训练,再跑多 rank 验证每条 key 只分配给一个 rank、训练成功后只由 rank 0 clear。 - -## 14. 最终路径摘要 - -```text -Producer(待实现) - vLLM 并行 prefill - → DraftFeatureSample + SampleMetadata - → encode_sample 得到 Tensor fields - → TQ kv_put(key, fields, tag) - -Consumer rank 0 - kv_list 只取 key/tag - → 过滤并按 sequence_no 排序 - → 切出 global batch - → broadcast 每个 rank 的 key/tag assignment - -每个 Consumer rank - 取 assignments[rank] - → kv_batch_get 本 rank keys - → decode_sample 得到 DraftFeatureSample - → 原 prepare_training_batch_from_samples - → 原 DSpark training_step_from_batch - -全部 rank - 同步确认 optimizer step 成功 - → rank 0 kv_clear(global_keys) - → 原 metrics/checkpoint - → 下一批 - -Producer 发布 EOS - → rank 0 确认没有完整 global batch - → 清理不足一批的尾样本 - → broadcast stop - → 各 rank 关闭本地 TQ client 并退出 -``` diff --git a/docs/standalone_tq_foundation_implementation.md b/docs/standalone_tq_foundation_implementation.md deleted file mode 100644 index b3e0dea3..00000000 --- a/docs/standalone_tq_foundation_implementation.md +++ /dev/null @@ -1,1030 +0,0 @@ -# Standalone TQ 公共基础层实现说明 - -Last updated: 08/21/2026 - -## 1. 文档范围和已验证结论 - -本文解释当前仓库已经实现并测试通过的 TQ 公共基础层: - -```text -verl_speco/transport/drafter_sample_protocol.py -verl_speco/integration/transferqueue_bridge.py -verl_speco/config/speco_base.yaml -verl_speco/tq_owner.py -tests/unit/test_drafter_sample_protocol.py -tests/unit/test_transferqueue_bridge.py -pyproject.toml -``` - -当前已经实现: - -1. Producer 和 Consumer 共用的样本 key、tag、fields 和 metadata 协议; -2. 普通进程连接 Ray 集群; -3. TQ Owner 创建 named `TransferQueueController`; -4. 独立 Client 发现并连接同一个 Controller; -5. 单样本 put、元数据 list、批量 get 和批量 clear; -6. Owner 与 Client 不同的关闭边界; -7. 独立 Owner 入口和共享 Hydra 配置; -8. mock 单元测试和真实双进程 TQ 0.1.7 smoke test。 - -当前还没有实现: - -1. Producer 读取 JSONL、并发访问多个 vLLM endpoint 的完整 pipeline; -2. `feature_store.type=tq` 工厂分支; -3. `TQFeatureStore` 和 `TQFeatureDataLoader`; -4. rank 0 选择 global keys、各 rank 读取 local keys; -5. TQ batch 接入 DSpark optimizer step; -6. optimizer step 成功后的 rank 0 clear。 - -因此当前代码已经证明“两个独立进程能通过 Ray 连接同一个 TQ,并按共享协议批量传输多个样本”,但尚未接到正式 Producer 和 DSpark Consumer 主循环。 - -## 2. 运行时角色和术语 - -### 2.1 Ray head - -Ray head 是 Ray 集群的控制节点。TQ 0.1.7 使用 Ray 管理 named Controller 和 storage actors。 - -Ray 不传输 standalone hidden-state payload。当前代码没有对这些样本调用: - -```python -ray.put(hidden_states) -``` - -### 2.2 TQ Owner - -TQ Owner 是普通 Python OS 进程,入口为: - -```text -python -m verl_speco.tq_owner -``` - -它不是 Ray actor。它先调用 `ray.init(address=...)` 加入 Ray,再调用带完整配置的 `tq.init(config)` 创建 TQ Controller 和 storage。 - -Owner 是唯一允许调用全局 `tq.close()` 的进程。 - -### 2.3 Named TransferQueueController - -TQ 0.1.7 内部创建: - -```python -TransferQueueController.options( - name="TransferQueueController" -).remote(...) -``` - -`name` 将 Controller actor 注册到 Ray actor registry。其他进程加入同一 Ray address 和 namespace 后,通过: - -```python -ray.get_actor("TransferQueueController") -``` - -取得 actor handle,再读取 TQ backend 配置。 - -Controller 保存控制信息和数据位置;使用 MooncakeStore 时,大 tensor 本身保存在 MooncakeStore。 - -### 2.4 TQ Client - -Owner、Producer、每个 Consumer rank 都在各自 OS 进程内拥有独立 TQ Client。 - -普通 Client 也传入相同的 native 配置: - -```python -tq.init(native_config) -``` - -发现 named Controller并初始化本进程的 storage manager。Client 不是 Ray actor,Producer 和 torchrun rank 也不需要改成 Ray actor。 - -### 2.5 Partition、key、tag 和 fields - -当前固定 partition: - -```text -speco_drafter_features -``` - -TQ 中一条记录逻辑上是: - -```text -partition_id -└── key - ├── tag:轻量 dict,由 kv_list 发现 - └── fields:Tensor payload,由 kv_batch_get 读取 -``` - -## 3. 共享配置如何工作 - -共享配置定义在 `verl_speco/config/speco_base.yaml`: - -```yaml -transfer_queue: - enable: false - package_version: "0.1.7" - ray: - address: null - namespace: speco-drafter - partition_id: speco_drafter_features - run_id: null - schema_version: 1 - connect_timeout_seconds: 120 - poll_interval_seconds: 0.5 - drop_last: true - controller: - polling_mode: true - backend: - storage_backend: SimpleStorage - SimpleStorage: - total_storage_size: 100000 - num_data_storage_units: 8 - MooncakeStore: - auto_init: false - metadata_server: localhost:50050 - master_server_address: localhost:50051 - local_hostname: "" - protocol: tcp - global_segment_size: 4294967296 - local_buffer_size: 1073741824 - device_name: "" -``` - -### 3.1 Ray 连接字段 - -```yaml -ray: - address: 10.0.0.1:6379 - namespace: speco-drafter -``` - -它们决定当前进程连接哪个 Ray 集群,以及在哪个 namespace 查找 `TransferQueueController`。Owner、Producer 和所有 Consumer ranks 必须使用相同值。 - -### 3.2 SPECO 协议字段 - -```yaml -partition_id: speco_drafter_features -run_id: dspark-20260819-a -schema_version: 1 -``` - -这些字段用于路由和校验,不属于 TQ 原生配置。`run_id` 用于隔离不同 pipeline 的记录。 - -### 3.3 TQ 原生字段 - -```yaml -controller: ... -backend: ... -``` - -只有这些字段应传给 `tq.init(full_config)`。bridge 的 `_native_tq_config()` 会删除: - -```text -enable -package_version -ray -partition_id -run_id -schema_version -connect_timeout_seconds -poll_interval_seconds -drop_last -``` - -对象变化为: - -```text -完整 SPECO transfer_queue dict -→ _native_tq_config() -→ controller/backend等TQ字段 -→ OmegaConf DictConfig -→ tq.init(same native config) -``` - -## 4. Bridge 的进程内状态 - -`verl_speco/integration/transferqueue_bridge.py` 在每个 OS 进程内分别维护: - -```python -_state = { - "enabled": False, - "configured": False, - "initialized": False, - "config": None, - "owner": False, - "ray_initialized_here": False, - "ray_address": None, - "ray_namespace": None, -} -``` - -该 dict 不跨进程共享。Owner、Producer、每个 rank 分别拥有自己的 `_state`。 - -| 字段 | 含义 | -|---|---| -| `enabled` | 当前进程配置是否开启 TQ | -| `configured` | 是否调用过 `configure_transfer_queue()` | -| `initialized` | 当前进程是否执行过 `tq.init()` | -| `config` | 当前进程保存的普通 dict 配置 | -| `owner` | 当前进程是否创建了全局 Controller/Storage | -| `ray_initialized_here` | bridge 是否负责调用了本进程的 `ray.init()` | -| `ray_address/namespace` | 本进程的 Ray 连接信息 | - -`_state_lock` 只保护同一进程内多个线程同时初始化,不是分布式锁。 - -## 5. Owner 的完整启动数据流 - -Owner 入口是 `verl_speco/tq_owner.py`,直接加载 Hydra 主配置 -`verl_speco/config/speco_base.yaml`。`run_owner()` 会复制其中的 -`transfer_queue` 子配置,并仅在 Owner 进程的副本中强制设置 -`enable=true`;普通训练任务看到的共享默认值仍是 `false`。 - -### 阶段 1:读取配置 - -执行者:Owner OS 进程。 - -入口: - -```python -run_owner(config) -``` - -取得: - -```python -training_cfg = config.actor_rollout_ref.rollout.drafter.training -tq_cfg = training_cfg.transfer_queue -``` - -然后调用: - -```python -configure_transfer_queue(training_cfg) -``` - -该函数只把 OmegaConf 转成普通 dict并更新当前进程 `_state`,不会连接 Ray,也不会创建 TQ。 - -### 阶段 2:连接 Ray - -Owner 调用: - -```python -connect_ray_cluster(ray_address, namespace) -``` - -内部执行: - -```python -if not ray.is_initialized(): - ray.init(address=ray_address, namespace=namespace) -``` - -边界类型是 Ray control-plane connection。此时还没有传输训练 tensor。 - -### 阶段 3:创建 Controller 和 Storage - -Owner 调用: - -```python -start_transfer_queue_owner(tq_cfg) -``` - -执行顺序: - -1. `_extract_tq_config()` 得到普通 dict; -2. 检查 `enable=true`; -3. 检查 `TransferQueue` 包可用; -4. 防止本进程重复初始化; -5. `_native_tq_config()` 删除 SPECO 字段; -6. `_as_tq_config()` 转 OmegaConf; -7. `tq.init(native_config)` 创建 Controller、Storage 和 Owner Client; -8. 设置 `_state.owner=True`、`initialized=True`。 - -Ray 中形成: - -```text -Ray cluster / namespace -├── named actor: TransferQueueController -└── storage backend - ├── SimpleStorage actors - └── 或 MooncakeStore connection/process -``` - -### 阶段 4:发布 owner-ready - -调用: - -```python -publish_owner_ready(run_id, schema_version) -``` - -生成: - -```python -key = "control:v1::owner-ready" -fields = {"marker": torch.tensor([1], dtype=torch.uint8)} -tag = { - "record_type": "control", - "status": "owner_ready", - "schema_version": 1, - "run_id": run_id, -} -``` - -这是一条控制记录,不进入训练 batch。 - -### 阶段 5:常驻和关闭 - -Owner 安装 `SIGINT/SIGTERM` handler,并等待: - -```python -stop_event.wait() -``` - -收到信号后调用: - -```python -close_transfer_queue_owner() -``` - -执行: - -```text -验证owner身份 -→ tq.close() -→ 清理Controller/Storage -→ ray.shutdown() -``` - -Owner 必须在 Producer 和 Consumer 退出后才能关闭。 - -## 6. 普通 Client 如何连接同一个 TQ - -Producer 和每个 Consumer rank 后续使用相同顺序: - -```python -configure_transfer_queue(training_cfg) -connect_ray_cluster(ray_address, namespace) -connect_transfer_queue_client() -``` - -`connect_transfer_queue_client()` 最终调用: - -```python -tq.init(same_native_config) -``` - -TQ 0.1.7 内部通过: - -```python -ray.get_actor("TransferQueueController") -``` - -找到 Owner 创建的 Controller,读取 backend 配置,然后创建当前进程的 TQ Client。 - -对象和边界变化: - -```text -actor名称字符串 -→ Ray actor registry -→ Controller actor handle -→ Controller.get_config.remote() -→ TQ DictConfig -→ 当前进程TransferQueueClient -→ 同一个SimpleStorage/MooncakeStore -``` - -## 7. 一条具体样本的初始对象 - -真实 smoke test使用: - -```python -sample = DraftFeatureSample( - algorithm="DSPARK", - input_ids=torch.tensor([1, 2, 3]), # int64[3], CPU - loss_mask=torch.tensor([0.0, 1.0, 1.0]), # float32[3], CPU - position_ids=torch.tensor([0, 1, 2]), # int64[3], CPU - hidden_states=torch.arange( - 12, dtype=torch.float32 - ).reshape(3, 4), # float32[3,4], CPU -) -``` - -同时构造: - -```python -meta = SampleMetadata( - schema_version=1, - run_id="codex-batch-smoke", - sample_id="smoke-0000", - sequence_no=0, - algorithm="DSPARK", - target_model_id="smoke-target", - target_model_revision="smoke-revision", - tokenizer_fingerprint="smoke-tokenizer", - target_layer_ids=[0], - hidden_states_layout="dflash_aux", - hidden_dtype="float32", - hidden_shape=[3, 4], - feature_length=3, - full_sequence_length=3, - feature_start=0, - feature_end=3, - use_logits=False, -) -``` - -`DraftFeatureSample` 和 `SampleMetadata` 都是进程内 Python 对象,不直接经过 TQ。 - -## 8. Key 的生成和两个同名函数 - -共享协议调用: - -```python -make_sample_key(meta) -``` - -输出: - -```text -drafter:v1:codex-batch-smoke:000000000000:smoke-0000 -``` - -字段顺序: - -```text -drafter / schema version / run_id / 12位sequence_no / sample_id -``` - -`sequence_no` 是输入文件 record 序号,不是训练 batch 编号。 - -bridge 为兼容 PR #48 还保留另一个: - -```python -transferqueue_bridge.make_sample_key( - global_step, - replica_rank, - request_id, -) -``` - -它生成: - -```text -speco::: -``` - -standalone Producer 必须从 `verl_speco.transport.drafter_sample_protocol` import `make_sample_key`,不能使用 bridge 中的 PR #48 旧函数。 - -## 9. Tag 如何生成 - -```python -tag = make_ready_tag(meta) -``` - -输出: - -```python -{ - "record_type": "sample", - "status": "ready", - "schema_version": 1, - "run_id": "codex-batch-smoke", - "sequence_no": 0, - "sample_id": "smoke-0000", - "algorithm": "DSPARK", -} -``` - -tag 只包含发现、过滤和排序需要的标量/字符串。`list_samples()` 只返回 key/tag,不读取 hidden states。 - -## 10. `encode_sample()` 如何生成 fields - -调用: - -```python -fields = encode_sample(sample, meta) -``` - -### 10.1 校验 - -执行: - -```text -SampleMetadata.validate() -DraftFeatureSample.validate(strict=True) -``` - -随后检查: - -1. hidden states 是一个 dense tensor; -2. ids/mask/position 长度等于 `feature_length`; -3. hidden 第一维等于 `feature_length`; -4. hidden shape 等于 metadata; -5. hidden dtype 等于 metadata; -6. feature window 长度正确。 - -### 10.2 Tensor 规范化 - -```text -input_ids → CPU contiguous int64[L] -loss_mask → CPU contiguous float32[L] -position_ids → CPU contiguous int64[L] -hidden_states → CPU contiguous,保持模型dtype -``` - -没有 `position_ids` 时生成 `torch.arange(L, dtype=int64)`。 - -### 10.3 Metadata JSON 编码 - -```text -SampleMetadata dataclass -→ dict -→ JSON UTF-8 bytes -→ torch.uint8[M] -``` - -实现等价于: - -```python -raw = json.dumps(metadata).encode("utf-8") -metadata_json = torch.tensor(list(raw), dtype=torch.uint8) -``` - -### 10.4 最终 fields - -```python -fields = { - "input_ids": int64[3], - "loss_mask": float32[3], - "position_ids": int64[3], - "hidden_states": float32[3,4], - "metadata_json": uint8[M], -} -``` - -如果存在,协议也保留 `last_hidden_states`、`target` 和 `target_logprobs`。 - -## 11. Bridge 如何写入 TQ - -调用: - -```python -put_sample(key, fields, tag=tag) -``` - -bridge 执行: - -1. 检查 TQ 已启用; -2. 丢弃 fields 中非 tensor 值; -3. 确保本进程已经使用相同 native 配置执行 `tq.init(config)`; -4. 取得配置中的 partition; -5. 调用: - -```python -tq.kv_put( - key=key, - partition_id="speco_drafter_features", - fields=fields, - tag=tag, -) -``` - -使用 MooncakeStore 时,大 tensor 路径是: - -```text -Producer CPU tensor -→ Producer TQ Client -→ MooncakeStore -``` - -不是 Ray `ObjectRef`。 - -## 12. Consumer 如何发现 key - -调用: - -```python -records = list_samples() -``` - -内部调用: - -```python -tq.kv_list(partition_id="speco_drafter_features") -``` - -标准化返回类型: - -```python -dict[str, dict[str, Any]] -``` - -示例: - -```python -{ - "drafter:v1:...:smoke-0000": { - "record_type": "sample", - "status": "ready", - "run_id": "codex-batch-smoke", - "sequence_no": 0, - ... - } -} -``` - -bridge 兼容 `key → tag` 和 `partition → key → tag` 两种 wrapper。这一步不读取 fields。 - -## 13. Consumer 如何批量取样本 - -输入: - -```python -keys = [key0, key1] -``` - -调用: - -```python -records = get_samples(keys) -``` - -bridge 只调用一次: - -```python -result = tq.kv_batch_get( - keys=keys, - partition_id="speco_drafter_features", -) -``` - -TQ 0.1.7 返回带 batch 维的 TensorDict。bridge 检查 `result.batch_size`,再执行: - -```python -rows = [result[index] for index in range(len(keys))] -``` - -每行转成普通 dict,最终返回: - -```python -[ - (key0, fields0), - (key1, fields1), -] -``` - -返回顺序与输入 keys 一致。重复 key 会提前报错。 - -## 14. `decode_sample()` 如何恢复训练对象 - -调用: - -```python -sample = decode_sample( - key, - tag, - fields, - expected_config, -) -``` - -### 14.1 Metadata 解码 - -```text -metadata_json uint8[M] -→ bytes -→ UTF-8 -→ json.loads -→ dict -→ SampleMetadata.from_dict -``` - -### 14.2 身份一致性 - -代码根据 metadata 重新生成 key,要求输入 key 完全相等;然后逐项校验 tag 的: - -```text -record_type/status/schema_version/run_id/sequence_no/sample_id/algorithm -``` - -所以 key、tag 和 payload metadata 不能来自不同样本。 - -### 14.3 Consumer 合同 - -Consumer 提供: - -```python -ExpectedFeatureConfig( - run_id="codex-batch-smoke", - schema_version=1, - algorithm="DSPARK", - target_model_id="smoke-target", - target_model_revision=None, - tokenizer_fingerprint=None, - target_layer_ids=None, - hidden_states_layout="dflash_aux", - hidden_dtype="float32", -) -``` - -值为 `None` 的字段不检查;其他字段必须完全一致。正式 Consumer 应填写 target checkpoint、tokenizer、layers、layout 和 dtype,避免使用错误 target 特征。 - -### 14.4 输出 - -完成 tensor 类型、长度、shape、dtype 校验后,构造: - -```python -DraftFeatureSample.from_dict(payload, strict=True) -``` - -输出可以交给现有 `trainer.prepare_training_batch_from_samples()`。当前 smoke test验证到这里,正式 `TQFeatureDataLoader` 尚未实现。 - -## 15. EOS 控制记录 - -调用: - -```python -key, fields, tag = make_eos_record(run_id, total_samples) -``` - -输出: - -```python -key = "control:v1::eos" -fields = {"marker": torch.tensor([1], dtype=torch.uint8)} -tag = { - "record_type": "control", - "status": "eos", - "schema_version": 1, - "run_id": run_id, - "total_samples": total_samples, -} -``` - -EOS 表示不会再发布新样本。Consumer 应先 drain ready samples 再退出。协议函数已实现,正式 Producer/Consumer 尚未调用。 - -## 16. Clear 和数据生命周期 - -bridge 提供: - -```python -clear_samples(keys) -``` - -内部调用: - -```python -tq.kv_clear(keys=keys, partition_id="speco_drafter_features") -``` - -基础层只执行删除,不决定删除时机。正式 Consumer 必须遵守: - -```text -rank 0选择global keys -→ 各rank读取local keys -→ 所有rank完成同一optimizer step -→ 汇总global success -→ rank 0 clear global keys -``` - -不能在 `get_samples()` 后立即 clear,因为 get 成功不代表训练 step 成功。 - -## 17. Client close 和 Owner close - -### 17.1 Client close - -Producer/rank 调用: - -```python -close_transfer_queue_client() -``` - -执行: - -```text -tq.get_client() -→ 当前进程client.close() -→ 如果bridge负责ray.init,则ray.shutdown() -``` - -它不调用全局 `tq.close()`,不会 kill Controller。调用后本进程不能继续使用 TQ。 - -### 17.2 Owner close - -Owner 调用: - -```python -close_transfer_queue_owner() -``` - -执行: - -```text -tq.close() -→ Controller/Storage全局清理 -→ ray.shutdown() -``` - -Owner 若误用 Client close,bridge 会抛 `RuntimeError`。 - -## 18. PR #48 兼容边界 - -bridge 继续保留: - -```python -init_transfer_queue(config) -get_sample(key) -close_transfer_queue() -``` - -PR #48 的 `SpecoTaskRunner` 已是 Ray actor,所以不调用 standalone 的 `connect_ray_cluster()`。旧 worker 继续使用单 key `get_sample()`。 - -standalone 后续使用新增的: - -```python -list_samples() -get_samples(keys) -clear_samples(keys) -``` - -因此没有修改 PR #48 现有调用点的函数签名。 - -## 19. 依赖和命令入口 - -`pyproject.toml` 新增: - -```toml -[project.optional-dependencies] -transfer-queue = ["TransferQueue==0.1.7"] -``` - -安装: - -```bash -pip install -e ".[transfer-queue]" -``` - -TQ 0.1.7 会安装 Ray,并要求 `numpy<2.0.0`。若现有包要求 NumPy 2,需要隔离环境或重新解决依赖。 - -Owner 命令: - -```text -verl-speco-tq-owner -``` - -## 20. 单元测试 - -协议测试 `tests/unit/test_drafter_sample_protocol.py` 覆盖: - -1. encode/decode round trip; -2. key 格式; -3. tag 身份冲突; -4. Consumer contract 冲突; -5. hidden shape 冲突; -6. EOS 格式。 - -bridge 测试 `tests/unit/test_transferqueue_bridge.py` 覆盖: - -1. Ray address/namespace 参数; -2. Owner 只向 TQ 传原生配置; -3. Client 使用相同 native 配置调用 `tq.init(config)`; -4. put/list/get-many/clear; -5. batch 返回顺序; -6. Client close 不调用全局 close; -7. Owner 不能误用 Client close。 - -运行: - -```bash -python -m pytest \ - tests/unit/test_drafter_sample_protocol.py \ - tests/unit/test_transferqueue_bridge.py \ - -q -``` - -## 21. 真实双进程 smoke test - -早期用于该验证的临时双进程工具已经移除;正式入口统一由 -`verl_speco.standalone_tq_training_launcher` 管理 Owner、Producer 和 Consumer 生命周期。 - -Owner 路径: - -```text -连接Ray -→ tq.init(full config) -→ 写sample 0和sample 1 -→ 等待client-done -→ clear done marker -→ 全局关闭 -``` - -Client 路径: - -```text -连接同一个Ray -→ tq.init(same native config) -→ kv_list发现两个key -→ 一次kv_batch_get([k0,k1]) -→ 拆成两个fields dict -→ 分别decode_sample -→ clear两个sample keys -→ 写client-done -→ 只关闭本地client -``` - -已验证输出: - -```text -OWNER_READY keys=[k0, k1] -CLIENT_OK samples=2 shape=(3, 4) -CLIENT_CLOSED_LOCAL_ONLY -OWNER_OBSERVED_SAMPLES_CLEARED -OWNER_CLOSED -``` - -这证明: - -1. 两个普通进程能连接同一个 TQ; -2. named Controller 发现有效; -3. 0.1.7 的 `kv_list/kv_batch_get/kv_clear` 参数有效; -4. TensorDict batch 能按 key 顺序拆开; -5. 共享协议能恢复 `DraftFeatureSample`; -6. Client close 不会杀掉 Owner; -7. Owner 能最终统一关闭。 - -## 22. 当前完整路径总结 - -```text -Owner -→ ray.init(address, namespace) -→ tq.init(native config) -→ named TransferQueueController - -普通Client -→ ray.init(same address, same namespace) -→ tq.init(same native config) -→ 找到同一个Controller - -DraftFeatureSample + SampleMetadata -→ make_sample_key -→ make_ready_tag -→ encode_sample -→ fields + metadata_json tensor -→ bridge.put_sample -→ tq.kv_put -→ SimpleStorage/MooncakeStore - -Consumer/测试Client -→ bridge.list_samples -→ key + tag -→ bridge.get_samples(keys) -→ tq.kv_batch_get -→ TensorDict batch -→ 每个key对应一个fields dict -→ decode_sample -→ DraftFeatureSample - -正式训练成功后(待实现) -→ bridge.clear_samples(global_keys) - -Client退出 -→ close_transfer_queue_client - -所有业务进程退出 -→ Owner close_transfer_queue_owner -→ tq.close -→ ray.shutdown -``` - -## 23. 下一阶段接入约束 - -后续代码不能重新定义协议或直接访问 TQ 私有对象。 - -Producer 应复用: - -```text -SampleMetadata -make_sample_key -make_ready_tag -encode_sample -bridge.put_sample -make_eos_record -``` - -Consumer 应复用: - -```text -bridge.list_samples -bridge.get_samples -decode_sample -bridge.clear_samples -``` - -下一阶段需要新增: - -```text -verl_speco/trainer/tq_feature_store.py -verl_speco/trainer/tq_sample_source.py -feature_store.py 的 type=tq 分支 -draft_training_loop.py 的流式训练分支 -Producer入口、输入读取和并发vLLM文件 -``` - -这些文件应建立在本文已经实现和真实验证过的连接、协议、KV 和关闭接口之上。 diff --git a/docs/standalone_tq_producer.md b/docs/standalone_tq_producer.md deleted file mode 100644 index bb766907..00000000 --- a/docs/standalone_tq_producer.md +++ /dev/null @@ -1,212 +0,0 @@ -# Standalone vLLM → TransferQueue Producer - -Last updated: 08/21/2026 - -本文解释 standalone Producer:它可以直接读取 verl 的 prompt-only Parquet(包括 -DAPO-Math-17k 的 chat-message `prompt`),也兼容已有 `prompt`/`response` 的 JSONL -或 Parquet。缺少 response 时由 target vLLM 生成,并在同一请求中提取 prompt 与 -output hidden states,之后把样本写到已存在的 TransferQueue(TQ)。 - -这条路径面向第一版 DSpark standalone 训练:Producer、TQ owner 和 Consumer -是三个独立 OS 进程;Ray 只用于让它们找到同一个 TQ Controller,hidden states -不通过 Ray object store 传输。 - -## 为什么需要这个 Producer - -此前仓库已经有两块基础能力: - -- `drafter_sample_protocol.py`:规定一条 TQ sample 的 key、tag、Tensor 字段和 - EOS record; -- `transferqueue_bridge.py` 与 `tq_owner.py`:负责连接 Ray/TQ、写读清理样本和 - owner 生命周期。 - -缺少的是把预先生成的文本变成 DSpark 训练特征并发布到 TQ 的独立进程。新增的 -Producer 补上这一段,不引入第二套协议或 feature store。 - -## 数据流 - -```text -verl prompt Parquet 或 prompt/response JSONL/Parquet - │ - │ 按文件顺序分配 sequence_no 和 sample_id - ▼ -Tokenizer - │ input_ids / loss_mask / feature window - ▼ -多个 vLLM endpoint(有界并发) - │ OpenAI completions 请求 → 临时 safetensors 文件 - ▼ -公共 hidden-state 转换函数 - │ DSpark DraftFeatureSample + SampleMetadata - ▼ -TransferQueue kv_put(一条输入记录对应一条 sample) - │ - ├─ put 成功:删除该请求的临时文件 - └─ 全部成功:写一个 EOS control record -``` - -Producer 在开始请求前会等到对应 `run_id` 的 `owner_ready` 控制记录。它不会创建 -Ray head、TQ Controller 或 storage backend;这些由 `verl-speco-tq-owner` 管理。 - -## 输入文件 - -输入可以是 JSONL 或 Parquet。`prompt` 可以是字符串,也可以是 verl 常用的 -`[{"role": ..., "content": ...}]` chat-message 列表。`response` 是可选字符串: -存在时直接 replay;不存在时由 target vLLM 生成。Parquet 通过 -`data.train_files` 直接传入,不需要转换。 - -```json -{"sample_id":"train-000017","prompt":"Question: 1 + 1 = ","response":"2"} -{"prompt":"Translate hello: ","response":"你好"} -``` - -- `sequence_no` 按非空行的文件顺序从 0 分配;并发完成顺序不会影响它。 -- `sample_id` 可选;省略时生成 `train-000000`、`train-000001` 等稳定值。 -- verl 数据的 `extra_info.index` 存在时会优先作为稳定 `sample_id`。 -- chat-message prompt 通过 target tokenizer 的 `apply_chat_template()` 编码,并加上 - generation prompt;不能把 `reward_model.ground_truth` 当作模型 response。 -- Producer tokenize `prompt` 和 `prompt + response`。后者必须以 prompt 的 token IDs - 为前缀;否则会报错,而不会猜测 response 的 loss-mask 边界。 -- `loss_mask` 中 prompt token 为 0,response token 为 1。 -- feature window 从 response 前一个 token 开始,长度由 - `max_feature_length` 限制;传给 vLLM 的 token IDs 截止于该 window 末端。 -- 其他 JSON 字段目前只作为 Producer 进程内来源元数据;第一版协议不会把它们写入 - TQ,所以 Consumer 不能读取这些字段。 - -## vLLM 与 hidden states - -对已有 response,Producer 使用 OpenAI-compatible completions API 做 prefill。 -对 prompt-only 数据,Producer 在一次请求中生成 response 并要求保存输出 hidden: - -```text -prompt= -max_tokens=<内部有界长度> -extra_body={ - "return_token_ids": true, - "kv_transfer_params": {"include_output_tokens": true} -} -``` - -响应必须同时满足: - -1. 若返回 `choices[0].prompt_token_ids`,它必须等于请求的 token IDs; -2. `kv_transfer_params.hidden_states_path` 必须存在; -3. 该文件必须含 `token_ids` 和形状为 `[seq, layers, hidden]` 的 `hidden_states`。 - -vLLM 0.23 已内置满足这个合同的 `ExampleHiddenStatesConnector`。不需要 SpeCo -Mooncake connector。在线服务必须关闭 chunked prefill,并显式配置一个 Producer -可见的临时目录。例如: - -```bash -export MODEL_PATH=/path/to/target-model -export HIDDEN_STATES_DIR=/dev/shm/speco-hidden-states -mkdir -p "${HIDDEN_STATES_DIR}" - -vllm serve "${MODEL_PATH}" \ - --host 0.0.0.0 \ - --port 8000 \ - --speculative-config \ - '{"method":"extract_hidden_states","num_speculative_tokens":1,"draft_model_config":{"hf_config":{"eagle_aux_hidden_state_layer_ids":[1,9,17,25,33,36]}}}' \ - --kv-transfer-config \ - "{\"kv_connector\":\"ExampleHiddenStatesConnector\",\"kv_role\":\"kv_producer\",\"kv_connector_extra_config\":{\"shared_storage_path\":\"${HIDDEN_STATES_DIR}\",\"use_synchronization_lock\":true}}" \ - --no-enable-chunked-prefill -``` - -上面的 layer IDs 只是 Qwen3-4B 示例。实际值必须按 target 模型和训练配置确定; -DSpark L1 开启时,vLLM 列表是 auxiliary layer IDs 加 final layer,而 Producer 的 -`TARGET_LAYER_IDS` 只填写 auxiliary 部分。 - -官方 connector 使用持久存在的 `.lock` 文件和 `flock` 协调异步落盘。Producer -读取前等待文件锁释放;TQ `put_sample` 成功后同时删除 safetensors 和 `.lock`。 - -`feature_from_vllm_payload()` 是从旧 replay 路径提取出的公共纯函数。它校验 token -对齐、选择 feature rows、拼接 auxiliary layers;DSpark L1 开启时额外拼接 final -hidden state。旧 replay 路径仍通过薄封装调用此函数,避免两套转换规则。 - -## TQ 写入和失败语义 - -每个输入 record 只写一个协议 key: - -```text -drafter:v1::<12位sequence_no>: -``` - -写入顺序是严格的: - -```text -加载临时 safetensors -→ 校验并转换 -→ TQ kv_put -→ 删除临时文件 -``` - -因此: - -- `kv_put` 失败时临时文件保留,且 Producer 不写 EOS; -- 任一请求、转换或写入失败会停止整条 Producer,不做自动重试或 endpoint 熔断; -- 只有所有 sample 都发布完成,才写 `control:v1::eos`; -- 进程退出时只调用 `close_transfer_queue_client()`,不会调用全局 `tq.close()`, - 不会销毁共享 Controller。只有 owner 可以关闭 TQ。 - -`max_pending_samples` 是简单背压:当前 run 的 ready sample 数达到该阈值时,新的 -vLLM 请求会暂停,等待 Consumer 清理已成功训练的 key。 - -## 配置与启动 - -默认 Producer 配置位于 -`speco.standalone_tq_producer`,TQ 连接配置仍位于 -`actor_rollout_ref.rollout.drafter.training.transfer_queue`。 - -必须设置的 Producer 字段: - -| 字段 | 含义 | -| --- | --- | -| `input_path` | 上述 JSONL 或 Parquet 文件 | -| `tokenizer_path` / `tokenizer_fingerprint` | 用于 tokenization 和 Consumer 合同校验 | -| `target_model_id` / `target_model_revision` | target checkpoint 身份 | -| `target_layer_ids` | auxiliary target layer IDs;DSpark L1 时 wire metadata 会额外写 `-1` 表示 final layer | -| `vllm_endpoints` / `vllm_model` | 一个或多个 OpenAI-compatible vLLM endpoint 与模型名 | - -必须与 owner/Consumer 一致的 TQ 字段: - -| 字段 | 固定要求 | -| --- | --- | -| `package_version` | `0.1.7` | -| `partition_id` | `speco_drafter_features` | -| `schema_version` | `1` | -| `run_id`、Ray address、Ray namespace | 三个进程必须相同 | - -单独调试时可以通过安装后的命令入口运行;正式训练由统一launcher启动Producer: - -```bash -verl-speco-tq-producer \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.enable=true \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.address= \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.run_id= \ - speco.standalone_tq_producer.input_path= \ - speco.standalone_tq_producer.tokenizer_path= \ - speco.standalone_tq_producer.tokenizer_fingerprint= \ - speco.standalone_tq_producer.target_model_id= \ - speco.standalone_tq_producer.target_model_revision= \ - speco.standalone_tq_producer.target_layer_ids='[2,8,14,20,26]' \ - speco.standalone_tq_producer.vllm_endpoints='[http://node0:8000/v1]' \ - speco.standalone_tq_producer.vllm_model= -``` - -完整生命周期顺序仍是:Ray/TQ backend → TQ owner → Consumer → Producer → Consumer -drain → owner shutdown。Producer 完成不代表训练完成,EOS 只表示不会再有新样本。 -正式独立训练入口 -`examples/run_qwen3-8b_drafter_separate_training.sh` 会通过 -`verl_speco.standalone_tq_training_launcher` 自动管理这套生命周期;上面的 Producer -脚本仅用于单独调试 Producer。 - -## 测试覆盖与未验证项 - -新增测试覆盖:JSONL/真实 Parquet 解析、DAPO chat prompt、target response generation、 -token 边界、多个 endpoint 的并发限制、ready 队列背压、 -成功时 sample 后 EOS 与临时文件删除、失败时无 EOS 且保留临时文件,以及旧 EAGLE3 -转换路径仍可复用公共函数。 - -这些测试使用 fake vLLM/TQ。真实 Ray + TransferQueue + vLLM 的多进程 -联调没有在当前环境执行;运行前仍需确认 vLLM 版本能返回上述 -`hidden_states_path` 以及 TQ 0.1.7 依赖环境可用。 diff --git a/docs/standalone_vllm_tq_dspark_training_plan.md b/docs/standalone_vllm_tq_dspark_training_plan.md deleted file mode 100644 index 7a865b8b..00000000 --- a/docs/standalone_vllm_tq_dspark_training_plan.md +++ /dev/null @@ -1,1130 +0,0 @@ -# 独立 vLLM Producer + TQ + DSpark Consumer 第一版方案 - -Last updated: 08/21/2026 - -## 1. 第一版要实现什么 - -只实现下面这条主链路: - -```text -verl prompt-only 数据或包含 prompt + response 的输入文件 -→ Producer 并发请求 vLLM prefill -→ Producer 将每条训练样本写入 TQ -→ Consumer 从同一个 TQ 取样本 -→ 独立 torchrun/FSDP DSpark 训练 -→ 一个 optimizer step 成功后删除该 step 的 TQ 样本 -``` - -第一版允许 **TQ 基础设施内部使用 Ray**,因为实测 `TransferQueue==0.1.7` 通过 Ray named actor 发现 `TransferQueueController`。Producer 和 Consumer 仍是普通 OS 进程,不改成 Ray actor;二者不使用 Ray RPC、Ray `ObjectRef` 或 Ray object store 传输训练样本。hidden-state payload 仍通过 TQ 的 `kv_put/kv_batch_get` 和配置的 MooncakeStore backend 传输。第一版不使用 DataProto/WorkerGroup,不生成长期 hidden-state feature store,也暂不实现复杂重试、自动重启和严格 checkpoint 数据恢复。 - -需要运行的组件: - -| 组件 | 数量 | 作用 | -|---|---:|---| -| vLLM server | 一个或多个 | 加载 target model,执行并行 prefill | -| Ray head | 1 个集群 | 保存 TQ named Controller/Storage actors;不承载 Producer/Consumer 业务 RPC | -| TQ owner | 1 个普通进程 | 连接 Ray,创建并保持 TQ controller/storage,最后统一关闭 TQ | -| Producer | 1 个进程 | 读文件、并发请求 vLLM、写 TQ | -| Consumer | 1 个 torchrun 任务 | 多个 rank 从 TQ 取数并训练 DSpark | - -Producer 和 Consumer 不互相调用,也不通过 HTTP 或 Ray 传样本。二者先连接同一个 Ray 集群,再通过携带相同 native 配置的 `tq.init(config)` 找到同一个 TQ,最后使用 TQ KV API 读写样本。已有 Controller 时 TransferQueue 0.1.7 会忽略后续配置并只连接;若 Client 意外先初始化,同一配置可避免默认 backend 抢先生效。 - -## 2. 共同的数据约定 - -这部分由两位开发者共同完成并先合入。建议文件: - -```text -verl_speco/transport/drafter_sample_protocol.py -tests/unit/test_drafter_sample_protocol.py -``` - -### 2.1 一个 key 对应一条样本 - -第一版固定: - -```text -一个输入文件 record -→ 一个 sequence_no -→ 一个 sample_id -→ 一个 TQ sample_key -→ 一个单样本 payload -``` - -`sequence_no` 是输入文件中的样本序号,不是 batch 编号。多个 sample keys 到 Consumer 后才组成训练 batch。 - -例如: - -```python -run_id = "dspark-20260818-a" -sequence_no = 17 -sample_id = "train-000017" - -partition_id = "speco_drafter_features" # 第一版沿用 PR 固定值 -sample_key = ( - "drafter:v1:dspark-20260818-a:" - "000000000017:train-000017" -) -``` - -### 2.2 Partition、key、tag 和 payload 的关系 - -TQ 中逻辑上是: - -```text -TQ 实例 -└── partition_id - └── sample_key - ├── tag - └── fields/payload -``` - -- TQ 实例:Producer 和 Consumer 共同连接的 controller/storage; -- `partition_id`:第一版沿用 PR 固定的 `speco_drafter_features`;每次启动新的 TQ owner,保证该 TQ 实例初始为空; -- `sample_key`:该分区中一条训练样本的地址; -- `tag`:轻量索引,Consumer 用 `kv_list()` 看见; -- `fields/payload`:真正的 Tensor 数据,Consumer 用 `kv_batch_get()` 读取。 - -Producer 写入: - -```python -tq.kv_put( - partition_id=partition_id, - key=sample_key, - fields=fields, - tag=tag, -) -``` - -Consumer 先发现 key: - -```python -all_records = tq.kv_list() -tags_by_key = all_records[partition_id] -``` - -这一步只拿 key 和 tag,不搬运 hidden states。 - -Consumer 再取数据: - -```python -result = tq.kv_batch_get( - partition_id=partition_id, - keys=selected_keys, -) -``` - -`kv_batch_get` 中的 batch 表示“一次读取多个独立 sample keys”,不是这些样本在 Producer 写入时就属于同一个对象。 - -### 2.3 Payload 字段 - -一个 sample key 对应的 `fields`: - -```python -fields = { - "input_ids": input_ids, # CPU int64[L] - "loss_mask": loss_mask, # CPU float32[L] - "position_ids": position_ids, # CPU int64[L] - "hidden_states": hidden_states, # CPU bf16[L,D] - "metadata_json": metadata_bytes, # CPU uint8[M] -} -``` - -| field | 含义 | Consumer 中的用途 | -|---|---|---| -| `input_ids` | 经过 feature window 选择后的 token IDs | DSpark token 输入 | -| `loss_mask` | 每个 token 是否参与 loss | 排除 prompt/padding/无效位置 | -| `position_ids` | 每个 row 对应的序列位置 | 位置编码和对齐校验 | -| `hidden_states` | target 指定层 hidden 沿最后一维拼接 | DSpark target/context 特征 | -| `metadata_json` | 模型、layers、layout、shape、样本身份 | Consumer 校验并恢复 metadata | - -符号: - -- `L`:这条训练 feature 保留的 token row 数; -- `H`:target model hidden size; -- `C`:DSpark context layer 数; -- L1 关闭:`D=C*H`,layout=`dflash_aux`; -- L1 开启:`D=C*H+H`,layout=`dflash_aux_plus_last`。 - -示例:`H=4096,C=5,L=1536`,开启 L1: - -```python -input_ids.shape == [1536] -loss_mask.shape == [1536] -position_ids.shape == [1536] -hidden_states.shape == [1536, 24576] -``` - -`metadata_json` 使用 JSON UTF-8 编码为 Tensor,因为参考 PR 的 bridge 只把 Tensor 放入 TQ fields: - -```python -raw = json.dumps(metadata, sort_keys=True).encode("utf-8") -metadata_bytes = torch.tensor(list(raw), dtype=torch.uint8) -``` - -### 2.4 Tag 字段 - -```python -tag = { - "record_type": "sample", - "status": "ready", - "schema_version": 1, - "run_id": "dspark-20260818-a", - "sequence_no": 17, - "sample_id": "train-000017", - "algorithm": "DSPARK", -} -``` - -tag 只存用于发现和筛选的 flat scalar/string。Consumer 只选择: - -```text -record_type=sample -status=ready -schema_version=1 -run_id=当前 run -algorithm=DSPARK -``` - -### 2.5 Metadata 字段 - -`metadata_json` 解码后至少包含: - -```python -metadata = { - "schema_version": 1, - "run_id": "dspark-20260818-a", - "sample_id": "train-000017", - "sequence_no": 17, - "algorithm": "DSPARK", - "target_model_id": "/models/Qwen3-8B", - "target_model_revision": "revision-or-checksum", - "tokenizer_fingerprint": "sha256:...", - "target_layer_ids": [2, 8, 14, 20, 26, -1], - "hidden_states_layout": "dflash_aux_plus_last", - "hidden_dtype": "bfloat16", - "hidden_shape": [1536, 24576], - "feature_length": 1536, - "full_sequence_length": 1800, - "feature_start": 264, - "feature_end": 1800, - "use_logits": False, -} -``` - -其中: - -- `target_model_id/revision`:hidden states 来自哪个 target checkpoint; -- `tokenizer_fingerprint`:Producer 使用的 tokenizer/template 版本; -- `target_layer_ids`:vLLM 返回和参与拼接的层; -- `hidden_states_layout`:Consumer 应如何拆 hidden 最后一维; -- `feature_length`:payload 中四个主要 Tensor 的第一维; -- `full_sequence_length`:完整 prompt+response 的 token 数; -- `[feature_start,feature_end)`:feature 在完整序列中的范围。 - -### 2.6 共享协议接口 - -Producer 和 Consumer 都 import 同一个模块,不各自手写字段名: - -```python -@dataclass(frozen=True) -class SampleMetadata: - schema_version: int - run_id: str - sample_id: str - sequence_no: int - algorithm: str - target_model_id: str - target_model_revision: str - tokenizer_fingerprint: str - target_layer_ids: list[int] - hidden_states_layout: str - hidden_dtype: str - hidden_shape: list[int] - feature_length: int - full_sequence_length: int - feature_start: int - feature_end: int - use_logits: bool - -def make_sample_key(meta: SampleMetadata) -> str: ... -def make_ready_tag(meta: SampleMetadata) -> dict: ... -def encode_sample(sample, meta: SampleMetadata) -> dict[str, Tensor]: ... -def decode_sample(key, tag, fields, expected_config) -> DraftFeatureSample: ... -def make_eos_record(run_id: str, total_samples: int): ... -``` - -Producer 使用 `SampleMetadata/make_sample_key/make_ready_tag/encode_sample`;Consumer 使用 `decode_sample`。 - -`SampleMetadata` Python 对象本身不经过 TQ: - -```text -Producer SampleMetadata -→ JSON -→ uint8 Tensor -→ TQ metadata_json -→ uint8 Tensor -→ JSON -→ Consumer metadata dict -``` - -`decode_sample()` 负责: - -1. 解码 `metadata_json`; -2. 校验 key、tag、metadata 中的 sample 身份一致; -3. 校验模型、tokenizer、layers 和 layout 与 Consumer 配置一致; -4. 校验 Tensor 必需字段、dtype 和 shape; -5. 返回现有 `DraftFeatureSample`。 - -### 2.7 EOS - -Producer 完成全部输入后写一个控制 record: - -```python -eos_key = f"control:v1:{run_id}:eos" -eos_fields = {"marker": torch.tensor([1], dtype=torch.uint8)} -eos_tag = { - "record_type": "control", - "status": "eos", - "schema_version": 1, - "run_id": run_id, - "total_samples": total_samples, -} -``` - -EOS 不进入训练 batch。Consumer 看到 EOS 后继续处理剩余 ready samples;`EOS 已出现且 ready 为空` 时结束。 - -## 3. TQ 怎么启动,Producer 和 Consumer 怎么连接同一个 TQ - -### 3.1 已验证的 TQ 0.1.7 连接机制 - -`TransferQueue==0.1.7` 没有提供“把 Controller 地址直接传给第二个进程”的高层连接接口。`tq.init(config)` 会先尝试以下已有 Controller 连接逻辑;存在时忽略传入配置,不存在时才用配置创建服务: - -```python -_TQ_CONTROLLER = ray.get_actor("TransferQueueController") -conf = ray.get(_TQ_CONTROLLER.get_config.remote()) -_maybe_create_tq_client(conf) -``` - -因此所有进程必须先加入同一个 Ray 集群。Ray 在本方案中只承担 TQ 控制面:保存 named `TransferQueueController`、返回 TQ 配置、管理 TQ 创建的 actor。业务进程不通过 Ray 发送 sample key 或 hidden-state tensor。 - -实际连接链路是: - -```text -TQ owner:ray.init(address) → tq.init(full_tq_config) → 创建 named Controller -Producer:ray.init(address) → tq.init(same config) → ray.get_actor() → 创建本地 TQ client -Consumer rank 0..N:ray.init(address) → tq.init(same config) → ray.get_actor() → 创建各自 TQ client -``` - -### 3.2 直接移植并扩展 PR #48 的 bridge - -参考文件: - -```text -C:/Users/xxyyrr/Desktop/上班/verl-SpeCo/ - verl_speco/integration/transferqueue_bridge.py -``` - -第一版不新增 `standalone_tq.py`,直接把该文件移植到目标项目并扩展。保留 PR 已有的 import 检查、进程内幂等状态、`kv_put`、`kv_batch_get`、TensorDict 解包和 owner-only shutdown;新增 Ray 显式连接、批量 list/get/clear 和本地 client 关闭接口。 - -目标文件必须实现以下函数,而不只是提供一个笼统的 transport class: - -```python -def configure_transfer_queue(config: Mapping[str, Any]) -> bool: ... -def connect_ray_cluster(ray_address: str, namespace: str | None = None) -> None: ... -def start_transfer_queue_owner(tq_config: Mapping[str, Any]) -> None: ... -def connect_transfer_queue_client() -> None: ... -def make_sample_key(run_id: str, sequence_no: int, sample_id: str) -> str: ... -def put_sample(key: str, fields: dict[str, Tensor], tag: dict[str, Any]) -> None: ... -def list_samples() -> dict[str, dict[str, Any]]: ... -def get_samples(keys: list[str]) -> list[tuple[str, dict[str, Tensor]]]: ... -def clear_samples(keys: list[str]) -> None: ... -def close_transfer_queue_client() -> None: ... -def close_transfer_queue_owner() -> None: ... -``` - -逐个函数的责任如下。 - -#### `configure_transfer_queue(config)` - -- 从 Hydra/OmegaConf 中读取 Ray address、可选 namespace、固定 partition、run ID、schema version 和 TQ native backend 配置; -- 转成普通 Python dict,保存在进程内 `_state`; -- 校验 `TransferQueue==0.1.7` 可 import; -- 不连接 Ray,不创建 TQ,不产生跨进程副作用; -- 返回该进程是否启用了 TQ。 - -#### `connect_ray_cluster(ray_address, namespace)` - -- 若 `ray.is_initialized()` 为 false,调用 `ray.init(address=ray_address, namespace=namespace)`; -- 若已经初始化,校验现有 Ray context 指向期望集群,不能悄悄连到另一个本地 Ray; -- Owner、Producer 和所有 torchrun ranks 都调用它; -- 此函数只建立当前 OS 进程到 Ray control plane 的连接。 - -#### `start_transfer_queue_owner(tq_config)` - -- 仅由 `tq_owner.py` 调用; -- 前置条件是 `connect_ray_cluster()` 已成功; -- 调用一次 `tq.init(OmegaConf.create(tq_config))`; -- 将 `_state.owner=True`、`_state.initialized=True`; -- TQ 0.1.7 会创建名为 `TransferQueueController` 的 Ray actor,并创建所选 storage backend; -- 重复调用必须报错,不能启动第二套同名 Controller。 - -#### `connect_transfer_queue_client()` - -- 由 Producer 和每个 Consumer rank 调用; -- 前置条件是当前进程已经连接 Ray; -- 调用 `tq.init(same native config)`,通过 `ray.get_actor("TransferQueueController")` 发现 owner;已有 Controller 时配置会被忽略,意外抢先时则以相同配置创建; -- 只创建当前进程的 TQ client,不创建新的 Controller; -- 成功后设置 `_state.initialized=True`;重复调用直接返回。 - -#### `put_sample/list_samples/get_samples/clear_samples` - -- 全部固定使用 `_SPECO_TQ_PARTITION = "speco_drafter_features"`; -- `put_sample()` 调用单样本 `tq.kv_put()`; -- `list_samples()` 调用 `tq.kv_list(partition_id=...)`,只返回 key/tag 元数据; -- `get_samples()` 一次调用 `tq.kv_batch_get(keys=...)`,再按输入 key 顺序解包为普通 dict; -- `clear_samples()` 调用 `tq.kv_clear(keys=..., partition_id=...)`; -- 这些函数不调用 Ray RPC,不把 tensor 放入 Ray object store。 - -#### `close_transfer_queue_client()` 与 `close_transfer_queue_owner()` - -TQ 0.1.7 的公共 `tq.close()` 会 kill 共享 Controller,所以两者不能写成同一个实现: - -- `close_transfer_queue_client()`:通过 0.1.7 已公开的 `tq.get_client()` 取得当前进程 client并调用它的 `close()`,然后 `ray.shutdown()`;绝不能调用会 kill Controller 的全局 `tq.close()`。这个 client close 只用于进程退出阶段,关闭后本进程不能再次调用 TQ; -- `close_transfer_queue_owner()`:仅当 `_state.owner=True` 时调用 `tq.close()`,清理 Controller/Storage,最后 `ray.shutdown()`; -- Producer 或任一训练 rank 提前退出都不能关闭全局 TQ。 - -### 3.3 共享配置 - -Owner、Producer 和 Consumer 必须使用相同的 Ray address/namespace,并通过同一个 named Controller 获得 backend 配置。第一版 partition 固定,不按任务动态创建: - -```yaml -transfer_queue: - enable: true - package_version: "0.1.7" - ray: - address: "ray-head-node:6379" - namespace: "speco-drafter" - partition_id: "speco_drafter_features" - run_id: "dspark-20260819-a" - schema_version: 1 - backend: - storage_backend: MooncakeStore - MooncakeStore: - auto_init: false - metadata_server: "node0:50050" - master_server_address: "node0:50051" - local_hostname: "" - protocol: tcp - global_segment_size: 4294967296 - local_buffer_size: 1073741824 - device_name: "" -``` - -`partition_id/run_id/schema_version` 用于过滤和校验样本;它们不能帮助进程发现 TQ。真正让三个任务连接到同一 TQ 的是“连接同一个 Ray 集群和 namespace,然后找到同名 Controller”。 - -依赖也必须锁定并单独验证。实机安装 `TransferQueue==0.1.7` 会安装 Ray,并要求 `numpy<2.0.0`;当前测试环境中的 `twinkle-kit` 要求 `numpy>=2.0.0`,两者冲突。开发时应使用专门的 TQ/训练环境或重新确认整套依赖约束,不能直接把 0.1.7 安装进已有生产环境后忽略 resolver warning。 - -### 3.4 `tq_owner.py` 要实现的入口和函数 - -新增: - -```text -verl_speco/tq_owner.py -``` - -`tq_owner.py` 建议明确实现: - -```python -def install_signal_handlers(stop_event: threading.Event) -> None: ... -def publish_owner_ready(run_id: str, schema_version: int) -> None: ... -def wait_until_stopped(stop_event: threading.Event) -> None: ... -def run_owner(config: DictConfig) -> int: ... -def main() -> None: ... -``` - -`run_owner()` 的执行顺序必须是: - -```text -configure_transfer_queue(config) -→ connect_ray_cluster(ray.address, ray.namespace) -→ start_transfer_queue_owner(full TQ native config) -→ put owner_ready 控制 record -→ 安装 SIGINT/SIGTERM handler -→ 保持 owner 进程存活 -→ 收到停止信号 -→ close_transfer_queue_owner() -``` - -Owner 必须常驻。TQ 0.1.7 创建 Controller 时没有设置 `lifetime="detached"`,不能在初始化后立即退出。 - -### 3.5 启动和关闭顺序 - -第一版由外部脚本管理全生命周期: - -```text -1. ray start --head,记录 Ray address -2. 启动 Mooncake metadata/master(若 auto_init=false) -3. 启动 TQ owner;owner 连接 Ray并调用 tq.init(full config) -4. 等待 owner_ready -5. 启动一个或多个 vLLM servers -6. 启动 Consumer;每个 torchrun rank 连接 Ray,然后 tq.init(same native config) -7. 启动 Producer;连接 Ray,然后 tq.init(same native config) -8. Producer 写 EOS,关闭本地 client并退出 -9. Consumer drain、保存 final checkpoint,所有 ranks 关闭本地 client并退出 -10. 给 TQ owner 发送 SIGTERM;仅 owner 执行 tq.close() -11. 等 owner 退出后执行 ray stop -12. 停止 Mooncake 服务 -``` - -外部脚本要用 `trap` 保证异常退出也按“Producer/Consumer → owner → Ray → Mooncake”的顺序清理。不能在 Consumer rank 的 `finally` 中调用全局 `tq.close()`。 - -## 4. Producer 要实现什么 - -### 4.1 Producer 完整顺序 - -```text -读取共享配置 -→ 连接 TQ并校验 owner_ready -→ 初始化 tokenizer -→ 初始化多个 vLLM endpoint clients -→ 流式读取输入文件 -→ 为每条输入分配 sequence_no/sample_id -→ 缺少 response 时由 target vLLM 生成;构造 input_ids/loss_mask -→ 并发请求 vLLM prefill -→ 读取 vLLM hidden-state 临时结果 -→ 转换成 DSpark DraftFeatureSample -→ 构造 SampleMetadata -→ encode_sample 得到 fields/tag/key -→ TQ kv_put 一条 sample -→ 删除该请求临时文件 -→ 所有输入完成后写 EOS -→ close_transfer_queue_client()并退出 -``` - -### 4.2 并发模型 - -Producer 是一个进程,内部并发请求多个 endpoint: - -```text -InputReader -→ bounded asyncio input_queue -→ N 个 RequestWorker -→ bounded publish_queue -→ TQ Publisher -``` - -- `vllm_endpoints` 是列表; -- 每个 endpoint 有独立 semaphore; -- 总并发由 `max_inflight_requests` 限制; -- input/publish queue 必须有上限; -- TQ `kv_put` 如果是同步 API,用 `asyncio.to_thread()` 调用; -- TQ ready 数达到 `max_pending_samples` 时暂停继续请求,形成简单背压。 - -`sequence_no` 在 InputReader 中分配,不按 vLLM 完成顺序分配。因此并发乱序不会改变 sample key。 - -### 4.3 vLLM 结果转换 - -复用现有 `TargetFeatureReplayer` 的: - -- OpenAI-compatible vLLM 请求; -- `prompt_token_ids` 校验; -- `kv_transfer_params.hidden_states_path`; -- safetensors 加载; -- `[seq,layers,hidden]` 校验; -- feature positions 选择; -- aux layers flatten; -- DSpark L1 时拼 final hidden。 - -不要复制 `_feature_from_vllm_payload()`;把纯转换逻辑提取成公共函数,让旧 file backend 和新 Producer 共同调用。 - -临时文件顺序: - -```text -加载 -→ 校验/转换 -→ TQ put 成功 -→ 删除 -``` - -第一版仍允许每个并发请求产生短期临时文件,但不生成全量 hidden-state 数据集。 - -### 4.4 Producer 文件分工 - -| 文件 | 实现内容 | -|---|---| -| `verl_speco/standalone_tq_producer.py` | CLI/Hydra 入口,连接 TQ,启动 asyncio pipeline,写 EOS和汇总指标 | -| `verl_speco/producer/input_reader.py` | 流式读输入、分配 `sequence_no/sample_id`、tokenize、构造 loss mask | -| `verl_speco/producer/vllm_feature_client.py` | endpoint pool、并发控制、HTTP 请求、临时结果读取和删除 | -| `verl_speco/trainer/target_feature_replay.py` | 提取可复用的 vLLM payload → `DraftFeatureSample` 纯转换函数 | -| `tests/unit/test_drafter_sample_protocol.py` | 协议编码、字段和 shape 测试(两人共同) | -| `tests/integration/test_tq_producer_smoke.py` | fake/短 vLLM → TQ sample + EOS | - -Producer 开发者同时负责移植/扩展 `integration/transferqueue_bridge.py` 和实现 `tq_owner.py`,因为这部分与 TQ 写入和连接直接相关。 - -### 4.5 Producer 各文件的函数级实现规格 - -#### `verl_speco/standalone_tq_producer.py` - -需要实现: - -```python -@dataclass -class ProducerStats: - input_count: int - published_count: int - failed_count: int - pending_bytes: int - -async def publish_one(result: PreparedFeature, transport) -> str: ... -async def run_producer(config: DictConfig) -> ProducerStats: ... -def validate_producer_config(config: DictConfig) -> None: ... -def main() -> None: ... -``` - -`main()` 只负责 Hydra/日志/退出码。`run_producer()` 是可测试的业务入口,执行:连接 Ray → 连接 TQ client → 创建 InputReader 和 vLLM client pool → 启动有界 asyncio pipeline → 等待所有请求及 `kv_put` 完成 → 发布 EOS → 关闭本地 client。它不能启动 Ray head、不能创建 TQ owner、不能调用全局 `tq.close()`。 - -`publish_one()` 接收已经完成转换的一条 `PreparedFeature`,调用共享协议的 `encode_sample()` 得到 `(key, fields, tag)`,再通过 bridge 的 `put_sample()` 发布。只有 `put_sample()` 成功返回,`published_count` 才增加,vLLM 临时文件才允许删除。 - -#### `verl_speco/producer/input_reader.py` - -需要实现: - -```python -@dataclass(frozen=True) -class InputRecord: - sequence_no: int - sample_id: str - prompt: str - response: str | None - source_metadata: dict[str, Any] - -def iter_input_records(path: str) -> Iterator[InputRecord]: ... -def tokenize_record(record: InputRecord, tokenizer, config) -> TokenizedRequest: ... -def build_loss_mask(input_ids: Tensor, prompt_length: int) -> Tensor: ... -``` - -`iter_input_records()` 流式读取 JSONL/Parquet,不把全文件载入内存,并按文件顺序分配稳定的 `sequence_no`。已有 response 时 `tokenize_record()` 直接拼接;prompt-only verl 数据通过 chat template 编码后由 target vLLM 生成 response,并设置 `include_output_tokens=true` 同步提取输出 hidden states。 - -#### `verl_speco/producer/vllm_feature_client.py` - -需要实现: - -```python -@dataclass(frozen=True) -class VllmEndpoint: - base_url: str - max_concurrency: int - -class VllmFeatureClientPool: - async def start(self) -> None: ... - async def prefill(self, request: TokenizedRequest) -> RawVllmFeature: ... - async def close(self) -> None: ... - -async def request_prefill(endpoint, request) -> VllmResponse: ... -def choose_endpoint(endpoints, state) -> VllmEndpoint: ... -def load_hidden_state_result(response) -> RawVllmFeature: ... -def delete_temporary_result(raw: RawVllmFeature) -> None: ... -``` - -`prefill()` 必须允许多个 coroutine 同时运行;全局 semaphore 限制总并发,每个 endpoint 另有独立 semaphore。`request_prefill()` 只负责 HTTP 请求和响应校验;`load_hidden_state_result()` 负责读取 vLLM 返回的临时 safetensors/path。临时结果的删除不放在 `load_hidden_state_result()`,而由 `publish_one()` 成功后触发。 - -#### `verl_speco/trainer/target_feature_replay.py` - -把当前类内部的纯转换部分抽成: - -```python -def feature_from_vllm_payload( - payload: RawVllmFeature, - request: TokenizedRequest, - feature_config: FeatureContract, -) -> DraftFeatureSample: ... -``` - -它不进行 HTTP、不访问 TQ、不删除文件,只完成 shape/layout/layer 校验、feature-position 选择、aux layer flatten 和可选 L1 final hidden 拼接。现有 file replay backend 和新 Producer 都调用这一函数,避免两套转换规则。 - -#### Producer 启动配置 - -统一 launcher 向 `verl_speco.standalone_tq_producer` 提供同一套: - -```text -RAY_ADDRESS / Ray namespace -run_id / schema_version / 固定 partition -Mooncake/TQ backend 配置 -输入文件和 tokenizer/model 配置 -vLLM endpoint 列表 -max_inflight_requests / per_endpoint_concurrency -``` - -正式运行不再保留单独的角色 shell wrapper。 - -## 5. Consumer 要实现什么 - -### 5.1 不新写另一套训练器 - -继续使用现有入口: - -```text -draft_train_launcher.py -→ draft_train.py -→ trainer/draft_training_loop.py -→ DrafterBaseTrainer -→ DSparkTrainerBackend -``` - -训练模式继续使用现有 standalone `training.mode=offline`,用 `feature_store.type=tq` 选择 TQ 数据源。不复制 FSDP、loss、optimizer 或 checkpoint 代码。 - -当前 `feature_store.type` 的作用是告诉 `build_feature_store_from_config()` 应创建哪一种训练数据来源: - -| 当前 type | 对象 | 数据来源 | -|---|---|---| -| `torch_shard` | `TorchShardFeatureStore` | 本地 `.pt` shard,当前 separate-training 示例使用它 | -| `token_replay` | `TokenReplayFeatureStore` | 保存 token replay 的 shard | -| `vllm_safetensors` / `safetensors` | `VllmSafetensorsFeatureStore` | 已生成的 safetensors feature | -| `jsonl_token_replay` / `jsonl` | `JsonlTokenReplayFeatureStore` | JSONL token/text replay | - -第一版新增: - -```yaml -actor_rollout_ref: - rollout: - drafter: - training: - mode: offline - feature_store: - type: tq - path: null - shuffle: false - repeat: false -``` - -这里 `offline` 的含义是“独立于 RL 的 standalone drafter training”,不等于数据必须来自磁盘。 - -不能只给现有工厂增加一个 `type=tq` 然后继续使用通用 `DraftFeatureDataLoader`。当前 loader 会: - -```python -keys = list(store.iter_keys(...)) -``` - -它假设数据集 keys 是一个静态快照;空 store 会直接结束,并且各 rank 独立枚举时可能在 Producer 持续写入的过程中看到不同快照。TQ 是流式、会新增并删除 keys 的数据源,所以 `type=tq` 必须选择专用的 `TQFeatureDataLoader`,由 rank 0 统一发现和分配 keys。 - -#### 5.1.1 当前磁盘 feature store 为什么每个 rank 都会取 keys - -`draft_train_launcher.py` 通过 `torchrun` 启动多个训练进程。每个进程都是一个 rank,并且每个 rank 都会独立进入 `run_standalone_draft_training()`、创建 `DraftFeatureDataLoader`、执行它的 `__iter__()`。当前 loader 的核心逻辑是: - -```python -keys = list( - self.store.iter_keys( - shuffle=self.shuffle, - seed=self.seed + epoch, - ) -) -rank_keys = keys[rank::world_size] - -for key in rank_keys: - batch.append(self.store.read(key)) -``` - -因此当前并不是 rank 0 枚举 keys 后再发送给其他 rank,而是所有 rank 都访问相同的静态 feature-store 路径: - -```text -rank 0:iter_keys() 得到完整静态列表 → keys[0::world_size] → read 自己的 samples -rank 1:iter_keys() 得到完整静态列表 → keys[1::world_size] → read 自己的 samples -... -``` - -例如 store 中固定存在: - -```python -keys = ["k0", "k1", "k2", "k3", "k4", "k5", "k6", "k7"] -``` - -当 `world_size=2` 时,两个 rank 都先得到上述完整列表,然后分别计算: - -```python -# rank 0 -rank_keys = keys[0::2] # ["k0", "k2", "k4", "k6"] - -# rank 1 -rank_keys = keys[1::2] # ["k1", "k3", "k5", "k7"] -``` - -这种实现能够成立,是因为 `.pt`、safetensors 或 replay 文件在训练期间是静态数据集。只要所有 rank 使用同一路径和相同 shuffle seed,`iter_keys()` 就会产生相同顺序的 key 快照,各 rank 可以无通信地算出互不重叠的子集。 - -#### 5.1.2 TQ 为什么必须改成 rank 0 发现 keys - -TQ 中的 keys 会在训练期间动态变化:Producer 持续 `kv_put`,训练成功后 Consumer 又执行 `kv_clear`。如果所有 rank 仍然各自调用 `kv_list`,不同调用时刻可能看到不同快照: - -```python -# rank 0 较早调用 -rank0_keys = ["k0", "k1", "k2", "k3"] - -# Producer 随后写入 k4、k5,rank 1 较晚调用 -rank1_keys = ["k0", "k1", "k2", "k3", "k4", "k5"] -``` - -各 rank 再独立切片后,可能得到不同数量的 local samples,进而无法保证它们以相同顺序进入 forward、backward 和梯度 collective。为此,TQ 专用 loader 必须把控制面和数据面分开: - -```text -控制面:rank 0 执行 kv_list,固定本 step 的 global_keys,并向各 rank 分发 local_keys -数据面:每个 rank 使用自己的 local_keys 直接执行 kv_batch_get,从 TQ/Mooncake 读取 hidden-state payload -``` - -rank 0 只发送较小的 key/tag 元数据,不读取并转发其他 rank 的 hidden-state tensor。所有 rank 仍然都会连接同一个 TQ,也都会调用批量读取接口;只有动态 key 的发现和本 step 的 global batch 决策集中在 rank 0。 - -因此 `feature_store.type=tq` 不是只替换底层 `read()`:它同时改变了 key 的发现、分配、等待和删除语义,需要专用 `TQFeatureDataLoader`,并要求训练循环在所有 rank 成功完成 optimizer step 后,由 rank 0 对本 step 的 `global_keys` 执行 `kv_clear`。 - -### 5.2 Consumer 完整顺序 - -```text -torchrun 启动多个 ranks -→ 每个 rank 初始化 torch.distributed -→ 每个 rank 连接同一个 TQ -→ rank 0 校验 owner_ready,并 broadcast 结果 -→ 初始化现有 DSpark trainer -→ rank 0 kv_list 查找 ready sample keys -→ rank 0 选一个 global batch并分给各 rank -→ 每个 rank kv_batch_get 自己的 local keys -→ decode_sample 得到 list[DraftFeatureSample] -→ prepare_training_batch_from_samples() -→ training_step_from_batch() -→ 所有 rank 汇总 success -→ 成功后 rank 0 kv_clear 这个 global batch 的 keys -→ 继续下一批 -→ 看到 EOS 且 ready 为空 -→ 保存 final checkpoint -→ 所有 ranks close_transfer_queue_client()并退出 -``` - -### 5.3 多 rank 如何分 key - -例如: - -```text -world_size=2 -batch_size_per_gpu=2 -global batch size=4 -``` - -rank 0 选出: - -```python -global_keys = ["k10", "k11", "k12", "k13"] -assignments = [ - ["k10", "k11"], # rank 0 - ["k12", "k13"], # rank 1 -] -``` - -通过 `dist.scatter_object_list` 或 broadcast 发送短字符串列表。每个 rank 直接从 TQ 读取自己的 payload,不通过 rank 0 转发 hidden states。 - -### 5.4 从 TQ 到训练 batch - -每个 rank: - -```python -records = tq_transport.get_samples(local_keys) - -samples = [ - decode_sample( - key=key, - tag=tags_by_key[key], - fields=fields, - expected_config=expected_contract, - ) - for key, fields in records -] - -batch = trainer.prepare_training_batch_from_samples( - samples, - step=optimizer_step, -) - -ok = await trainer.training_step_from_batch( - batch, - optimizer_step, -) -``` - -`records` 是多个独立单样本 payload;`samples` 是变长 `DraftFeatureSample` 列表。Consumer transport 层不要直接 `torch.stack()`,现有 trainer/backend 负责对齐和组 batch。 - -### 5.5 删除与结束 - -第一版采用简单逻辑: - -```text -所有 rank get/decode/train 都成功 -→ all_reduce(global_success)=True -→ rank 0 kv_clear(global_batch_keys) -``` - -任何 rank 失败都不 clear。程序报错退出,第一版不自动恢复。 - -EOS 后不足一个 global batch 的尾部,第一版使用 `drop_last=True`:记录数量,rank 0 clear 这些尾部 keys,然后正常结束。 - -checkpoint 仍使用现有 `save_interval_steps` 和 final checkpoint;第一版不保证崩溃后已 clear 数据能够严格重放。 - -### 5.6 Consumer 文件分工 - -| 文件 | 实现内容 | -|---|---| -| `verl_speco/trainer/feature_store.py` | 工厂增加 `type=tq`,返回 `TQFeatureStore`;TQ 类型不要求 `path` | -| `verl_speco/trainer/tq_feature_store.py` | 实现固定 partition 上的 list/get-many/clear/EOS,不伪装静态 `iter_keys()` | -| `verl_speco/trainer/tq_sample_source.py` | 实现专用 `TQFeatureDataLoader`:rank 0 动态 list/filter/sort、key 分配、各 rank get/decode、EOS/tail | -| `verl_speco/trainer/draft_training_loop.py` | 保持 `mode=offline`;`type=tq` 时跳过 `feature_store.path` 必填检查并选择专用 loader;调用现有 prepare/train,global success 后 clear,final checkpoint | -| `verl_speco/draft_train_launcher.py` | `feature_store.type=tq` 时不要求 feature-store path,透传 torchrun/TQ 配置 | -| `verl_speco/config/speco_base.yaml` | 新增 TQ connection、run、poll、batch 等默认配置 | -| `tests/unit/test_tq_sample_source.py` | key 过滤/排序/分 rank、EOS、decode 调用测试 | -| `tests/integration/test_dspark_tq_consumer_smoke.py` | 两 rank 取不同样本,完成训练并 clear | - -### 5.7 Consumer 各文件的函数级实现规格 - -#### `verl_speco/trainer/feature_store.py` - -修改现有工厂: - -```python -def build_feature_store_from_config(feature_store_cfg, read_only=False): - store_type = str(feature_store_cfg.get("type", "torch_shard")).lower() - if store_type == "tq": - return TQFeatureStore.from_config(feature_store_cfg) - ... -``` - -要求: - -- 保留 `torch_shard/token_replay/vllm_safetensors/jsonl_token_replay` 的当前行为; -- `type=tq` 时不读取 `path`; -- 工厂只创建对象,不在 import 阶段连接 Ray/TQ; -- `TQFeatureStore` 是流式数据源适配器,不强行实现有误导性的静态 `iter_keys()` 和逐条 `read()`。 - -#### `verl_speco/trainer/tq_feature_store.py` - -需要实现: - -```python -@dataclass(frozen=True) -class ReadyEntry: - key: str - tag: dict[str, Any] - -class TQFeatureStore: - @classmethod - def from_config(cls, cfg) -> "TQFeatureStore": ... - def connect(self) -> None: ... - def list_ready(self, run_id: str) -> list[ReadyEntry]: ... - def get_many(self, entries: list[ReadyEntry]) -> list[DraftFeatureSample]: ... - def clear_many(self, keys: list[str]) -> None: ... - def read_eos(self, run_id: str) -> EosMetadata | None: ... - def close_local(self) -> None: ... -``` - -`connect()` 调用共享 bridge 的 `connect_ray_cluster()` 和 `connect_transfer_queue_client()`。`list_ready()` 调用 `list_samples()` 后只保留 `tag.status=ready`、匹配 `run_id/schema_version` 的数据 key,并按 `(sequence_no,key)` 排序。`get_many()` 一次批量读取 fields,然后逐条调用共享 `decode_sample()`;返回顺序必须和 entries 相同。`clear_many()` 只允许 rank 0 在全局 step 成功后调用。`close_local()` 只关闭本 rank client,不能关闭 Controller。 - -#### `verl_speco/trainer/tq_sample_source.py` - -需要实现: - -```python -@dataclass -class TQLocalBatch: - local_keys: list[str] - local_samples: list[DraftFeatureSample] - global_keys: list[str] | None - -class TQFeatureDataLoader: - def __iter__(self) -> Iterator[TQLocalBatch]: ... - def _select_global_batch(self) -> tuple[list[str], dict[str, ReadyEntry]]: ... - def _broadcast_assignments(self, assignments) -> list[ReadyEntry]: ... - def _handle_eos_and_tail(self) -> bool: ... - def clear_completed_batch(self, global_keys: list[str] | None) -> None: ... -``` - -执行责任必须明确: - -- 所有 rank 创建 loader 并调用 `store.connect()`; -- 只有 rank 0 执行 `_select_global_batch()` 和 `kv_list`; -- rank 0 生成 `list[list[ReadyEntry]]` assignments,使用 `torch.distributed.scatter_object_list` 或 broadcast 分发小型 key/tag; -- 每个 rank 对自己的 local entries 调用 `store.get_many()`,hidden states 从 TQ/Mooncake 直接进入该 rank; -- loader yield `TQLocalBatch`,不能只 yield samples,因为训练后清理还需要 `global_keys`; -- 暂时为空时轮询等待,不能像静态 loader 一样结束;只有 EOS 已出现并且 ready 数据 drain 完才停止; -- `drop_last=True` 时由 rank 0 记录并清理不足 global batch 的尾部 key。 - -#### `verl_speco/trainer/draft_training_loop.py` - -需要新增或调整: - -```python -def build_training_source(config, rank, world_size): ... -def all_ranks_succeeded(local_ok: bool, device) -> bool: ... -async def run_tq_training_loop(trainer, loader, config) -> dict[str, Any]: ... -``` - -`build_training_source()` 根据 `feature_store.type` 分支:静态类型继续创建 `DraftFeatureDataLoader`;`tq` 创建 `TQFeatureDataLoader`,跳过 `feature_store.path` 必填检查,并禁止 `shuffle/repeat`。`run_tq_training_loop()` 对每个 `TQLocalBatch` 调用现有 `prepare_training_batch_from_samples()` 和 `training_step_from_batch()`;所有 rank 通过 collective 得到 global success 后,才让 rank 0 调用 `clear_completed_batch(global_keys)`。异常路径不 clear,finally 只执行 `store.close_local()`。 - -#### `verl_speco/draft_train_launcher.py` - -保留当前 `torch.distributed.run` 启动方式,只增加配置校验和环境透传: - -```python -def validate_tq_launch_config(overrides, launch_config) -> None: ... -def build_child_env(config) -> dict[str, str]: ... -``` - -它不启动 Ray head、不调用 `tq.init()`。职责是确认 `type=tq` 时提供了 Ray address/namespace,且不要求 feature-store path;随后把同一 Ray connection 配置传给每个 torchrun 子进程。每个子 rank 自己建立 Ray/TQ client,不能在 launcher 父进程建一个 client后期待 fork 继承。 - -#### `verl_speco/config/speco_base.yaml` - -增加默认字段: - -```yaml -feature_store: - type: torch_shard - path: null - shuffle: true - repeat: true - tq: - ray_address: null - ray_namespace: speco-drafter - partition_id: speco_drafter_features - run_id: null - schema_version: 1 - poll_interval_seconds: 0.5 - connect_timeout_seconds: 120 - drop_last: true -``` - -当 `type=tq` 时运行期覆盖为 `shuffle=false/repeat=false`。backend 的完整 owner 配置只需要 TQ owner 使用;Producer/Consumer client 从 named Controller 获取它,不各自重新创建 backend。 - -#### Consumer 测试必须覆盖的函数边界 - -- `test_feature_store_factory_builds_tq_without_path()`; -- `test_rank0_filters_and_sorts_ready_entries()`; -- `test_nonzero_rank_never_calls_kv_list()`; -- `test_assignments_are_disjoint_and_global_batch_complete()`; -- `test_each_rank_gets_only_local_keys()`; -- `test_decode_preserves_hidden_states_layout()`; -- `test_clear_only_after_all_ranks_success()`; -- `test_failure_does_not_clear()`; -- `test_eos_drains_ready_then_stops()`; -- `test_client_close_does_not_kill_owner()`。 - -## 6. 两个人怎么分工 - -### 共同先完成 - -1. `drafter_sample_protocol.py`; -2. Ray/TQ connection 配置字段; -3. 一个小型 golden sample; -4. 启动一个测试 Ray head 和 TQ owner,独立进程 A put、独立进程 B list/get/clear 的 smoke test; -5. 验证 client 退出不会 kill owner,只有 owner shutdown 才销毁 Controller。 - -### 开发者 A:Producer/TQ - -负责: - -```text -integration/transferqueue_bridge.py -tq_owner.py -standalone_tq_producer.py -producer/input_reader.py -producer/vllm_feature_client.py -target_feature_replay.py 的公共转换函数 -owner/producer 启动脚本 -Producer/TQ 测试 -``` - -开发者 A 的可交付接口不是“提供一个 TQ 类”,而是: - -```text -bridge:connect_ray_cluster/start_owner/connect_client/put/list/get_many/clear/client_close/owner_close -owner:run_owner/main/signal handler/owner_ready -producer:run_producer/publish_one/统计与 EOS -input reader:iter_input_records/tokenize_record/build_loss_mask -vLLM client:endpoint pool/request_prefill/load/delete -feature conversion:feature_from_vllm_payload -``` - -开发者 B 可以先针对这些接口写 fake transport,不需要等待真实 vLLM 和 Mooncake 联通。 - -### 开发者 B:Consumer/训练 - -负责: - -```text -feature_store.py 的 type=tq 工厂分支 -tq_feature_store.py -tq_sample_source.py / TQFeatureDataLoader -draft_training_loop.py 的 offline + type=tq 分支 -draft_train_launcher.py 配置适配 -speco_base.yaml Consumer 配置 -Consumer 启动脚本 -Consumer/DSpark 测试 -``` - -开发者 B 的可交付接口是: - -```text -feature-store factory:type=tq 分支 -TQFeatureStore:connect/list_ready/get_many/clear_many/read_eos/close_local -TQFeatureDataLoader:select global batch/distribute local keys/get/yield/EOS-tail -training loop:build source/train/global success/clear/final checkpoint -launcher:TQ 配置校验和 torchrun 子进程环境透传 -``` - -### 联调入口 - -建议再提供: - -```text -examples/run_dspark_tq_pipeline_local.sh -``` - -只用于单机联调,顺序启动: - -```text -ray start --head -→ Mooncake metadata/master -→ TQ owner(ray.init + tq.init(full config)) -→ owner_ready -→ vLLM health check -→ Consumer -→ Producer -→ 等 Producer/Consumer 退出 -→ SIGTERM TQ owner(owner 执行 tq.close) -→ ray stop -→ 停止 Mooncake -``` - -最小联调:Producer 发布 8 条样本,2 个 Consumer ranks、每 rank batch size 2,完成 2 个 optimizer steps,8 个 sample keys 被清理,EOS 后保存 final checkpoint。 - -## 7. 第一版验收标准 - -1. TQ owner、Producer 和所有 Consumer ranks 日志显示相同 Ray address/namespace、TQ Controller、partition 和 run ID。 -2. 在同一 Ray 集群中,独立 owner 创建 TQ 后,独立进程 A put,独立进程 B 能 list/get/clear。 -3. Producer 对多个 vLLM endpoints 并发请求,不串行访问。 -4. 一个输入 record 只生成一个 sample key 和一个 payload。 -5. `kv_list` 只拿 tag;`kv_batch_get` 才拿 hidden states。 -6. Consumer 各 rank 读取不同 local keys,不通过 rank 0 搬运 Tensor。 -7. `decode_sample` 能拒绝模型、layer、layout、dtype 或 shape 不匹配的数据。 -8. 所有 rank 训练成功后才 clear 当前 global batch。 -9. Producer 先完成时,Consumer 能 drain 后再退出。 -10. 不产生长期 hidden-state feature store。 -11. Producer 或任一 Consumer rank 退出不会销毁 TQ Controller;只有 owner 调用全局 `tq.close()`。 -12. Ray object store 中不承载 hidden-state payload,训练 tensor 通过 TQ/Mooncake 路径读取。 - -## 8. 后续建议:第一版跑通后再做 - -以下内容不进入第一版开发: - -- Producer HTTP/TQ 复杂重试和 endpoint 熔断; -- Producer 发布 journal,避免重启后重复生成已 clear 样本; -- 自动重启;第一版失败后人工停止整条 pipeline并使用新 `run_id` 重跑; -- Consumer 从最新 checkpoint 自动恢复; -- checkpoint 成功后再 clear 的严格提交窗口; -- TQ owner/storage 整体丢失后的数据重建; -- 多个独立 Consumer 竞争同一 partition; -- lease、ack、超时回收和 exactly-once; -- 动态扩缩容; -- vLLM server 直接写 TQ。 - -第一版先保证:同一个 TQ 能连通、Producer 能并发生产、Consumer 能正确取数训练、每步成功后能及时清理。 diff --git a/docs/transferqueue_integration_plan.md b/docs/transferqueue_integration_plan.md deleted file mode 100644 index cac16ac1..00000000 --- a/docs/transferqueue_integration_plan.md +++ /dev/null @@ -1,145 +0,0 @@ -# verl-SpeCo TransferQueue 落地方案 - -Last updated: 08/21/2026 - -> 目标:在**不修改上游 verl**的前提下,把 SpeCo online 训练里的逐样本特征流 -> 从「`SpecoRayPPOTrainer` driver 中转 + Ray object store」改为「TransferQueue -> 直传」,干掉 driver 这个数据瓶颈,并解锁流式消费与跨副本负载均衡。 -> -> 约束:仅 hook,与 SpeCo 现有 hook 模式一致;TQ 作为独立库使用,**不复用** verl -> 的 `main_ppo_sync` TQ 集成。 - ---- - -## 0. 现状:SpeCo online 特征流的 controller 瓶颈 - -SpeCo online 路径以 `SpecoRayPPOTrainer` 为 hub,所有跨进程大张量都被 driver 串行 -中转,介质是 Ray object store(`ray.put`/`ray.get`/`parallel_put`)。这与 verl 引入 -TQ 想解决的痛点 1:1 对应,只是 verl 干掉的是 `RayPPOTrainer`,我们要干掉的是 -SpeCo 在它之上加的 drafter 管线中转。 - -| # | 流向 | 当前机制 | 是否逐样本 | hook 位置(SpeCo 侧) | -|---|---|---|---|---| -| **a1** | target hidden states(SGLang 采集)-> drafter | `drafter_sample` 塞进 `DataProto.non_tensor_batch` -> driver pop/bucket -> `parallel_put` -> drafter `ray.get` | ✅ | `speco_ray_trainer.py` `generate_sequences_with_speco`;`sglang_adapter.py` `pop_drafter_samples`/`bucket_drafter_samples_by_replica`;`sglang_runtime.py` 组装 `drafter_sample` | -| **a2** | target hidden states(old-logprob hook)-> drafter | actor 前向 hook 截行 -> `ray.put` chunk -> driver 重打包 -> 分发 | ✅ | `oldlogprob_runtime.py` `_install_oldlogprob_hidden_hooks`/`_put_oldlogprob_hidden_refs`;`speco_ray_trainer.py` `_speco_collect_oldlogprob_features` | -| **b2** | target top-logprobs -> drafter(`use_logits=true`) | 随 a1 同一 side-channel | ✅ | `sglang_runtime.py` `target_logprobs`/`hidden_raw_target_logprobs` | -| **d** | rollout tokens -> drafter 训练集 | 随 a1 同一 side-channel(online)/`torch.save` 分片(offline) | ✅ | `sglang_runtime.py`;`speco_worker.py` `_store_rollout_sample` | -| b1 | target **lm_head 权重**(行)-> drafter `TargetHead` | ONE_TO_ALL Ray 分发 | ❌ 参数广播 | `rollout_publish.py` `export_actor_lm_head_weight`/`get_actor_lm_head_weight`;`speco_ray_trainer.py` `_speco_sync_target_lm_head_weight` | -| c | drafter 权重 -> rollout 引擎 | `ray.put` -> driver -> actor;vLLM 末段 ZMQ+SHM,SGLang 进程内 | ❌ 参数广播 | `speco_worker.py` `maybe_publish`;`rollout_publish.py` `update_draft_weights`;`vllm_runtime.py` `BucketedWeightSender` | - -**关键事实**:hidden states 跨进程前一律 CPU 物化(`oldlogprob_runtime.py`、 -`sglang_runtime.py`、`feature_store.py` 均 `.cpu()`),a1 路径下 driver 进程的 -host memory 会真正承载整批 hidden states 并做一次 Ray store 往返。这正是 TQ 要 -消除的往返。 - ---- - -## 1. 为什么不把"替换 feature_store"作为第一刀 - -`TorchShardFeatureStore`(`feature_store.py`)是 `torch.save` 分片 + JSONL manifest, -**只服务于 `collect_only`/`offline`**,不参与 online 热路径。替换它能统一离线存储 -抽象、换更快的分布式后端,但**不解决 controller 瓶颈**,性能收益有限。降级为 -可选尾项(见 §5 P3)。 - ---- - -## 2. 目标方案:TQ 直传逐样本特征流(a1 / a2 / b2 / d) - -### 2.1 角色映射 - -| TQ 角色 | SpeCo 对应 | -|---|---| -| Producer(写) | rollout worker(SGLang 路径,a1/b2/d)/ actor worker(old-logprob 路径,a2)——均在 SpeCo 既有 hook 内 | -| Consumer(读) | drafter worker `collect_rollout_features`(SpeCo 侧) | -| TransferQueueController(control plane) | SpeCo launcher 启动一个 Ray actor;drafter 经 `Sampler`/`StreamingDataLoader` 拉取 | -| Storage backend | `SimpleStorage`(ZMQ,跨节点 CPU 内存);进阶可切 `MooncakeStore`(RDMA,GPU-DRAM) | - -### 2.2 partition / key / 字段设计 - -- `partition_id`:`speco_train`(验证集用 `speco_val`)。 -- `key`:`{uid}_{session_id}_{index}`,与 verl TQ 一致;`uid` SpeCo 已有。 -- `tags`:`global_steps`、`source`∈{`rollout`,`oldlogprob`}、`replica_rank`/`owner_rank`、`status`、`prompt_len`/`response_len`/`seq_len`。ReplayBuffer/负载均衡按 tag 匹配。 -- `fields`(列):`input_ids`、`loss_mask`、`position_ids`、`hidden_states`、`last_hidden_states`/`target`、`target_logprobs`、`hidden_positions`、`prompts`、`responses`。与 `DraftFeatureSample`(`feature_store.py`)字段对齐,便于 online/offline 复用。 - -### 2.3 数据流(目标) - -``` -rollout/actor worker (SpeCo hook) - │ 生成/截取 hidden states 后,就地 tq.kv_batch_put(samples) - ▼ -TransferQueue (SimpleStorage, 跨节点 CPU 内存;可选 MooncakeStore RDMA) - │ control plane 按 sample 粒度追踪 ready 状态,Sampler 跨 drafter 副本均衡 - ▼ -drafter worker - │ tq.kv_batch_get / StreamingDataLoader 消费 → 喂入既有 DataBuffer / collect_online_data - ▼ -drafter 训练 (不变) -``` - -driver 只下发触发与轻量 key/meta,**不再承载 hidden states**。 - ---- - -## 3. 落地改动点(全部在 SpeCo 侧,hook-only) - -### 3.1 启动与配置 -- `draft_train_launcher.py` / `main.py`:`tq.init(config.transfer_queue)`;起 `TransferQueueController.remote(Sampler)`。 -- `config/speco_base.yaml`:新增 `drafter.transfer_queue` 块(backend、partition、enable 开关)。参考 verl `ppo_trainer.yaml` 的 `transfer_queue:` 结构,但**独立配置**,不复用 verl 的。 - -### 3.2 Producer 侧 -- **a1/b2/d(SGLang)**:`sglang_runtime.py` 组装 `drafter_sample` 处(~1594-1648),增加 `tq.kv_batch_put`;返回给 driver 的 `drafter_sample` 只保留 key/meta(或整段不再走 DataProto side-channel,driver 仅触发)。 -- **a2(old-logprob)**:`oldlogprob_runtime.py` `_put_oldlogprob_hidden_refs`(~216),把 `ray.put(hidden_chunk)` 换成 `tq.kv_batch_put`;`OLD_LOGPROB_HIDDEN_CHUNK_REFS_KEY` 改为 TQ key 列表。 - -### 3.3 Consumer 侧 -- `speco_worker.py` `collect_rollout_features`(~665):把 `_resolve_ray_object_ref`/`_resolve_hidden_state_chunks`(`ray.get`)换成 `tq.kv_batch_get`;`_dispatch_nd_compute`(~159)的 `parallel_put` 退化为只传 key(或 drafter 直接从 TQ Sampler 拉,driver 不参与分发)。 -- drafter 内部 `DataBuffer`/`collect_online_data`(`base_trainer.py`)保持不变,只是数据来源由 `ray.get` 改为 TQ get。 - -### 3.4 Driver 侧 -- `speco_ray_trainer.py`:`_speco_collect_rollout_features_rpc`/`speco_collect_rollout_features`(~351)、`_speco_collect_oldlogprob_features`(~1114)不再搬数据,只做触发/传 key;`bucket_drafter_samples_by_replica` 可由 TQ `Sampler` 替代(逐步迁移,先保留作回退)。 - -### 3.5 不改动 -- **b1(lm_head 权重)、c(drafter 权重)**:保持现状。与 verl 上游一致(权重不走 TQ),且 c 的 vLLM 末段已有专用 ZMQ+SHM 通道。 -- verl 本体:零改动。 - ---- - -## 4. 收益与边界(诚实评估) - -### 收益 -1. **去掉 driver 对 hidden states 的 host-memory 中转 + Ray store 往返**:producer 直存 TQ,consumer 直取,driver 不再承载整批特征。 -2. **流式消费**:drafter 在样本 ready 时即可消费,不必等整批 `generate_sequences` 返回,采集与训练可重叠。 -3. **跨 drafter 副本负载均衡**:TQ `Sampler`/`RankAwareSampler` 替代手写 `bucket_drafter_samples_by_replica`/`owner_rank` 分配。 -4. **(若采纳 P3)统一 online/collect_only/offline 存储**:同一 TQ partition,`collect_only` 写、`offline` 读,消掉 on-disk 分片层。 - -### 边界 / 不解决的事 -- 只优化**特征采集**这一子阶段,**不加速** rollout 本身、actor update、reward;e2e 增益取决于该子阶段在 step 中的占比。 SpeCo README 的 20% rollout / 11% e2e 提升来自 acceptance length,与本方案是不同机制,不要混为一谈。 -- **权重同步(b1/c)不放进 TQ**,与 verl 上游保持一致。 -- hidden states 跨进程前**仍需 CPU 物化**(现状如此);要避免物化需切 `MooncakeStore` RDMA,属进阶项。 -- 引入 TQ 依赖与一个 control-plane Ray actor,增加少量运维面。 - -### 风险 -- TQ 与 SpeCo 现有 `owner_rank`/`replica_rank` 路由语义需对齐(Sampler 要复刻「按 owner 分桶」语义,否则样本会错配 drafter 副本)。 -- old-logprob 的 chunk 拆分(`hidden_states_ref_chunks`)映射到 TQ 列式存储时,需保证 chunk meta 与 key 的一致性。 -- 回退路径:保留 `enable_transfer_queue=False` 时走原 Ray 路径,渐进切换。 - ---- - -## 5. 分阶段实施 - -| 阶段 | 范围 | 产出 | -|---|---|---| -| **P0** | a1(SGLang hidden states)走 TQ 直传;drafter `kv_batch_get` 消费;driver 仅触发 | 验证 controller-bypass 闭环 + 正确性 | -| **P1** | a2(old-logprob hidden states)走 TQ;chunk 拆分映射 TQ 列 | 覆盖第二条采集路径 | -| **P2** | b2(top-logprobs)+ d(tokens)随 a1 同 partition 传输;Sampler 替代手写 bucket | 完整特征流 + 跨副本均衡 | -| **P3(可选)** | `TorchShardFeatureStore` → TQ partition,统一 online/collect_only/offline | 离线工作流统一 | - -每个阶段保留 `enable_transfer_queue` 开关与原 Ray 路径回退。 - ---- - -## 6. 待确认决策 - -1. **TQ backend**:`SimpleStorage`(CPU 内存,默认)起步,还是直接上 `MooncakeStore`(RDMA,省 CPU 物化)?后者依赖 RDMA 网络,建议 P0 用 SimpleStorage。 -2. **drafter 消费模式**:`kv_batch_get`(主动拉,改动小)还是 `StreamingDataLoader`(全自动流式,改动大、收益高)?建议 P0 用前者,P2 再考虑后者。 -3. **driver 角色**:P0 先保留 driver 传 key(最小改动),还是直接让 drafter 从 TQ Sampler 自取(driver 彻底退出数据路径)?前者风险低,建议 P0 用前者。 -4. **是否做 P3**:离线统一是否在本次范围内,还是单独立项。 From 960d6728f17b541b62ec250647bab2ddee251102 Mon Sep 17 00:00:00 2001 From: xxyyrr598 <1615681145@qq.com> Date: Thu, 27 Aug 2026 15:51:25 +0800 Subject: [PATCH 44/50] refactor(tq): drop mooncake transfer path and consolidate dspark separate training script --- README.md | 86 +---- ...en3-8b_drafter_dspark_separate_training.sh | 150 +++++++- .../run_qwen3-8b_drafter_separate_training.sh | 165 --------- tests/config/test_speco_config_overlay.py | 4 - tests/examples/test_example_scripts.py | 10 +- .../test_check_example_naming.py | 2 +- tests/unit/test_mooncake_transfer.py | 73 ---- tests/unit/test_target_feature_pipeline.py | 2 - verl_speco/config/draft_trainer.yaml | 17 +- verl_speco/config/speco_base.yaml | 9 - .../mooncake_hidden_states_connector.py | 327 ------------------ verl_speco/trainer/draft_training_loop.py | 14 +- verl_speco/trainer/mooncake_transfer.py | 214 ------------ verl_speco/trainer/target_feature_pipeline.py | 44 +-- verl_speco/trainer/target_feature_replay.py | 143 +------- 15 files changed, 163 insertions(+), 1097 deletions(-) delete mode 100644 examples/run_qwen3-8b_drafter_separate_training.sh delete mode 100644 tests/unit/test_mooncake_transfer.py delete mode 100644 verl_speco/integration/mooncake_hidden_states_connector.py delete mode 100644 verl_speco/trainer/mooncake_transfer.py diff --git a/README.md b/README.md index 371c8dff..bae7873a 100644 --- a/README.md +++ b/README.md @@ -305,7 +305,7 @@ TransferQueue, and a consumer trains the drafter independently of PPO. Quickstart: ```bash -bash examples/run_qwen3-8b_drafter_separate_training.sh +bash examples/run_qwen3-8b_drafter_dspark_separate_training.sh ``` Set the same model, dataset, drafter, checkpoint, GPU, and optimization values @@ -355,90 +355,6 @@ python -m verl_speco.inspect_feature_store /path/to/features \ --strict-exit ``` -### Producer/Mooncake pipeline - -For token-only stores, standalone training can overlap target-model inference, -hidden-state transfer, and drafter optimization. vLLM acts as the producer, -Mooncake transfers the extracted tensors without per-sample files, and a -bounded rank-local queue prefetches complete batches while FSDP trains the -current batch. This path is disabled by default and does not change online -PPO/drafter training. - -Install the Mooncake package for the target hardware and start a Mooncake -master. For example, use `mooncake-transfer-engine` on CUDA or -`mooncake-transfer-engine-npu` on Ascend. Export the same connection settings -in the vLLM and training processes. The connector requires vLLM 0.23 or newer: - -```bash -# CUDA; use mooncake-transfer-engine-npu instead on Ascend. -pip install 'mooncake-transfer-engine>=0.3.10.post1' - -export MOONCAKE_MASTER_SERVER=127.0.0.1:50051 -export MOONCAKE_METADATA_SERVER=http://127.0.0.1:8090/metadata -export MOONCAKE_PROTOCOL=tcp -export MOONCAKE_GLOBAL_SEGMENT_SIZE=$((16 * 1024 * 1024 * 1024)) -export MOONCAKE_LOCAL_BUFFER_SIZE=$((2 * 1024 * 1024 * 1024)) - -mooncake_master \ - --enable_http_metadata_server=true \ - --http_metadata_server_host=0.0.0.0 \ - --http_metadata_server_port=8090 -``` - -Start vLLM with SpeCo's store-only hidden-state connector. The final layer must -be appended to the auxiliary layer list when the selected drafter loss needs -last-hidden-state supervision: - -```bash -vllm serve /path/to/target_model \ - --port 8000 \ - --tensor-parallel-size 2 \ - --speculative-config '{"method":"extract_hidden_states","num_speculative_tokens":1,"draft_model_config":{"hf_config":{"eagle_aux_hidden_state_layer_ids":[1,9,17,25,33,36]}}}' \ - --kv-transfer-config '{"kv_connector":"SpeCoMooncakeHiddenStatesConnector","kv_connector_module_path":"verl_speco.integration.mooncake_hidden_states_connector","kv_role":"kv_producer"}' \ - --no-enable-chunked-prefill -``` - -Enable the pipeline in the standalone command: - -```bash -python -m verl_speco.draft_train_launcher \ - speco.draft_training.nproc_per_node=4 \ - actor_rollout_ref.model.path=/path/to/target_model \ - actor_rollout_ref.rollout.drafter.speculative_algorithm=DSPARK \ - actor_rollout_ref.rollout.drafter.training.mode=offline \ - actor_rollout_ref.rollout.drafter.training.feature_store.type=jsonl_token_replay \ - actor_rollout_ref.rollout.drafter.training.feature_store.path=/path/to/data.jsonl \ - actor_rollout_ref.rollout.drafter.training.target_feature_replay.backend=vllm_mooncake \ - actor_rollout_ref.rollout.drafter.training.target_feature_replay.vllm_endpoint=http://127.0.0.1:8000/v1 \ - actor_rollout_ref.rollout.drafter.training.target_feature_pipeline.enabled=true \ - actor_rollout_ref.rollout.drafter.training.target_feature_pipeline.concurrency=16 \ - actor_rollout_ref.rollout.drafter.training.target_feature_pipeline.transfer_concurrency=8 \ - actor_rollout_ref.rollout.drafter.training.target_feature_pipeline.producer_prefetch_depth=4 \ - actor_rollout_ref.rollout.drafter.training.target_feature_pipeline.prefetch_depth=2 -``` - -To spread replay requests across independent vLLM deployments, replace the -single endpoint override with an endpoint pool: - -```bash -'actor_rollout_ref.rollout.drafter.training.target_feature_replay.vllm_endpoints=["http://host1:8000/v1","http://host2:8000/v1"]' \ -actor_rollout_ref.rollout.drafter.training.target_feature_replay.endpoint_cooldown=5 -``` - -The standalone producer routes each request to the healthy endpoint with the -fewest in-flight requests. A failed request is retried on another endpoint when -the configured retry budget permits it. The original `vllm_endpoint` option -remains supported and is used when `vllm_endpoints` is unset. - -`concurrency` and `transfer_concurrency` are global budgets divided across all -training ranks. Start with request concurrency between 16 and 32, then increase -only while vLLM throughput rises. -`producer_prefetch_depth` bounds batches that have outstanding HTTP work; -`transfer_concurrency` bounds simultaneous Mooncake GETs. A -`prefetch_depth` of 2 normally hides transfer latency without retaining too -many large hidden-state batches. The standalone metrics include producer queue -depth, consumer wait time, vLLM request time, and Mooncake GET time. - ## Configuration SPECO-specific options live under: diff --git a/examples/run_qwen3-8b_drafter_dspark_separate_training.sh b/examples/run_qwen3-8b_drafter_dspark_separate_training.sh index 2cd1e8eb..1a08a43b 100644 --- a/examples/run_qwen3-8b_drafter_dspark_separate_training.sh +++ b/examples/run_qwen3-8b_drafter_dspark_separate_training.sh @@ -13,9 +13,153 @@ # See the License for the specific language governing permissions and # limitations under the License. set -euo pipefail +set -x -# Backward-compatible DSpark entry point. Both example names run the same -# standalone vLLM -> TransferQueue/Mooncake -> DSpark training pipeline. script_dir=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd) +repo_root=$(cd -- "${script_dir}/.." && pwd) +cd "${repo_root}" -exec bash "${script_dir}/run_qwen3-8b_drafter_separate_training.sh" "$@" +# Standalone DSpark draft-model training using an already-running hidden-state +# vLLM. Start tools/run_qwen3-8b_drafter_hidden_state_vllm.sh in another terminal +# first. This process owns Ray/TQ, Producer and Consumer, but it must not own +# the target vLLM so inference and training can use different accelerators. + +project_name=${PROJECT_NAME:-verl_dspark_drafter} +exp_name=${EXP_NAME:-qwen3_8b_dspark_separate_training} + +draft_train_gpus_per_node=${TRAIN_GPUS:-2} + +MODEL_PATH=${MODEL_PATH:-/path/to/Qwen3-8B} +# Ordinary verl prompt Parquet is supported; target vLLM generates responses. +TRAIN_FILE=${TRAIN_FILE:-/path/to/train_file.parquet} +# Optional. Leave empty to initialize DSpark from the target-model/config +# fallback; set it only when loading or resuming an existing drafter. +DRAFTER_PATH=${DRAFTER_PATH:-} +DRAFT_CKPTS_DIR=${DRAFT_CKPTS_DIR:-/path/to/dspark_draft_checkpoints} + +PYTHON_BIN=${PYTHON_BIN:-python3} +DEVICE_ENV=${DEVICE_ENV:-ASCEND_RT_VISIBLE_DEVICES} +TRAIN_DEVICES=${TRAIN_DEVICES:-2,3} +SPECO_VLLM_ENDPOINTS=${SPECO_VLLM_ENDPOINTS:-'[http://127.0.0.1:8000/v1,http://127.0.0.1:8001/v1]'} +VLLM_READY_TIMEOUT_SECONDS=${VLLM_READY_TIMEOUT_SECONDS:-120} + +# Producer -> vLLM concurrency and bounded queues. MAX_INFLIGHT_REQUESTS is the +# process-wide request limit; PER_ENDPOINT_CONCURRENCY applies independently to +# every URL in SPECO_VLLM_ENDPOINTS. +VLLM_REQUEST_TIMEOUT=${VLLM_REQUEST_TIMEOUT:-120} +VLLM_MAX_INFLIGHT_REQUESTS=${VLLM_MAX_INFLIGHT_REQUESTS:-16} +VLLM_PER_ENDPOINT_CONCURRENCY=${VLLM_PER_ENDPOINT_CONCURRENCY:-4} +PRODUCER_INPUT_QUEUE_SIZE=${PRODUCER_INPUT_QUEUE_SIZE:-32} +PRODUCER_PUBLISH_QUEUE_SIZE=${PRODUCER_PUBLISH_QUEUE_SIZE:-16} +PRODUCER_MAX_PENDING_SAMPLES=${PRODUCER_MAX_PENDING_SAMPLES:-1024} +PRODUCER_PENDING_POLL_INTERVAL=${PRODUCER_PENDING_POLL_INTERVAL:-0.5} +PRODUCER_MAX_SEQUENCE_LENGTH=${PRODUCER_MAX_SEQUENCE_LENGTH:-8192} +PRODUCER_MAX_FEATURE_LENGTH=${PRODUCER_MAX_FEATURE_LENGTH:-512} +PRODUCER_GENERATION_MAX_TOKENS=${PRODUCER_GENERATION_MAX_TOKENS:-512} + +# Standalone trainer. +MAX_STEPS=${MAX_STEPS:-10} +SAVE_INTERVAL_STEPS=${SAVE_INTERVAL_STEPS:-5} +SAVE_FINAL_CHECKPOINT=${SAVE_FINAL_CHECKPOINT:-true} +BATCH_SIZE_PER_GPU=${BATCH_SIZE_PER_GPU:-2} +LEARNING_RATE=${LEARNING_RATE:-1e-6} +LR_WARMUP_STEPS=${LR_WARMUP_STEPS:-0} +LR_SCHEDULER_TYPE=${LR_SCHEDULER_TYPE:-constant} +LR_DECAY_STEPS=${LR_DECAY_STEPS:-100} +MIN_LR_RATIO=${MIN_LR_RATIO:-0.1} +PARAM_OFFLOAD=${PARAM_OFFLOAD:-true} +OPTIMIZER_OFFLOAD=${OPTIMIZER_OFFLOAD:-true} + +# DSpark architecture, sampling and losses. TARGET_LAYER_IDS must match the +# auxiliary layers exposed by both hidden-state vLLM services. +DSPARK_BLOCK_SIZE=${DSPARK_BLOCK_SIZE:-7} +DSPARK_NUM_ANCHORS=${DSPARK_NUM_ANCHORS:-32} +DSPARK_MAX_WINDOW=${DSPARK_MAX_WINDOW:-512} +DSPARK_LOSS_MODE=${DSPARK_LOSS_MODE:-full_vocab} +DSPARK_SAMPLED_CE_NEGATIVES=${DSPARK_SAMPLED_CE_NEGATIVES:-0} +DSPARK_LOSS_DECAY_GAMMA=${DSPARK_LOSS_DECAY_GAMMA:-7} +DSPARK_NUM_TARGET_LAYERS=${DSPARK_NUM_TARGET_LAYERS:-5} +DSPARK_NUM_HIDDEN_LAYERS=${DSPARK_NUM_HIDDEN_LAYERS:-5} +DSPARK_TARGET_LAYER_IDS=${DSPARK_TARGET_LAYER_IDS:-'[1,9,17,25,33]'} +DSPARK_MARKOV_RANK=${DSPARK_MARKOV_RANK:-256} +DSPARK_MARKOV_HEAD_TYPE=${DSPARK_MARKOV_HEAD_TYPE:-vanilla} +DSPARK_CE_LOSS_ALPHA=${DSPARK_CE_LOSS_ALPHA:-0.1} +DSPARK_L1_LOSS_ALPHA=${DSPARK_L1_LOSS_ALPHA:-0.45} +DSPARK_L1_CHUNK_SIZE=${DSPARK_L1_CHUNK_SIZE:-0} +# The current DSpark trainer rejects nonzero confidence loss because target +# acceptance labels are not part of the standalone feature protocol yet. +DSPARK_CONFIDENCE_LOSS_ALPHA=${DSPARK_CONFIDENCE_LOSS_ALPHA:-0.0} +DSPARK_DEBUG_LOG=${DSPARK_DEBUG_LOG:-false} +DSPARK_DEBUG_LOG_FIRST_N=${DSPARK_DEBUG_LOG_FIRST_N:-2} +DSPARK_DEBUG_LOG_INTERVAL=${DSPARK_DEBUG_LOG_INTERVAL:-100} + +export "${DEVICE_ENV}=${TRAIN_DEVICES}" +export SPECO_VLLM_ENDPOINTS + +# Fail before entering the unified launcher when the separately managed vLLM +# is absent. Otherwise a localhost endpoint would make the launcher start its +# fallback vLLM inside the training process and on the training devices. +if ! "${PYTHON_BIN}" tools/wait_for_vllm_endpoints.py \ + --endpoints "${SPECO_VLLM_ENDPOINTS}" \ + --timeout-seconds "${VLLM_READY_TIMEOUT_SECONDS}"; then + echo "Start tools/run_qwen3-8b_drafter_hidden_state_vllm.sh first" >&2 + exit 1 +fi + +PYTHONUNBUFFERED=1 "${PYTHON_BIN}" -m verl_speco.standalone_tq_training_launcher \ + speco.draft_training.num_gpus_per_node=${draft_train_gpus_per_node} \ + speco.draft_training.nnodes=1 \ + speco.draft_training.standalone=True \ + data.train_files=${TRAIN_FILE} \ + actor_rollout_ref.model.path=${MODEL_PATH} \ + actor_rollout_ref.actor.strategy=fsdp2 \ + actor_rollout_ref.actor.fsdp_config.param_offload=${PARAM_OFFLOAD} \ + actor_rollout_ref.actor.fsdp_config.optimizer_offload=${OPTIMIZER_OFFLOAD} \ + actor_rollout_ref.rollout.tensor_model_parallel_size=1 \ + actor_rollout_ref.rollout.drafter.enable=True \ + actor_rollout_ref.rollout.drafter.enable_drafter_training=True \ + actor_rollout_ref.rollout.drafter.model_path=${DRAFTER_PATH} \ + actor_rollout_ref.rollout.drafter.checkpoint_path=${DRAFT_CKPTS_DIR} \ + actor_rollout_ref.rollout.drafter.speculative_algorithm=DSPARK \ + actor_rollout_ref.rollout.drafter.training.mode=offline \ + actor_rollout_ref.rollout.drafter.training.max_steps=${MAX_STEPS} \ + actor_rollout_ref.rollout.drafter.training.save_interval_steps=${SAVE_INTERVAL_STEPS} \ + actor_rollout_ref.rollout.drafter.training.save_final_checkpoint=${SAVE_FINAL_CHECKPOINT} \ + actor_rollout_ref.rollout.drafter.training.batch_size_per_gpu=${BATCH_SIZE_PER_GPU} \ + actor_rollout_ref.rollout.drafter.training.lr=${LEARNING_RATE} \ + actor_rollout_ref.rollout.drafter.training.lr_warmup_steps=${LR_WARMUP_STEPS} \ + actor_rollout_ref.rollout.drafter.training.lr_scheduler_type=${LR_SCHEDULER_TYPE} \ + actor_rollout_ref.rollout.drafter.training.lr_decay_steps=${LR_DECAY_STEPS} \ + actor_rollout_ref.rollout.drafter.training.min_lr_ratio=${MIN_LR_RATIO} \ + actor_rollout_ref.rollout.drafter.training.use_logits=False \ + actor_rollout_ref.rollout.drafter.training.dspark_block_size=${DSPARK_BLOCK_SIZE} \ + actor_rollout_ref.rollout.drafter.training.dspark_num_anchors=${DSPARK_NUM_ANCHORS} \ + actor_rollout_ref.rollout.drafter.training.dspark_max_window=${DSPARK_MAX_WINDOW} \ + actor_rollout_ref.rollout.drafter.training.dspark_loss_mode=${DSPARK_LOSS_MODE} \ + actor_rollout_ref.rollout.drafter.training.dspark_sampled_ce_negatives=${DSPARK_SAMPLED_CE_NEGATIVES} \ + actor_rollout_ref.rollout.drafter.training.dspark_loss_decay_gamma=${DSPARK_LOSS_DECAY_GAMMA} \ + actor_rollout_ref.rollout.drafter.training.dspark_num_target_layers=${DSPARK_NUM_TARGET_LAYERS} \ + actor_rollout_ref.rollout.drafter.training.dspark_num_hidden_layers=${DSPARK_NUM_HIDDEN_LAYERS} \ + actor_rollout_ref.rollout.drafter.training.dspark_target_layer_ids=${DSPARK_TARGET_LAYER_IDS} \ + actor_rollout_ref.rollout.drafter.training.dspark_markov_rank=${DSPARK_MARKOV_RANK} \ + actor_rollout_ref.rollout.drafter.training.dspark_markov_head_type=${DSPARK_MARKOV_HEAD_TYPE} \ + actor_rollout_ref.rollout.drafter.training.dspark_ce_loss_alpha=${DSPARK_CE_LOSS_ALPHA} \ + actor_rollout_ref.rollout.drafter.training.dspark_l1_loss_alpha=${DSPARK_L1_LOSS_ALPHA} \ + actor_rollout_ref.rollout.drafter.training.dspark_l1_chunk_size=${DSPARK_L1_CHUNK_SIZE} \ + actor_rollout_ref.rollout.drafter.training.dspark_confidence_loss_alpha=${DSPARK_CONFIDENCE_LOSS_ALPHA} \ + actor_rollout_ref.rollout.drafter.training.dspark_debug_log=${DSPARK_DEBUG_LOG} \ + actor_rollout_ref.rollout.drafter.training.dspark_debug_log_first_n=${DSPARK_DEBUG_LOG_FIRST_N} \ + actor_rollout_ref.rollout.drafter.training.dspark_debug_log_interval=${DSPARK_DEBUG_LOG_INTERVAL} \ + speco.standalone_tq_producer.request_timeout=${VLLM_REQUEST_TIMEOUT} \ + speco.standalone_tq_producer.max_inflight_requests=${VLLM_MAX_INFLIGHT_REQUESTS} \ + speco.standalone_tq_producer.per_endpoint_concurrency=${VLLM_PER_ENDPOINT_CONCURRENCY} \ + speco.standalone_tq_producer.input_queue_size=${PRODUCER_INPUT_QUEUE_SIZE} \ + speco.standalone_tq_producer.publish_queue_size=${PRODUCER_PUBLISH_QUEUE_SIZE} \ + speco.standalone_tq_producer.max_pending_samples=${PRODUCER_MAX_PENDING_SAMPLES} \ + speco.standalone_tq_producer.pending_poll_interval_seconds=${PRODUCER_PENDING_POLL_INTERVAL} \ + speco.standalone_tq_producer.max_sequence_length=${PRODUCER_MAX_SEQUENCE_LENGTH} \ + speco.standalone_tq_producer.max_feature_length=${PRODUCER_MAX_FEATURE_LENGTH} \ + speco.standalone_tq_producer.generation_max_tokens=${PRODUCER_GENERATION_MAX_TOKENS} \ + trainer.project_name=${project_name} \ + trainer.experiment_name=${exp_name} \ + "$@" diff --git a/examples/run_qwen3-8b_drafter_separate_training.sh b/examples/run_qwen3-8b_drafter_separate_training.sh deleted file mode 100644 index 1a08a43b..00000000 --- a/examples/run_qwen3-8b_drafter_separate_training.sh +++ /dev/null @@ -1,165 +0,0 @@ -#!/usr/bin/env bash -# Copyright 2026 Bytedance Ltd. and/or its affiliates -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -set -euo pipefail -set -x - -script_dir=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd) -repo_root=$(cd -- "${script_dir}/.." && pwd) -cd "${repo_root}" - -# Standalone DSpark draft-model training using an already-running hidden-state -# vLLM. Start tools/run_qwen3-8b_drafter_hidden_state_vllm.sh in another terminal -# first. This process owns Ray/TQ, Producer and Consumer, but it must not own -# the target vLLM so inference and training can use different accelerators. - -project_name=${PROJECT_NAME:-verl_dspark_drafter} -exp_name=${EXP_NAME:-qwen3_8b_dspark_separate_training} - -draft_train_gpus_per_node=${TRAIN_GPUS:-2} - -MODEL_PATH=${MODEL_PATH:-/path/to/Qwen3-8B} -# Ordinary verl prompt Parquet is supported; target vLLM generates responses. -TRAIN_FILE=${TRAIN_FILE:-/path/to/train_file.parquet} -# Optional. Leave empty to initialize DSpark from the target-model/config -# fallback; set it only when loading or resuming an existing drafter. -DRAFTER_PATH=${DRAFTER_PATH:-} -DRAFT_CKPTS_DIR=${DRAFT_CKPTS_DIR:-/path/to/dspark_draft_checkpoints} - -PYTHON_BIN=${PYTHON_BIN:-python3} -DEVICE_ENV=${DEVICE_ENV:-ASCEND_RT_VISIBLE_DEVICES} -TRAIN_DEVICES=${TRAIN_DEVICES:-2,3} -SPECO_VLLM_ENDPOINTS=${SPECO_VLLM_ENDPOINTS:-'[http://127.0.0.1:8000/v1,http://127.0.0.1:8001/v1]'} -VLLM_READY_TIMEOUT_SECONDS=${VLLM_READY_TIMEOUT_SECONDS:-120} - -# Producer -> vLLM concurrency and bounded queues. MAX_INFLIGHT_REQUESTS is the -# process-wide request limit; PER_ENDPOINT_CONCURRENCY applies independently to -# every URL in SPECO_VLLM_ENDPOINTS. -VLLM_REQUEST_TIMEOUT=${VLLM_REQUEST_TIMEOUT:-120} -VLLM_MAX_INFLIGHT_REQUESTS=${VLLM_MAX_INFLIGHT_REQUESTS:-16} -VLLM_PER_ENDPOINT_CONCURRENCY=${VLLM_PER_ENDPOINT_CONCURRENCY:-4} -PRODUCER_INPUT_QUEUE_SIZE=${PRODUCER_INPUT_QUEUE_SIZE:-32} -PRODUCER_PUBLISH_QUEUE_SIZE=${PRODUCER_PUBLISH_QUEUE_SIZE:-16} -PRODUCER_MAX_PENDING_SAMPLES=${PRODUCER_MAX_PENDING_SAMPLES:-1024} -PRODUCER_PENDING_POLL_INTERVAL=${PRODUCER_PENDING_POLL_INTERVAL:-0.5} -PRODUCER_MAX_SEQUENCE_LENGTH=${PRODUCER_MAX_SEQUENCE_LENGTH:-8192} -PRODUCER_MAX_FEATURE_LENGTH=${PRODUCER_MAX_FEATURE_LENGTH:-512} -PRODUCER_GENERATION_MAX_TOKENS=${PRODUCER_GENERATION_MAX_TOKENS:-512} - -# Standalone trainer. -MAX_STEPS=${MAX_STEPS:-10} -SAVE_INTERVAL_STEPS=${SAVE_INTERVAL_STEPS:-5} -SAVE_FINAL_CHECKPOINT=${SAVE_FINAL_CHECKPOINT:-true} -BATCH_SIZE_PER_GPU=${BATCH_SIZE_PER_GPU:-2} -LEARNING_RATE=${LEARNING_RATE:-1e-6} -LR_WARMUP_STEPS=${LR_WARMUP_STEPS:-0} -LR_SCHEDULER_TYPE=${LR_SCHEDULER_TYPE:-constant} -LR_DECAY_STEPS=${LR_DECAY_STEPS:-100} -MIN_LR_RATIO=${MIN_LR_RATIO:-0.1} -PARAM_OFFLOAD=${PARAM_OFFLOAD:-true} -OPTIMIZER_OFFLOAD=${OPTIMIZER_OFFLOAD:-true} - -# DSpark architecture, sampling and losses. TARGET_LAYER_IDS must match the -# auxiliary layers exposed by both hidden-state vLLM services. -DSPARK_BLOCK_SIZE=${DSPARK_BLOCK_SIZE:-7} -DSPARK_NUM_ANCHORS=${DSPARK_NUM_ANCHORS:-32} -DSPARK_MAX_WINDOW=${DSPARK_MAX_WINDOW:-512} -DSPARK_LOSS_MODE=${DSPARK_LOSS_MODE:-full_vocab} -DSPARK_SAMPLED_CE_NEGATIVES=${DSPARK_SAMPLED_CE_NEGATIVES:-0} -DSPARK_LOSS_DECAY_GAMMA=${DSPARK_LOSS_DECAY_GAMMA:-7} -DSPARK_NUM_TARGET_LAYERS=${DSPARK_NUM_TARGET_LAYERS:-5} -DSPARK_NUM_HIDDEN_LAYERS=${DSPARK_NUM_HIDDEN_LAYERS:-5} -DSPARK_TARGET_LAYER_IDS=${DSPARK_TARGET_LAYER_IDS:-'[1,9,17,25,33]'} -DSPARK_MARKOV_RANK=${DSPARK_MARKOV_RANK:-256} -DSPARK_MARKOV_HEAD_TYPE=${DSPARK_MARKOV_HEAD_TYPE:-vanilla} -DSPARK_CE_LOSS_ALPHA=${DSPARK_CE_LOSS_ALPHA:-0.1} -DSPARK_L1_LOSS_ALPHA=${DSPARK_L1_LOSS_ALPHA:-0.45} -DSPARK_L1_CHUNK_SIZE=${DSPARK_L1_CHUNK_SIZE:-0} -# The current DSpark trainer rejects nonzero confidence loss because target -# acceptance labels are not part of the standalone feature protocol yet. -DSPARK_CONFIDENCE_LOSS_ALPHA=${DSPARK_CONFIDENCE_LOSS_ALPHA:-0.0} -DSPARK_DEBUG_LOG=${DSPARK_DEBUG_LOG:-false} -DSPARK_DEBUG_LOG_FIRST_N=${DSPARK_DEBUG_LOG_FIRST_N:-2} -DSPARK_DEBUG_LOG_INTERVAL=${DSPARK_DEBUG_LOG_INTERVAL:-100} - -export "${DEVICE_ENV}=${TRAIN_DEVICES}" -export SPECO_VLLM_ENDPOINTS - -# Fail before entering the unified launcher when the separately managed vLLM -# is absent. Otherwise a localhost endpoint would make the launcher start its -# fallback vLLM inside the training process and on the training devices. -if ! "${PYTHON_BIN}" tools/wait_for_vllm_endpoints.py \ - --endpoints "${SPECO_VLLM_ENDPOINTS}" \ - --timeout-seconds "${VLLM_READY_TIMEOUT_SECONDS}"; then - echo "Start tools/run_qwen3-8b_drafter_hidden_state_vllm.sh first" >&2 - exit 1 -fi - -PYTHONUNBUFFERED=1 "${PYTHON_BIN}" -m verl_speco.standalone_tq_training_launcher \ - speco.draft_training.num_gpus_per_node=${draft_train_gpus_per_node} \ - speco.draft_training.nnodes=1 \ - speco.draft_training.standalone=True \ - data.train_files=${TRAIN_FILE} \ - actor_rollout_ref.model.path=${MODEL_PATH} \ - actor_rollout_ref.actor.strategy=fsdp2 \ - actor_rollout_ref.actor.fsdp_config.param_offload=${PARAM_OFFLOAD} \ - actor_rollout_ref.actor.fsdp_config.optimizer_offload=${OPTIMIZER_OFFLOAD} \ - actor_rollout_ref.rollout.tensor_model_parallel_size=1 \ - actor_rollout_ref.rollout.drafter.enable=True \ - actor_rollout_ref.rollout.drafter.enable_drafter_training=True \ - actor_rollout_ref.rollout.drafter.model_path=${DRAFTER_PATH} \ - actor_rollout_ref.rollout.drafter.checkpoint_path=${DRAFT_CKPTS_DIR} \ - actor_rollout_ref.rollout.drafter.speculative_algorithm=DSPARK \ - actor_rollout_ref.rollout.drafter.training.mode=offline \ - actor_rollout_ref.rollout.drafter.training.max_steps=${MAX_STEPS} \ - actor_rollout_ref.rollout.drafter.training.save_interval_steps=${SAVE_INTERVAL_STEPS} \ - actor_rollout_ref.rollout.drafter.training.save_final_checkpoint=${SAVE_FINAL_CHECKPOINT} \ - actor_rollout_ref.rollout.drafter.training.batch_size_per_gpu=${BATCH_SIZE_PER_GPU} \ - actor_rollout_ref.rollout.drafter.training.lr=${LEARNING_RATE} \ - actor_rollout_ref.rollout.drafter.training.lr_warmup_steps=${LR_WARMUP_STEPS} \ - actor_rollout_ref.rollout.drafter.training.lr_scheduler_type=${LR_SCHEDULER_TYPE} \ - actor_rollout_ref.rollout.drafter.training.lr_decay_steps=${LR_DECAY_STEPS} \ - actor_rollout_ref.rollout.drafter.training.min_lr_ratio=${MIN_LR_RATIO} \ - actor_rollout_ref.rollout.drafter.training.use_logits=False \ - actor_rollout_ref.rollout.drafter.training.dspark_block_size=${DSPARK_BLOCK_SIZE} \ - actor_rollout_ref.rollout.drafter.training.dspark_num_anchors=${DSPARK_NUM_ANCHORS} \ - actor_rollout_ref.rollout.drafter.training.dspark_max_window=${DSPARK_MAX_WINDOW} \ - actor_rollout_ref.rollout.drafter.training.dspark_loss_mode=${DSPARK_LOSS_MODE} \ - actor_rollout_ref.rollout.drafter.training.dspark_sampled_ce_negatives=${DSPARK_SAMPLED_CE_NEGATIVES} \ - actor_rollout_ref.rollout.drafter.training.dspark_loss_decay_gamma=${DSPARK_LOSS_DECAY_GAMMA} \ - actor_rollout_ref.rollout.drafter.training.dspark_num_target_layers=${DSPARK_NUM_TARGET_LAYERS} \ - actor_rollout_ref.rollout.drafter.training.dspark_num_hidden_layers=${DSPARK_NUM_HIDDEN_LAYERS} \ - actor_rollout_ref.rollout.drafter.training.dspark_target_layer_ids=${DSPARK_TARGET_LAYER_IDS} \ - actor_rollout_ref.rollout.drafter.training.dspark_markov_rank=${DSPARK_MARKOV_RANK} \ - actor_rollout_ref.rollout.drafter.training.dspark_markov_head_type=${DSPARK_MARKOV_HEAD_TYPE} \ - actor_rollout_ref.rollout.drafter.training.dspark_ce_loss_alpha=${DSPARK_CE_LOSS_ALPHA} \ - actor_rollout_ref.rollout.drafter.training.dspark_l1_loss_alpha=${DSPARK_L1_LOSS_ALPHA} \ - actor_rollout_ref.rollout.drafter.training.dspark_l1_chunk_size=${DSPARK_L1_CHUNK_SIZE} \ - actor_rollout_ref.rollout.drafter.training.dspark_confidence_loss_alpha=${DSPARK_CONFIDENCE_LOSS_ALPHA} \ - actor_rollout_ref.rollout.drafter.training.dspark_debug_log=${DSPARK_DEBUG_LOG} \ - actor_rollout_ref.rollout.drafter.training.dspark_debug_log_first_n=${DSPARK_DEBUG_LOG_FIRST_N} \ - actor_rollout_ref.rollout.drafter.training.dspark_debug_log_interval=${DSPARK_DEBUG_LOG_INTERVAL} \ - speco.standalone_tq_producer.request_timeout=${VLLM_REQUEST_TIMEOUT} \ - speco.standalone_tq_producer.max_inflight_requests=${VLLM_MAX_INFLIGHT_REQUESTS} \ - speco.standalone_tq_producer.per_endpoint_concurrency=${VLLM_PER_ENDPOINT_CONCURRENCY} \ - speco.standalone_tq_producer.input_queue_size=${PRODUCER_INPUT_QUEUE_SIZE} \ - speco.standalone_tq_producer.publish_queue_size=${PRODUCER_PUBLISH_QUEUE_SIZE} \ - speco.standalone_tq_producer.max_pending_samples=${PRODUCER_MAX_PENDING_SAMPLES} \ - speco.standalone_tq_producer.pending_poll_interval_seconds=${PRODUCER_PENDING_POLL_INTERVAL} \ - speco.standalone_tq_producer.max_sequence_length=${PRODUCER_MAX_SEQUENCE_LENGTH} \ - speco.standalone_tq_producer.max_feature_length=${PRODUCER_MAX_FEATURE_LENGTH} \ - speco.standalone_tq_producer.generation_max_tokens=${PRODUCER_GENERATION_MAX_TOKENS} \ - trainer.project_name=${project_name} \ - trainer.experiment_name=${exp_name} \ - "$@" diff --git a/tests/config/test_speco_config_overlay.py b/tests/config/test_speco_config_overlay.py index e14a5607..2b29e464 100644 --- a/tests/config/test_speco_config_overlay.py +++ b/tests/config/test_speco_config_overlay.py @@ -87,12 +87,8 @@ def test_overlay_has_expected_default_drafter_shape() -> None: assert standalone_training.target_feature_replay.endpoint_cooldown == 5 assert standalone_training.target_feature_pipeline.enabled is False assert standalone_training.target_feature_pipeline.concurrency == 16 - assert standalone_training.target_feature_pipeline.transfer_concurrency == 8 assert standalone_training.target_feature_pipeline.producer_prefetch_depth == 4 assert standalone_training.target_feature_pipeline.prefetch_depth == 2 - assert ( - standalone_training.target_feature_replay.mooncake.protocol == "tcp" - ) def test_overlay_composes_with_release_upstream_verl(tmp_path: Path) -> None: diff --git a/tests/examples/test_example_scripts.py b/tests/examples/test_example_scripts.py index 663b95c2..2b23b5e0 100644 --- a/tests/examples/test_example_scripts.py +++ b/tests/examples/test_example_scripts.py @@ -69,7 +69,7 @@ def test_example_keeps_speco_entrypoint_and_required_drafter_switches( def test_standalone_tq_training_example_uses_unified_launcher() -> None: source = ( - ROOT / "examples" / "run_qwen3-8b_drafter_separate_training.sh" + ROOT / "examples" / "run_qwen3-8b_drafter_dspark_separate_training.sh" ).read_text(encoding="utf-8") assert "-m verl_speco.standalone_tq_training_launcher" in source @@ -98,14 +98,6 @@ def test_standalone_tq_hidden_state_vllm_uses_separate_devices() -> None: assert '"kv_connector":"ExampleHiddenStatesConnector"' in source -def test_standalone_tq_compatibility_example_delegates_to_formal_entry() -> None: - source = ( - ROOT / "examples" / "run_qwen3-8b_drafter_dspark_separate_training.sh" - ).read_text(encoding="utf-8") - - assert 'run_qwen3-8b_drafter_separate_training.sh" "$@"' in source - - def test_vllm_eagle3_example_keeps_runtime_agnostic_training_switches() -> None: source = (ROOT / "examples" / "run_qwen3-8b_drafter_eagle3_vllm.sh").read_text( encoding="utf-8" diff --git a/tests/special_sanity/test_check_example_naming.py b/tests/special_sanity/test_check_example_naming.py index cac73db5..81708416 100644 --- a/tests/special_sanity/test_check_example_naming.py +++ b/tests/special_sanity/test_check_example_naming.py @@ -69,7 +69,7 @@ def test_missing_actor_backend_rejected(): def test_separate_training_entrypoint_passes(): - assert _violations("run_qwen3-8b_drafter_separate_training.sh") == [] + assert _violations("run_qwen3-8b_drafter_dspark_separate_training.sh") == [] def test_separate_training_may_name_its_drafter_backends(): diff --git a/tests/unit/test_mooncake_transfer.py b/tests/unit/test_mooncake_transfer.py deleted file mode 100644 index 2cb342ea..00000000 --- a/tests/unit/test_mooncake_transfer.py +++ /dev/null @@ -1,73 +0,0 @@ -# Copyright 2026 Bytedance Ltd. and/or its affiliates -# -# Licensed under the Apache License, Version 2.0 (the "License"); - -import sys -import types - -import pytest - -torch = pytest.importorskip("torch") -pytest.importorskip("safetensors") - -from verl_speco.trainer.mooncake_transfer import ( # noqa: E402 - MooncakeTensorStore, - MooncakeTransferConfig, - _parse_size, -) - - -class _RawStore: - objects = {} - - def setup(self, **kwargs): - self.setup_kwargs = kwargs - return 0 - - def put(self, key, payload): - self.objects[key] = payload - return 0 - - def get(self, key): - return self.objects.get(key) - - def remove(self, key, force): - self.objects.pop(key, None) - - def close(self): - return None - - -def test_mooncake_tensor_store_roundtrip(monkeypatch): - store_module = types.ModuleType("mooncake.store") - store_module.MooncakeDistributedStore = _RawStore - mooncake_module = types.ModuleType("mooncake") - mooncake_module.store = store_module - monkeypatch.setitem(sys.modules, "mooncake", mooncake_module) - monkeypatch.setitem(sys.modules, "mooncake.store", store_module) - - config = MooncakeTransferConfig( - local_hostname="localhost", - metadata_server="P2PHANDSHAKE", - master_server_address="127.0.0.1:50051", - global_segment_size=_parse_size("64MB"), - local_buffer_size=_parse_size("128MB"), - protocol="tcp", - device_name="", - get_timeout=1, - get_poll_interval=0.01, - ) - store = MooncakeTensorStore(config) - expected = { - "token_ids": torch.arange(4), - "hidden_states": torch.arange(24, dtype=torch.bfloat16).reshape(2, 3, 4), - } - - metadata = store.put("sample", expected) - actual = store.get("sample") - - assert metadata["tensor_shapes"]["hidden_states"] == (2, 3, 4) - assert torch.equal(actual["token_ids"], expected["token_ids"]) - assert torch.equal(actual["hidden_states"], expected["hidden_states"]) - store.remove("sample") - assert "sample" not in _RawStore.objects diff --git a/tests/unit/test_target_feature_pipeline.py b/tests/unit/test_target_feature_pipeline.py index 67e8a227..6063ca45 100644 --- a/tests/unit/test_target_feature_pipeline.py +++ b/tests/unit/test_target_feature_pipeline.py @@ -37,7 +37,6 @@ def test_target_feature_producer_preserves_batch_order_and_prefetches(): _Replayer(), rank=0, concurrency=2, - transfer_concurrency=1, producer_prefetch_depth=2, prefetch_depth=2, queue_timeout=2, @@ -67,7 +66,6 @@ def test_target_feature_producer_propagates_background_failure(): _FailingReplayer(), rank=0, concurrency=1, - transfer_concurrency=1, producer_prefetch_depth=1, prefetch_depth=1, queue_timeout=2, diff --git a/verl_speco/config/draft_trainer.yaml b/verl_speco/config/draft_trainer.yaml index e17d7862..1a600607 100644 --- a/verl_speco/config/draft_trainer.yaml +++ b/verl_speco/config/draft_trainer.yaml @@ -64,18 +64,6 @@ actor_rollout_ref: endpoint_cooldown: 5 on_generate: delete require_arange_positions: true - mooncake: - local_hostname: ${oc.env:MOONCAKE_LOCAL_HOSTNAME,localhost} - metadata_server: ${oc.env:MOONCAKE_METADATA_SERVER,http://localhost:8090/metadata} - master_server_address: ${oc.env:MOONCAKE_MASTER_SERVER,localhost:50051} - # Consumer ranks only need receive buffers. The vLLM producer owns - # the large segment through its MOONCAKE_GLOBAL_SEGMENT_SIZE env. - global_segment_size: 64MB - local_buffer_size: 1GB - protocol: tcp - device_name: "" - get_timeout: 120 - get_poll_interval: 0.02 offline_generation: input_type: token_replay input_path: null @@ -87,13 +75,12 @@ actor_rollout_ref: enabled: false path: null max_size_gb: 0 - # Overlap vLLM generation and Mooncake/file transfer with FSDP training. - # This pipeline is standalone-only and accepts vLLM replay backends. + # Overlap vLLM file replay with FSDP training. + # This pipeline is standalone-only. target_feature_pipeline: enabled: false # Global budgets. Standalone divides them across torchrun ranks. concurrency: 16 - transfer_concurrency: 8 producer_prefetch_depth: 4 prefetch_depth: 2 queue_timeout: 300 diff --git a/verl_speco/config/speco_base.yaml b/verl_speco/config/speco_base.yaml index 5ae6e7e5..1eca96b6 100644 --- a/verl_speco/config/speco_base.yaml +++ b/verl_speco/config/speco_base.yaml @@ -248,12 +248,3 @@ actor_rollout_ref: SimpleStorage: total_storage_size: 100000 num_data_storage_units: 8 - MooncakeStore: - auto_init: false - metadata_server: localhost:50050 - master_server_address: localhost:50051 - local_hostname: "" - protocol: tcp - global_segment_size: 4294967296 - local_buffer_size: 1073741824 - device_name: "" diff --git a/verl_speco/integration/mooncake_hidden_states_connector.py b/verl_speco/integration/mooncake_hidden_states_connector.py deleted file mode 100644 index 9d17ffec..00000000 --- a/verl_speco/integration/mooncake_hidden_states_connector.py +++ /dev/null @@ -1,327 +0,0 @@ -# Copyright 2026 Bytedance Ltd. and/or its affiliates -# -# Licensed under the Apache License, Version 2.0 (the "License"); -"""vLLM connector that publishes extracted target features through Mooncake. - -Load this module with ``kv_connector_module_path``. It is intentionally kept -outside normal SpeCo imports because its API is tied to vLLM V1 internals. -""" - -from __future__ import annotations - -import logging -import os -import re -import socket -from dataclasses import dataclass, field -from typing import TYPE_CHECKING, Any - -import torch -from vllm.config import VllmConfig, get_layers_from_vllm_config -from vllm.distributed.kv_transfer.kv_connector.v1.base import ( - KVConnectorBase_V1, - KVConnectorMetadata, - KVConnectorRole, - SupportsHMA, -) -from vllm.v1.attention.backend import AttentionMetadata -from vllm.v1.core.sched.output import SchedulerOutput - -from verl_speco.trainer.mooncake_transfer import ( - MooncakeTensorStore, - MooncakeTransferConfig, -) - -if TYPE_CHECKING: - from vllm.v1.core.kv_cache_manager import KVCacheBlocks - from vllm.v1.kv_cache_interface import KVCacheConfig - from vllm.v1.request import Request - -logger = logging.getLogger(__name__) - - -def _validate_vllm_version() -> None: - try: - from importlib.metadata import version - - from packaging.version import Version - - installed = Version(version("vllm")) - except Exception: # noqa: BLE001 - return - if installed < Version("0.23.0"): - raise RuntimeError( - "SpeCoMooncakeHiddenStatesConnector requires vLLM >= 0.23.0; " - f"found {installed}. The connector uses the V1 HMA hidden-state API." - ) - - -def _safe_key(key: str) -> str: - value = re.sub(r"[^a-zA-Z0-9_-]", "_", key) - return f"k{value}" if value and value[0].isdigit() else value - - -def _slot_mapping( - block_ids: list[int], page_size: int, num_tokens: int, device: torch.device -) -> torch.Tensor: - blocks = torch.tensor(block_ids, dtype=torch.int64, device=device) - offsets = torch.arange(page_size, dtype=torch.int64, device=device) - return (blocks.unsqueeze(1) * page_size + offsets).flatten()[:num_tokens] - - -@dataclass -class _RequestMetadata: - request_id: str - token_ids: torch.Tensor - block_ids: list[int] = field(default_factory=list) - - -@dataclass -class SpeCoMooncakeConnectorMetadata(KVConnectorMetadata): - requests: list[_RequestMetadata] = field(default_factory=list) - - def add(self, request_id: str, token_ids: list[int], block_ids: list[int]) -> None: - self.requests.append( - _RequestMetadata( - request_id=request_id, - token_ids=torch.tensor(token_ids, dtype=torch.long), - block_ids=list(block_ids), - ) - ) - - -class SpeCoMooncakeHiddenStatesConnector(KVConnectorBase_V1, SupportsHMA): - """Store-only connector for vLLM ``extract_hidden_states`` output.""" - - @property - def prefer_cross_layer_blocks(self) -> bool: - return False - - def __init__( - self, - vllm_config: VllmConfig, - role: KVConnectorRole, - kv_cache_config: "KVCacheConfig | None" = None, - ): - _validate_vllm_version() - super().__init__( - vllm_config=vllm_config, - role=role, - kv_cache_config=kv_cache_config, - ) - speculative = vllm_config.speculative_config - if speculative is None: - raise ValueError( - "SpeCoMooncakeHiddenStatesConnector requires extract_hidden_states" - ) - hf_config = speculative.draft_model_config.hf_config - self._layer_ids = list( - getattr(hf_config, "eagle_aux_hidden_state_layer_ids", []) - ) - self._hidden_size = int(vllm_config.model_config.get_hidden_size()) - self._training_layers = max(len(self._layer_ids) - 1, 1) - self._cache_layers: list[str] = [] - self._cache_group_id = self._find_cache_group(kv_cache_config) - self._active_requests: dict[str, Any] = {} - self._request_blocks: dict[str, list[int]] = {} - self._response_metadata: dict[str, dict[str, Any]] = {} - configured_prefix = os.getenv("SPECO_MOONCAKE_KEY_PREFIX") - self._key_prefix = _safe_key( - configured_prefix or f"{socket.gethostname()}_{os.getpid()}" - ) - self._store: MooncakeTensorStore | None = None - self._store_setup_attempted = False - self._tp_rank: int | None = None - - @staticmethod - def _find_cache_group(kv_cache_config: "KVCacheConfig | None") -> int | None: - if kv_cache_config is None: - return None - for index, group in enumerate(kv_cache_config.kv_cache_groups): - if any("cache_only_layers" in name for name in group.layer_names): - return index - return None - - def _get_tp_rank(self) -> int: - if self._tp_rank is None: - try: - from vllm.distributed import get_tensor_model_parallel_rank - - self._tp_rank = int(get_tensor_model_parallel_rank()) - except Exception: # noqa: BLE001 - self._tp_rank = 0 - return self._tp_rank - - def _ensure_store(self) -> MooncakeTensorStore | None: - if self._store_setup_attempted: - return self._store - self._store_setup_attempted = True - if self._get_tp_rank() != 0: - return None - try: - store = MooncakeTensorStore(MooncakeTransferConfig.from_mapping()) - store.setup() - self._store = store - except Exception: # noqa: BLE001 - logger.exception("Failed to initialize SpeCo Mooncake connector") - return self._store - - def start_load_kv(self, *args: Any, **kwargs: Any) -> None: - return None - - def wait_for_layer_load(self, layer_name: str) -> None: - return None - - def wait_for_save(self) -> None: - return None - - def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]) -> None: - from vllm.model_executor.models.extract_hidden_states import ( - CacheOnlyAttentionLayer, - ) - - layers = get_layers_from_vllm_config( - self._vllm_config, CacheOnlyAttentionLayer, list(kv_caches) - ) - self._cache_layers = list(layers) - if len(self._cache_layers) != 1: - raise RuntimeError( - "Expected one extract_hidden_states cache layer, got " - f"{self._cache_layers}" - ) - - def save_kv_layer( - self, - layer_name: str, - kv_layer: torch.Tensor, - attn_metadata: AttentionMetadata, - **kwargs: Any, - ) -> None: - if layer_name not in self._cache_layers: - return - from vllm.model_executor.models.extract_hidden_states import ( - CacheOnlyAttentionMetadata, - ) - - if not isinstance(attn_metadata, CacheOnlyAttentionMetadata): - raise TypeError( - "Expected CacheOnlyAttentionMetadata for extracted hidden states" - ) - metadata = self._get_connector_metadata() - if not isinstance(metadata, SpeCoMooncakeConnectorMetadata): - raise TypeError("Unexpected connector metadata type") - store = self._ensure_store() - if store is None: - return - page_size = int(kv_layer.shape[1]) - for request in metadata.requests: - num_tokens = int(request.token_ids.numel()) - positions = _slot_mapping( - request.block_ids, page_size, num_tokens, kv_layer.device - ) - if int(positions.numel()) < num_tokens: - continue - all_hidden = kv_layer.flatten(0, 1)[positions][:num_tokens].reshape( - num_tokens, -1 - ) - split_at = self._training_layers * self._hidden_size - training_hidden = all_hidden[:, :split_at].reshape( - num_tokens, self._training_layers, self._hidden_size - ) - last_hidden = all_hidden[:, -self._hidden_size :].unsqueeze(1) - hidden_states = torch.cat((training_hidden, last_hidden), dim=1).to( - torch.bfloat16 - ) - key = f"{self._key_prefix}_{_safe_key(request.request_id)}" - result = store.put( - key, - { - "hidden_states": hidden_states, - "token_ids": request.token_ids, - }, - ) - response = self._response_metadata.get(request.request_id) - if response is not None: - response.update(result) - - def get_num_new_matched_tokens( - self, request: "Request", num_computed_tokens: int - ) -> tuple[int | None, bool]: - return 0, False - - def update_state_after_alloc( - self, - request: "Request", - blocks: "KVCacheBlocks", - num_external_tokens: int, - ) -> None: - if num_external_tokens != 0: - raise ValueError("SpeCo Mooncake connector is store-only") - - def build_connector_meta( - self, scheduler_output: SchedulerOutput - ) -> KVConnectorMetadata: - metadata = SpeCoMooncakeConnectorMetadata() - for request in scheduler_output.scheduled_new_reqs: - token_ids = request.prompt_token_ids or [] - group_id = self._cache_group_id - if group_id is None: - group_id = max( - range(len(request.block_ids)), - key=lambda index: len(request.block_ids[index]), - ) - self._cache_group_id = group_id - blocks = list(request.block_ids[group_id]) - metadata.add(request.req_id, token_ids, blocks) - self._active_requests[request.req_id] = request - self._request_blocks[request.req_id] = blocks - self._response_metadata[request.req_id] = { - "mooncake_key": (f"{self._key_prefix}_{_safe_key(request.req_id)}"), - "input_ids_list": token_ids, - "tensor_shapes": { - "hidden_states": ( - len(token_ids), - self._training_layers + 1, - self._hidden_size, - ), - "token_ids": (len(token_ids),), - }, - "tensor_dtypes": { - "hidden_states": "bfloat16", - "token_ids": "int64", - }, - } - - cached = scheduler_output.scheduled_cached_reqs - for index, request_id in enumerate(cached.req_ids): - if request_id not in self._active_requests: - continue - new_blocks = cached.new_block_ids[index] - if new_blocks is not None: - self._request_blocks[request_id].extend( - new_blocks[self._cache_group_id] - ) - request = self._active_requests[request_id] - metadata.add( - request_id, - request.prompt_token_ids or [], - self._request_blocks[request_id], - ) - return metadata - - def request_finished( - self, request: "Request", block_ids: list[int] - ) -> tuple[bool, dict[str, Any] | None]: - request_id = request.request_id - self._active_requests.pop(request_id, None) - self._request_blocks.pop(request_id, None) - return False, self._response_metadata.pop(request_id, None) - - def request_finished_all_groups( - self, request: "Request", block_ids: tuple[list[int], ...] - ) -> tuple[bool, dict[str, Any] | None]: - return self.request_finished(request, block_ids[0] if block_ids else []) - - @classmethod - def get_required_kvcache_layout(cls, vllm_config: VllmConfig) -> str | None: - return "NHD" diff --git a/verl_speco/trainer/draft_training_loop.py b/verl_speco/trainer/draft_training_loop.py index a1b50698..f20243aa 100644 --- a/verl_speco/trainer/draft_training_loop.py +++ b/verl_speco/trainer/draft_training_loop.py @@ -264,7 +264,7 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: ): raise ValueError( "target_feature_pipeline.enabled=true requires a vLLM replay " - "backend (vllm_file or vllm_mooncake)" + "backend (vllm_file)" ) from verl_speco.trainer.target_feature_pipeline import ( TargetFeatureProducer, @@ -280,15 +280,6 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: ), 1, ), - transfer_concurrency=int( - max( - math.ceil( - int(pipeline_cfg.get("transfer_concurrency", 8) or 8) - / world_size - ), - 1, - ) - ), producer_prefetch_depth=int( pipeline_cfg.get("producer_prefetch_depth", 4) or 4 ), @@ -1054,7 +1045,6 @@ def _log_standalone_step_metrics(metrics: dict[str, float], *, rank: int) -> Non ("replay/cache_hit_ratio", "cache_hit"), ("replay/target_forward_time_total", "target_forward_total"), ("replay/vllm_request_time_total", "vllm_request_total"), - ("replay/mooncake_get_time_total", "mooncake_get_total"), ("producer/consumer_wait_time_total", "producer_wait_total"), ("producer/ready_queue_size", "ready_batches"), ): @@ -1175,7 +1165,7 @@ def _next_batch_across_ranks( ) -> Any | None: """Fetch one batch and make producer failures visible to every rank. - Producer and Mooncake errors happen before the FSDP training step. Every + Producer and replay errors happen before the FSDP training step. Every rank therefore reports its fetch result through the same collective before any rank is allowed to enter model collectives. This prevents healthy ranks from waiting in FSDP after another rank has already started cleanup. diff --git a/verl_speco/trainer/mooncake_transfer.py b/verl_speco/trainer/mooncake_transfer.py deleted file mode 100644 index 0f4a6c51..00000000 --- a/verl_speco/trainer/mooncake_transfer.py +++ /dev/null @@ -1,214 +0,0 @@ -# Copyright 2026 Bytedance Ltd. and/or its affiliates -# -# Licensed under the Apache License, Version 2.0 (the "License"); -"""Small optional Mooncake client used by standalone target-feature replay. - -The payload is stored as one safetensors object. A single object makes the -producer response atomic and avoids the file creation/locking protocol used by -``vllm_file``. Mooncake is imported lazily so normal and online training do -not acquire a runtime dependency on it. -""" - -from __future__ import annotations - -import logging -import os -import socket -import time -from dataclasses import dataclass -from typing import Any - -import torch - -logger = logging.getLogger(__name__) - - -def _parse_size(value: str | int) -> int: - if isinstance(value, int): - return value - text = str(value).strip().upper() - multipliers = { - "TB": 1024**4, - "GB": 1024**3, - "MB": 1024**2, - "KB": 1024, - "T": 1024**4, - "G": 1024**3, - "M": 1024**2, - "K": 1024, - "B": 1, - } - for suffix in sorted(multipliers, key=len, reverse=True): - if text.endswith(suffix): - return int(float(text[: -len(suffix)]) * multipliers[suffix]) - return int(text) - - -@dataclass(frozen=True) -class MooncakeTransferConfig: - local_hostname: str - metadata_server: str - master_server_address: str - global_segment_size: int - local_buffer_size: int - protocol: str - device_name: str - get_timeout: float - get_poll_interval: float - - @classmethod - def from_mapping(cls, config: Any | None = None) -> "MooncakeTransferConfig": - config = config or {} - - def value(name: str, default: Any) -> Any: - getter = getattr(config, "get", None) - if callable(getter): - return getter(name, default) - return getattr(config, name, default) - - master = str( - value( - "master_server_address", - os.getenv("MOONCAKE_MASTER_SERVER", "127.0.0.1:50051"), - ) - ) - master_host = master.rsplit(":", 1)[0] - return cls( - local_hostname=str( - value( - "local_hostname", - os.getenv("MOONCAKE_LOCAL_HOSTNAME", socket.gethostname()), - ) - ), - metadata_server=str( - value( - "metadata_server", - os.getenv( - "MOONCAKE_METADATA_SERVER", - f"http://{master_host}:8090/metadata", - ), - ) - ), - master_server_address=master, - global_segment_size=_parse_size( - value( - "global_segment_size", - os.getenv("MOONCAKE_GLOBAL_SEGMENT_SIZE", "4GB"), - ) - ), - local_buffer_size=_parse_size( - value( - "local_buffer_size", - os.getenv("MOONCAKE_LOCAL_BUFFER_SIZE", "1GB"), - ) - ), - protocol=str(value("protocol", os.getenv("MOONCAKE_PROTOCOL", "tcp"))), - device_name=str( - value("device_name", os.getenv("MOONCAKE_DEVICE_NAME", "")) - ), - get_timeout=float(value("get_timeout", 120.0)), - get_poll_interval=max(float(value("get_poll_interval", 0.02)), 0.001), - ) - - def export_environment(self) -> None: - os.environ["MOONCAKE_LOCAL_HOSTNAME"] = self.local_hostname - os.environ["MOONCAKE_METADATA_SERVER"] = self.metadata_server - os.environ["MOONCAKE_MASTER_SERVER"] = self.master_server_address - os.environ["MOONCAKE_GLOBAL_SEGMENT_SIZE"] = str(self.global_segment_size) - os.environ["MOONCAKE_LOCAL_BUFFER_SIZE"] = str(self.local_buffer_size) - os.environ["MOONCAKE_PROTOCOL"] = self.protocol - os.environ["MOONCAKE_DEVICE_NAME"] = self.device_name - if self.protocol.lower() == "tcp": - os.environ.setdefault("MC_STORE_MEMCPY", "0") - - -class MooncakeTensorStore: - """Store and retrieve a tensor dictionary as one Mooncake object.""" - - def __init__(self, config: MooncakeTransferConfig): - self.config = config - self._store: Any | None = None - - def setup(self) -> None: - if self._store is not None: - return - self.config.export_environment() - try: - from mooncake.store import MooncakeDistributedStore - except ImportError as exc: - raise RuntimeError( - "Mooncake replay requires mooncake-transfer-engine " - "(use mooncake-transfer-engine-npu on Ascend)" - ) from exc - store = MooncakeDistributedStore() - result = store.setup( - local_hostname=self.config.local_hostname, - metadata_server=self.config.metadata_server, - global_segment_size=self.config.global_segment_size, - local_buffer_size=self.config.local_buffer_size, - protocol=self.config.protocol, - rdma_devices=self.config.device_name, - master_server_addr=self.config.master_server_address, - ) - if result not in (None, 0): - raise RuntimeError(f"Mooncake client setup failed with code {result}") - self._store = store - - def put(self, key: str, tensors: dict[str, torch.Tensor]) -> dict[str, Any]: - self.setup() - assert self._store is not None - from safetensors.torch import save - - cpu_tensors = { - name: tensor.detach().to("cpu").contiguous() - for name, tensor in tensors.items() - } - payload = save(cpu_tensors) - result = self._store.put(key, payload) - if result not in (None, 0): - raise RuntimeError(f"Mooncake put failed for {key!r}: code={result}") - return { - "mooncake_key": key, - "tensor_shapes": { - name: tuple(tensor.shape) for name, tensor in cpu_tensors.items() - }, - "tensor_dtypes": { - name: str(tensor.dtype).removeprefix("torch.") - for name, tensor in cpu_tensors.items() - }, - "payload_bytes": len(payload), - } - - def get(self, key: str) -> dict[str, torch.Tensor]: - self.setup() - assert self._store is not None - from safetensors.torch import load - - deadline = time.monotonic() + self.config.get_timeout - while True: - payload = self._store.get(key) - if payload is not None: - return dict(load(bytes(payload))) - if time.monotonic() >= deadline: - raise TimeoutError( - f"Mooncake object {key!r} was unavailable for " - f"{self.config.get_timeout:.1f}s" - ) - time.sleep(self.config.get_poll_interval) - - def remove(self, key: str) -> None: - if self._store is None: - return - try: - remove = getattr(self._store, "remove", None) - if callable(remove): - remove(key, True) - else: - self._store.batch_remove([key], force=True) - except Exception: # noqa: BLE001 - logger.warning("Failed to remove Mooncake object %s", key, exc_info=True) - - def close(self) -> None: - if self._store is not None and hasattr(self._store, "close"): - self._store.close() - self._store = None diff --git a/verl_speco/trainer/target_feature_pipeline.py b/verl_speco/trainer/target_feature_pipeline.py index fbf9f6c0..284d87bf 100644 --- a/verl_speco/trainer/target_feature_pipeline.py +++ b/verl_speco/trainer/target_feature_pipeline.py @@ -42,7 +42,6 @@ def __init__( *, rank: int, concurrency: int, - transfer_concurrency: int, producer_prefetch_depth: int, prefetch_depth: int, queue_timeout: float, @@ -51,7 +50,6 @@ def __init__( self.replayer = replayer self.rank = int(rank) self.concurrency = max(int(concurrency), 1) - self.transfer_concurrency = max(int(transfer_concurrency), 1) self.producer_prefetch_depth = max(int(producer_prefetch_depth), 1) self.prefetch_depth = max(int(prefetch_depth), 1) self.queue_timeout = max(float(queue_timeout), 1.0) @@ -61,10 +59,6 @@ def __init__( max_workers=self.concurrency, thread_name_prefix=f"speco-request-r{self.rank}", ) - self._transfer_executor = ThreadPoolExecutor( - max_workers=self.transfer_concurrency, - thread_name_prefix=f"speco-transfer-r{self.rank}", - ) self._thread = threading.Thread( target=self._run, name=f"speco-target-producer-r{self.rank}", @@ -80,10 +74,9 @@ def __init__( self._thread.start() logger.info( "[target producer rank=%s] started request_concurrency=%s " - "transfer_concurrency=%s producer_prefetch_depth=%s prefetch_depth=%s", + "producer_prefetch_depth=%s prefetch_depth=%s", self.rank, self.concurrency, - self.transfer_concurrency, self.producer_prefetch_depth, self.prefetch_depth, ) @@ -100,20 +93,10 @@ def submit_next() -> bool: except StopIteration: return False started = time.perf_counter() - if self.replayer.backend == "vllm_mooncake": - futures = [ - self._request_executor.submit( - self.replayer.produce_mooncake_descriptor, sample - ) - for sample in samples - ] - else: - futures = [ - self._request_executor.submit( - self.replayer.materialize, [sample] - ) - for sample in samples - ] + futures = [ + self._request_executor.submit(self.replayer.materialize, [sample]) + for sample in samples + ] pending.append((started, futures)) return True @@ -127,22 +110,10 @@ def submit_next() -> bool: self.producer_seconds += time.perf_counter() - started submit_next() transfer_started = time.perf_counter() - if self.replayer.backend == "vllm_mooncake": - transfer_futures = [ - self._transfer_executor.submit( - self.replayer.consume_pipeline_mooncake_descriptor, - descriptor, - ) - for descriptor in produced - ] - batch = [future.result() for future in transfer_futures] - else: - batch = [item for group in produced for item in group] + batch = [item for group in produced for item in group] self.transfer_seconds += time.perf_counter() - transfer_started self.produced_batches += 1 self.produced_samples += len(batch) - if self.replayer.backend == "vllm_mooncake": - self.replayer.record_pipeline_materialized(len(batch)) self._put(batch) self._put(_END) except BaseException as exc: # noqa: BLE001 @@ -169,7 +140,7 @@ def __next__(self) -> list[DraftFeatureSample]: except queue.Empty as exc: raise TimeoutError( "Timed out waiting for target-feature producer; inspect the vLLM " - "and Mooncake logs for a stalled request or missing object" + "logs for a stalled request or missing hidden-state file" ) from exc self.consumer_wait_seconds += time.perf_counter() - started if value is _END: @@ -194,5 +165,4 @@ def metrics(self) -> dict[str, float]: def close(self) -> None: self._stop.set() self._request_executor.shutdown(wait=True, cancel_futures=True) - self._transfer_executor.shutdown(wait=True, cancel_futures=True) self._thread.join(timeout=5.0) diff --git a/verl_speco/trainer/target_feature_replay.py b/verl_speco/trainer/target_feature_replay.py index 6ed068bb..43ab361e 100644 --- a/verl_speco/trainer/target_feature_replay.py +++ b/verl_speco/trainer/target_feature_replay.py @@ -42,13 +42,6 @@ class HiddenStateAlignmentError(ValueError): """The returned hidden states cannot cover the requested training sample.""" -@dataclass(frozen=True) -class MooncakeReplayDescriptor: - sample: DraftReplaySample - prompt_ids: list[int] - key: str - - @dataclass(frozen=True) class FeatureContract: """Explicit inputs for converting one vLLM payload into a training sample.""" @@ -548,12 +541,10 @@ def __init__( .strip() .lower() ) - if self.backend == "mooncake": - self.backend = "vllm_mooncake" - if self.backend not in {"torch", "vllm_file", "vllm_mooncake"}: + if self.backend not in {"torch", "vllm_file"}: raise ValueError( f"Unsupported target_feature_replay.backend={self.backend!r}; " - "expected 'torch', 'vllm_file', or 'vllm_mooncake'" + "expected 'torch' or 'vllm_file'" ) configured_model_path = _config_value(self.replay_cfg, "model_path", None) model_path = configured_model_path or self.draft_config.model.path @@ -682,20 +673,15 @@ def __init__( for index, endpoint in enumerate(self.vllm_endpoints) ] self._vllm_clients_initialized = False - self.mooncake_store: Any | None = None self._client_lock = threading.Lock() self._endpoint_lock = threading.Lock() self._metrics_lock = threading.Lock() - self._pending_keys_lock = threading.Lock() - self._pending_mooncake_keys: set[str] = set() self.cache_hits = 0 self.cache_misses = 0 self.materialized_samples = 0 self.target_forward_seconds = 0.0 self.vllm_request_seconds = 0.0 self.vllm_requests = 0 - self.mooncake_get_seconds = 0.0 - self.mooncake_gets = 0 self._warned_replay_algorithm_mismatch = False self._warned_replay_layer_mismatch = False self._warned_replay_layout_mismatch = False @@ -873,8 +859,6 @@ def _ensure_model(self) -> None: def _materialize_one(self, sample: DraftReplaySample) -> DraftFeatureSample: if self.backend == "vllm_file": return self._materialize_one_vllm_file(sample) - if self.backend == "vllm_mooncake": - return self._materialize_one_vllm_mooncake(sample) return self._materialize_one_torch(sample) def _materialize_one_torch(self, sample: DraftReplaySample) -> DraftFeatureSample: @@ -1043,117 +1027,6 @@ def _materialize_one_vllm_file( logger.warning("Failed to delete vLLM hidden-states file %s", path) return feature - def _materialize_one_vllm_mooncake( - self, sample: DraftReplaySample - ) -> DraftFeatureSample: - return self.consume_mooncake_descriptor( - self._request_mooncake_descriptor(sample) - ) - - def produce_mooncake_descriptor( - self, sample: DraftReplaySample | DraftFeatureSample - ) -> MooncakeReplayDescriptor | DraftFeatureSample: - if isinstance(sample, DraftFeatureSample): - return sample - if not isinstance(sample, DraftReplaySample): - raise TypeError( - "Mooncake producer expected DraftReplaySample or " - f"DraftFeatureSample, got {type(sample)!r}" - ) - if self.use_logits: - raise NotImplementedError( - "target_feature_replay.backend=vllm_mooncake does not yet " - "support training.use_logits=true; use backend=torch." - ) - self._validate_target_path(sample) - cache_key = self._cache_key(sample) - with self._cache_lock: - cached = self.cache.get(cache_key) if self.cache is not None else None - if cached is not None: - with self._metrics_lock: - self.cache_hits += 1 - return cached - with self._metrics_lock: - self.cache_misses += 1 - return self._request_mooncake_descriptor(sample) - - def _request_mooncake_descriptor( - self, sample: DraftReplaySample - ) -> MooncakeReplayDescriptor: - self._validate_vllm_positions(sample) - feature_positions = sample.feature_positions.detach().cpu().long() - feature_end = int(feature_positions[-1].item()) + 1 - prompt_ids = sample.input_ids[:feature_end].detach().cpu().long().tolist() - response = self._request_vllm_response(prompt_ids) - params = getattr(response, "kv_transfer_params", None) - if params is None: - raise ValueError("vLLM response missing kv_transfer_params") - key = params.get("mooncake_key") - if not key: - raise ValueError( - "vLLM response missing mooncake_key; start vLLM with " - "SpeCoMooncakeHiddenStatesConnector" - ) - key = str(key) - with self._pending_keys_lock: - self._pending_mooncake_keys.add(key) - return MooncakeReplayDescriptor(sample, prompt_ids, key) - - def consume_mooncake_descriptor( - self, descriptor: MooncakeReplayDescriptor | DraftFeatureSample - ) -> DraftFeatureSample: - if isinstance(descriptor, DraftFeatureSample): - return descriptor - store = self._ensure_mooncake_store() - transfer_started = time.perf_counter() - payload = store.get(descriptor.key) - with self._metrics_lock: - self.mooncake_get_seconds += time.perf_counter() - transfer_started - self.mooncake_gets += 1 - try: - feature = self._feature_from_vllm_payload( - descriptor.sample, - payload, - prompt_ids=descriptor.prompt_ids, - source="token_replay_vllm_mooncake", - ) - return feature - finally: - if self.vllm_on_generate == "delete": - store.remove(descriptor.key) - with self._pending_keys_lock: - self._pending_mooncake_keys.discard(descriptor.key) - - def consume_pipeline_mooncake_descriptor( - self, descriptor: MooncakeReplayDescriptor | DraftFeatureSample - ) -> DraftFeatureSample: - feature = self.consume_mooncake_descriptor(descriptor) - if isinstance(descriptor, MooncakeReplayDescriptor) and self.cache is not None: - with self._cache_lock: - self.cache.put(self._cache_key(descriptor.sample), feature) - return feature - - def record_pipeline_materialized(self, count: int) -> None: - with self._metrics_lock: - self.materialized_samples += int(count) - - def _ensure_mooncake_store(self): - if self.mooncake_store is not None: - return self.mooncake_store - with self._client_lock: - if self.mooncake_store is None: - from verl_speco.trainer.mooncake_transfer import ( - MooncakeTensorStore, - MooncakeTransferConfig, - ) - - mooncake_cfg = _config_value(self.replay_cfg, "mooncake", {}) or {} - self.mooncake_store = MooncakeTensorStore( - MooncakeTransferConfig.from_mapping(mooncake_cfg) - ) - self.mooncake_store.setup() - return self.mooncake_store - def _validate_vllm_positions(self, sample: DraftReplaySample) -> None: if not self.vllm_require_arange_positions: return @@ -1555,9 +1428,6 @@ def metrics(self) -> dict[str, float]: metrics[f"{prefix}_request_time_total"] = float( state.request_seconds ) - if self.backend == "vllm_mooncake": - metrics["replay/mooncake_gets_total"] = float(self.mooncake_gets) - metrics["replay/mooncake_get_time_total"] = float(self.mooncake_get_seconds) total = self.cache_hits + self.cache_misses if total > 0: metrics["replay/cache_hit_ratio"] = self.cache_hits / float(total) @@ -1579,15 +1449,6 @@ def close(self) -> None: exc_info=True, ) state.client = None - if self.mooncake_store is not None: - with self._pending_keys_lock: - pending_keys = tuple(self._pending_mooncake_keys) - self._pending_mooncake_keys.clear() - if self.vllm_on_generate == "delete": - for key in pending_keys: - self.mooncake_store.remove(key) - self.mooncake_store.close() - self.mooncake_store = None if self.model is None: return try: From b238e3601bb0e452f18638a730f12542203185ff Mon Sep 17 00:00:00 2001 From: xyr <1615681145@qq.com> Date: Tue, 1 Sep 2026 10:00:16 +0800 Subject: [PATCH 45/50] feat: add standalone resume and TQ producer/consumer updates --- ...sync_vllm_mooncake_dspark_training_plan.md | 2902 +++++++++++++++++ ...eature_sample_tq_protocol_refactor_plan.md | 549 ++++ docs/standalone_tq_consumer_implementation.md | 713 ++++ docs/standalone_tq_drafter_resume_plan.md | 448 +++ ...standalone_tq_foundation_implementation.md | 1030 ++++++ docs/standalone_tq_producer.md | 212 ++ docs/standalone_tq_training_parameters.md | 169 + ...standalone_vllm_tq_dspark_training_plan.md | 1130 +++++++ docs/transferqueue_integration_plan.md | 145 + docs/vllm_direct_hidden_state_cotrain_plan.md | 666 ++++ tests/unit/test_standalone_resume.py | 55 + .../test_standalone_tq_training_launcher.py | 12 + tests/unit/test_tq_consumer.py | 1 + tests/unit/test_tq_producer.py | 37 + verl_speco/config/speco_base.yaml | 3 + verl_speco/standalone_tq_producer.py | 23 +- verl_speco/standalone_tq_training_launcher.py | 28 +- verl_speco/trainer/draft_training_loop.py | 114 +- verl_speco/trainer/standalone_resume.py | 131 + verl_speco/trainer/tq_sample_source.py | 10 + 20 files changed, 8368 insertions(+), 10 deletions(-) create mode 100644 docs/async_vllm_mooncake_dspark_training_plan.md create mode 100644 docs/draft_feature_sample_tq_protocol_refactor_plan.md create mode 100644 docs/standalone_tq_consumer_implementation.md create mode 100644 docs/standalone_tq_drafter_resume_plan.md create mode 100644 docs/standalone_tq_foundation_implementation.md create mode 100644 docs/standalone_tq_producer.md create mode 100644 docs/standalone_tq_training_parameters.md create mode 100644 docs/standalone_vllm_tq_dspark_training_plan.md create mode 100644 docs/transferqueue_integration_plan.md create mode 100644 docs/vllm_direct_hidden_state_cotrain_plan.md create mode 100644 tests/unit/test_standalone_resume.py create mode 100644 verl_speco/trainer/standalone_resume.py diff --git a/docs/async_vllm_mooncake_dspark_training_plan.md b/docs/async_vllm_mooncake_dspark_training_plan.md new file mode 100644 index 00000000..525da0fb --- /dev/null +++ b/docs/async_vllm_mooncake_dspark_training_plan.md @@ -0,0 +1,2902 @@ +# verl-SpeCo 异步 vLLM → TransferQueue(Mooncake 后端)→ DSpark 流式训练方案 + +Last updated: 08/21/2026 + +## 1. 文档范围 + +本文只讨论当前 `verl-SpeCo-ls` 项目的独立草稿模型训练入口: + +```text +examples/run_qwen3-8b_drafter_separate_training.sh + → python -m verl_speco.standalone_tq_training_launcher + ├─→ TransferQueue owner + ├─→ vLLM hidden-state producer + └─→ TransferQueue consumer + → python -m verl_speco.draft_train_launcher + → torch.distributed.run + → python -m verl_speco.draft_train + → run_standalone_draft_training() +``` + +目标是训练 `mode=train/offline` 的草稿模型,不是 RL 与草稿模型一起训练,也不是先生成全量 hidden states 再长期保存到磁盘。 + +输入既可以是已有 response 的 replay 文件,也可以是 verl prompt-only Parquet。新流水线需要: + +1. Producer 读取 prompt;缺少 response 时由 target vLLM 生成; +2. Producer 并行请求一个或多个 vLLM endpoint 做 prefill; +3. hidden states 写入 TransferQueue(简称 TQ),TQ 的数据后端使用 Mooncake; +4. DSpark 的各训练 rank 从 TQ/Mooncake 并行读取自己的 local batch; +5. 所有 rank 完成同一个 optimizer step 后,清理这一批 TQ 数据; +6. 不使用自研 Stream Coordinator;第一版复用 PR #48 已验证的 TQ KV 传输模式,由 TQ 保存 payload 和 tags,rank 0 通过 `kv_list` 发现 ready key 并组成 global batch。 + +本文不照搬 AngelSpec 的进程组织。实现依据是旁边只读参考仓库 `verl-SpeCo` 的 PR #48 分支,尤其是: + +```text +verl_speco/integration/transferqueue_bridge.py +verl_speco/integration/sglang_runtime.py +verl_speco/integration/oldlogprob_runtime.py +verl_speco/workers/speco_worker.py +verl_speco/integration/task_runner.py +``` + +参考仓库只用于理解和复用设计,不修改其代码。真正实现仍放在本项目 `verl-SpeCo-ls`,服务于 standalone drafter training。 + +建议按下面顺序阅读: + +1. 第一部分先介绍 verl/SpeCo 中的进程、`TokenOutput`、`DataProto`、WorkerGroup、replica 和 owner rank; +2. 再完整解释 PR #48 的 SGLang TQ 路径; +3. 再解释 PR #48 的 old-logprob TQ 路径; +4. 再解释 hidden 恢复后如何进入 online buffer 和 optimizer step; +5. 第二部分才讨论如何把相同 TQ 传输能力改造成独立 Producer/DSpark Consumer。 + +## 第一部分:PR #48 原始 TQ 流程 + +这一部分只解释只读参考仓库 `../verl-SpeCo` 的 PR #48。这里的 Producer、Consumer、Ray driver、数据结构和生命周期全部是 PR #48 当前代码的行为,不是本文为 standalone 设计的行为。 + +### 1A. 阅读 PR #48 前必须知道的项目对象 + +#### SGLang server + +SGLang server 是 rollout 推理进程。它接收 prompt,执行 target model 推理并生成 response。SpeCo 的 SGLang patch 还会在推理过程中收集指定层 hidden states。 + +它不是 drafter trainer,也不执行 optimizer step。在 PR #48 中它是 hidden-state Producer。 + +#### TokenOutput + +`TokenOutput` 是一次 rollout request 的返回对象。核心字段可以概括为: + +```python +TokenOutput( + token_ids=list[int], + log_probs=..., + routed_experts=..., + extra_fields={ + "global_steps": int, + "drafter_sample": dict | None, + }, +) +``` + +`token_ids` 是生成结果;`extra_fields` 是 SpeCo 添加的旁路字段。`drafter_sample` 不影响正常 response 返回,它用于把草稿训练所需信息从 rollout 侧带回训练控制层。 + +#### DataProto 和 non_tensor_batch + +verl 将一批 rollout 结果整理成 `DataProto`。它通常分为: + +```python +DataProto( + batch=TensorDict(...), + non_tensor_batch={...}, + meta_info={...}, +) +``` + +- `batch`:规则的批量 tensor,例如 prompts、responses、attention mask; +- `non_tensor_batch`:不能直接组成规则 dense batch 的 Python 对象或 object array; +- `meta_info`:批次级配置和指标。 + +每个 request 的 `TokenOutput.extra_fields["drafter_sample"]` 最终会被 rollout/agent-loop 聚合到: + +```python +gen_batch_output.non_tensor_batch["drafter_sample"] +``` + +所以 SGLang 侧写的是单 request `TokenOutput`,driver 侧拿到的是批量 `DataProto`。 + +#### RayPPOTrainer driver + +`SpecoRayPPOTrainer` 所在进程是控制进程,简称 driver。它负责按训练 step 依次调用 rollout、old-logprob、actor update,以及向各 WorkerGroup 发 RPC。 + +driver 不执行 drafter 模型 forward。PR #48 之前,它会接触包含 hidden tensor 的 `drafter_sample`;PR #48 之后,它只处理中小 tensor、标量 metadata 和 TQ key。 + +#### WorkerGroup + +WorkerGroup 是 verl 对一组 Ray worker 的调用封装。代码: + +```python +self.drafter_wg.collect_rollout_features(buckets) +``` + +不是普通本地函数调用,而是根据注册的 dispatch rule,将参数分发到多张卡上的 `SpecoWorker.collect_rollout_features()`。 + +#### Rollout replica + +rollout replica 是一组共同承载一个 rollout 模型副本的进程/GPU。一个任务可能存在多个 data-parallel rollout replica。`replica_rank` 标识当前 sample 是哪个 rollout replica 生成的: + +```text +replica_rank = 0, 1, 2, ... +``` + +#### Drafter training replica、DP rank 和 SP rank + +drafter 训练也可能按 data parallel 和 sequence parallel 组织: + +```text +drafter replica / DP rank 0 + ├─ SP rank 0 + └─ SP rank 1 + +drafter replica / DP rank 1 + ├─ SP rank 0 + └─ SP rank 1 +``` + +同一个 drafter DP replica 内的 SP ranks 共同执行一个模型副本的训练。`replica_rank` 用来把 rollout replica 产生的数据路由到对应 drafter DP replica。 + +#### Owner rank + +`collect_rollout_features` 注册了: + +```python +@register( + dispatch_mode=make_nd_compute_dispatch_fn( + mesh_name="drafter_owner_route" + ) +) +``` + +每个 drafter DP replica 的 `SP rank 0` 被标记为 collect leader/owner。driver 传入的是按 replica 分好的 bucket,dispatch 层负责把 bucket 发到对应训练组。一个 replica 内可能有多个 rank 需要共同训练,但只有指定 leader 负责汇总 RPC 返回。 + +这里的 owner route 是 PR #48 为什么不能简单让任意 worker 随机拿 sample 的原因:sample 必须进入与当前 drafter device mesh 一致的训练组。 + +## 2. PR #48 改造前的 online 特征流程 + +PR #48 改造的是 SpeCo 的 online drafter feature transport。hidden states 有两条主要来源。 + +### 2.1 SGLang rollout hidden 路径 + +改造前: + +```text +SGLang server + → 生成 drafter_sample,其中直接包含 hidden_states CPU tensor + → TokenOutput.extra_fields + → RayPPOTrainer driver 收集 drafter_sample + → driver 按 drafter replica/owner 分桶 + → Ray dispatch / object store + → SpecoWorker.collect_rollout_features(samples) + → _store_rollout_sample() + → online drafter buffer/train +``` + +此时 `drafter_sample` 类似: + +```python +{ + "input_ids": int64[1, total_len], + "prompts": int64[1, prompt_len], + "responses": int64[1, response_len], + "hidden_states": bf16[1, hidden_rows, hidden_dim], + "hidden_positions": int64[1, hidden_rows], + "target_logprobs": tensor | None, + "global_step": 42, + "replica_rank": 1, +} +``` + +问题是整个字典经过 driver,而 `hidden_states` 是其中最大的字段。driver 的 host memory 和 Ray object store 都会承载这些 tensor。 + +### 2.2 old-logprob hook hidden 路径 + +另一条路径在 actor old-logprob forward 中捕获 hidden states。改造前,大 chunk 通过: + +```python +chunk_ref = ray.put(hidden_chunk) +``` + +sample 不直接带 tensor,而是带: + +```python +{ + "hidden_states_ref_chunks": [ + { + "ref": ray_object_ref, + "start": 0, + "length": 512, + }, + ], +} +``` + +drafter worker 再 `ray.get(ref)`,根据 `start/length` 切出每条 sample 所需行。 + +### 2.3 PR #48 要改变的边界 + +PR #48 没有改变: + +- rollout 什么时候产生 sample; +- driver 如何触发 drafter worker; +- drafter worker 如何调用 `_store_rollout_sample()`; +- drafter model 的训练逻辑; +- drafter 权重发布。 + +它只改变大 tensor 的跨进程介质: + +```text +改造前:Producer → Ray driver/object store → Consumer +改造后:Producer → TQ storage → Consumer + key 仍走原 Ray 控制路径 +``` + +## 3. PR #48 改造后的完整 TQ 流程 + +### 3.0 总览 + +PR #48 的目标不是让 TQ 自己产生训练 batch,而是把原来经过 Ray driver/Ray object store 的大 hidden tensor 攁到 TQ。原来的控制路径继续存在,只是控制路径上从“大 tensor”变成“小 key”。 + +整体结构是: + +```text + 原 Ray 控制路径 + drafter_sample / chunk ref +Producer ───────────────────── key ───────────────────▶ Consumer + │ │ + │ kv_put(large tensor) │ kv_batch_get(key) + ▼ ▼ +TransferQueue storage ─────────────────────────────────────┘ +``` + +因此 PR #48 同时保留两条通道: + +```text +控制通道:Producer → Ray driver → drafter worker +数据通道:Producer → TQ storage → drafter worker +``` + +控制通道负责告诉 Consumer “本次应该处理哪个样本、对应哪个 TQ key”;数据通道负责传输 hidden states 等大 tensor。 + +#### 3.0.1 配置放在哪里 + +PR #48 在 drafter training 配置下增加: + +```yaml +actor_rollout_ref: + rollout: + drafter: + training: + transfer_queue: + enable: false + backend: + storage_backend: SimpleStorage + SimpleStorage: + total_storage_size: 100000 + num_data_storage_units: 8 +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/config/speco_base.yaml +``` + +这里的 `transfer_queue` 是 SpeCo drafter 自己的配置,不是 standalone `feature_store.type`,也不是简单把上游 verl 的 `transfer_queue.enable` 打开。 + +#### 3.0.2 TaskRunner 创建整套 TQ + +RL 任务启动时,`SpecoTaskRunner.run()` 在创建 workers 之前调用: + +```python +from verl_speco.integration.transferqueue_bridge import ( + close_transfer_queue, + init_transfer_queue, +) + +transfer_queue_started = init_transfer_queue(config) +try: + trainer.init_workers() + trainer.fit() +finally: + if transfer_queue_started: + close_transfer_queue() +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/integration/task_runner.py:319 +``` + +`init_transfer_queue(config)` 内部读取: + +```python +config.actor_rollout_ref.rollout.drafter.training.transfer_queue +``` + +然后执行: + +```python +tq.init(_to_plain_dict(tq_cfg)) +``` + +并记录: + +```python +_state["initialized"] = True +_state["owner"] = True +``` + +这里的 owner 是“创建/拥有 TQ 生命周期的进程”。只有 owner 在任务结束时执行 `tq.close()`。 + +关键顺序是: + +```text +SpecoTaskRunner +→ tq.init(完整配置) +→ trainer.init_workers() +→ Ray workers 启动 +``` + +也就是说,PR #48 假设 TQ Controller/storage 已经由 TaskRunner 在 Ray 集群环境中建立,后启动的 worker 只需要连接。 + +#### 3.0.3 每个 Producer/Consumer 进程怎么连接 TQ + +bridge 中的 `_ensure_initialized()` 是进程级懒初始化: + +```python +def _ensure_initialized(): + if _state["initialized"]: + return + + with _state_lock: + if _state["initialized"]: + return + + tq.init() + _state["initialized"] = True +``` + +注意这里是: + +```python +tq.init() +``` + +不是: + +```python +tq.init(config) +``` + +无参初始化的含义是连接 TaskRunner 已创建的同一套 TQ。它依赖 PR #48 所处的 Ray 运行环境完成服务发现。 + +因此 PR #48 不是“每个 worker 各创建一套 TQ”,而是: + +```text +TaskRunner:tq.init(config),创建一次 +SGLang producer:tq.init(),连接 +actor producer:tq.init(),连接 +drafter consumer:tq.init(),连接 +``` + +#### 3.0.4 SGLang Producer 怎么写 hidden states + +SGLang 已经完成 rollout,并组装出 `drafter_sample` 后,PR #48 执行: + +```python +configure_transfer_queue(training_cfg) + +if is_transfer_queue_enabled(): + tq_key = make_sample_key( + collection_global_steps, + self.replica_rank, + request_id, + ) + + tq_payload = { + "hidden_states": hidden_states.unsqueeze(0).cpu(), + } + + if target_logprobs is not None: + tq_payload["target_logprobs"] = ( + target_logprobs.unsqueeze(0).cpu() + ) + + put_sample( + tq_key, + tq_payload, + tag={ + "global_step": collection_global_steps, + "replica_rank": self.replica_rank, + }, + ) + + drafter_sample["hidden_states_tq_key"] = tq_key + drafter_sample["hidden_states"] = None +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/integration/sglang_runtime.py:2020 +``` + +这里发生了两条不同的数据流: + +```text +大 tensor:SGLang → TQ +小字典/key:SGLang → 原 Ray/driver 路径 → drafter worker +``` + +写进 TQ 后将: + +```python +drafter_sample["hidden_states"] = None +``` + +是为了避免 hidden states 继续经过 driver/Ray object store。driver 仍收到 sample,但大 tensor 已替换成: + +```python +drafter_sample["hidden_states_tq_key"] +``` + +#### 3.0.5 `put_sample()` 实际怎么写 + +bridge 中: + +```python +def put_sample(key, tensor_dict, *, tag=None): + payload = { + k: v + for k, v in tensor_dict.items() + if torch.is_tensor(v) + } + + _ensure_initialized() + + tq.kv_put( + key=key, + partition_id="speco_drafter_features", + fields=payload, + tag=tag or {}, + ) +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/integration/transferqueue_bridge.py:206 +``` + +这里可以明确看到: + +- PR #48 使用 TQ 高层 KV API; +- 一个 key 对应一个 sample; +- `fields` 是 tensor 字典; +- `tag` 是小 metadata; +- partition 当前写死为 `speco_drafter_features`; +- 写入前 tensor 已 `.cpu()`; +- 写入失败直接抛异常,不静默回退。 + +key 的生成代码是: + +```python +def make_sample_key(global_step, replica_rank, request_id): + return f"speco:{global_step}:{replica_rank}:{request_id}" +``` + +这个 key 对 RL rollout 是合理的,因为它用 step、rollout replica 和 request ID 标识一次在线采样。 + +#### 3.0.5.1 PR #48 中“数据”和“元数据”分别长什么样 + +PR #48 一条 SGLang 样本在写 TQ 之前,`drafter_sample` 同时包含大 tensor 和控制信息。简化后类似: + +```python +drafter_sample = { + # 普通训练输入,仍走原 sample/Ray 控制路径 + "input_ids": int64[1, prompt_len + response_len], + "prompts": int64[1, prompt_len], + "responses": int64[1, response_len], + + # 大 tensor,开启 TQ 后从这个字典移除 + "hidden_states": bf16[1, hidden_rows, hidden_dim], + "target_logprobs": fp32[1, rows, topk_or_vocab] | None, + + # hidden 与 token 对齐所需的小字段 + "hidden_positions": int64[1, hidden_rows] | None, + "hidden_position_start": int, + "hidden_position_end": int, + "hidden_window_start": int, + "hidden_window_end": int, + + # 控制信息 + "global_step": int, + "replica_rank": int, +} +``` + +执行 `put_sample()` 时,并不是把整个 `drafter_sample` 放进 TQ。PR #48 只抽取占用大的 tensor: + +```python +tq_payload = { + "hidden_states": bf16[1, hidden_rows, hidden_dim], + "target_logprobs": fp32[1, rows, ...], # 可选 + "hidden_raw_target_logprobs": ..., # 可选 + "hidden_raw_target_logprobs_positions": ..., # 可选 +} +``` + +这就是 TQ 的 data payload。它被传给: + +```python +tq.kv_put(fields=tq_payload) +``` + +另外还有 TQ tag: + +```python +tag = { + "global_step": 42, + "replica_rank": 1, +} +``` + +tag 是 TQ 侧轻量 metadata,用于描述/检索对象,不承载 hidden tensor。 + +写入完成后,仍经 Ray 传递的轻量 `drafter_sample` 变成: + +```python +drafter_sample = { + "input_ids": int64[1, total_len], + "prompts": int64[1, prompt_len], + "responses": int64[1, response_len], + + "hidden_states": None, + "target_logprobs": None, + "hidden_states_tq_key": "speco:42:1:req-007", + + "hidden_positions": int64[1, hidden_rows] | None, + "hidden_position_start": 128, + "hidden_position_end": 640, + "global_step": 42, + "replica_rank": 1, +} +``` + +因此 PR #48 实际存在三类对象: + +| 对象 | 内容 | 传输路径 | 作用 | +|---|---|---|---| +| TQ fields/payload | `hidden_states` 等大 tensor | Producer → TQ storage → Consumer | 避免 driver 搬运大 tensor | +| TQ tag | `global_step`、`replica_rank` | TQ control/index metadata | 描述该 key | +| 轻量 `drafter_sample` | tokens、位置、标量 metadata、`hidden_states_tq_key` | 原 Ray driver 路径 | 告诉 Consumer 取哪个 key,以及如何解释取回的 tensor | + +代码实现解耦的关键不是“所有内容都进 TQ”,而是: + +```python +drafter_sample["hidden_states_tq_key"] = tq_key +drafter_sample["hidden_states"] = None +``` + +第一行给 Consumer 留下寻址信息;第二行阻止大 tensor 继续沿旧路径传输。 + +#### 3.0.5.2 Consumer 如何把两部分重新合成一个训练样本 + +Consumer 最初拿到的是轻量 sample: + +```python +sample["hidden_states"] is None +sample["hidden_states_tq_key"] == "speco:42:1:req-007" +``` + +它执行: + +```python +payload = get_sample(sample["hidden_states_tq_key"]) +sample["hidden_states"] = payload["hidden_states"] +``` + +合并后: + +```python +sample = { + "input_ids": ..., + "prompts": ..., + "responses": ..., + "hidden_positions": ..., + "hidden_states": bf16[1, hidden_rows, hidden_dim], + "hidden_states_tq_key": "speco:42:1:req-007", + ... +} +``` + +后面的 `_store_rollout_sample()` 看到的结构与关闭 TQ 时基本一致,所以训练主体不需要增加 TQ 分支。TQ bridge 只改变“大 tensor 从哪里恢复”,不改变 drafter trainer 的输入语义。 + +#### 3.0.6 old-logprob Producer 怎么写 chunk + +PR #48 的另一条 Producer 路径来自 actor old-logprob hidden hook。原来是: + +```python +chunk_ref = ray.put(hidden_chunk) +``` + +开启 TQ 后改成: + +```python +tq_key = make_sample_key( + global_step, + owner, + f"chunk{len(chunk_refs)}", +) + +put_sample( + tq_key, + {"hidden": hidden_chunk}, + tag={ + "global_step": global_step, + "owner": owner, + }, +) + +chunk_ref = tq_key +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/integration/oldlogprob_runtime.py:533 +``` + +后面的 driver 不需要知道 ref 是 Ray ObjectRef 还是 TQ string key,它只把 ref 当不透明 token 继续传递。 + +#### 3.0.7 Consumer 怎么根据 key 读取 + +drafter worker 收到原来的 sample 小字典后: + +```python +tq_key = sample.get("hidden_states_tq_key") + +if tq_key is not None and self._speco_tq_enabled: + payload = get_sample(tq_key) + + for field in ( + "hidden_states", + "target_logprobs", + "hidden_raw_target_logprobs", + "hidden_raw_target_logprobs_positions", + ): + if payload.get(field) is not None: + sample[field] = payload[field] + + if sample.get("hidden_states") is None: + raise RuntimeError( + "TQ key exists but hidden_states payload is missing" + ) +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/workers/speco_worker.py:848 +``` + +恢复 tensor 后,后面的逻辑仍使用原 `sample`/`batch`,drafter trainer 不需要知道 tensor 来自 Ray 还是 TQ。 + +#### 3.0.8 `get_sample()` 实际怎么读 + +```python +def get_sample(key): + _ensure_initialized() + + result = tq.kv_batch_get( + keys=[key], + partition_id="speco_drafter_features", + ) + + value = _extract_value(result, key) + return _tensordict_to_dict(value) +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/integration/transferqueue_bridge.py:237 +``` + +`_extract_value()` 兼容三种返回形态: + +```python +if isinstance(result, dict): + return result.get(key) +if isinstance(result, (list, tuple)): + return result[0] +return result +``` + +这是因为不同 TQ 版本/后端返回包装可能不同。 + +#### 3.0.9 为什么需要 `_densify_tq_tensor()` + +PR #48 后续修复发现,TQ 把单样本 tensor 放进 TensorDict 后,`kv_batch_get` 可能返回 NestedTensor,并额外带 batch 维。旧代码要执行: + +```python +tensor[start:start + length] +``` + +但 NestedTensor 不支持在 jagged dim 直接 slice。因此加入: + +```python +def _densify_tq_tensor(tensor): + if tensor.is_nested: + parts = [ + part + for part in tensor.unbind() + if part.numel() > 0 + ] + tensor = torch.cat(parts, dim=0) + + if tensor.dim() == 3: + tensor = tensor.squeeze(0) + elif tensor.dim() == 1: + tensor = tensor.unsqueeze(0) + + return tensor.contiguous() +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/workers/speco_worker.py:72 +``` + +standalone Consumer 同样必须做这个转换,不能假定 `kv_batch_get` 返回普通 dense `[seq, hidden]`。 + +#### 3.0.10 为什么需要 per-step cache + +old-logprob 路径里,一个 owner hidden chunk 可能被约 16 个 sample 共同引用。如果每个 sample 都: + +```python +get_sample(same_tq_key) +``` + +就会重复传输同一个数百 MB chunk。PR #48 在每次 `collect_rollout_features()` 开始时创建: + +```python +self._tq_chunk_cache = {} +``` + +解析 ref 时: + +```python +cache_key = ref if isinstance(ref, str) else id(ref) + +if cache_key not in cache: + cache[cache_key] = _resolve_tq_or_ray_ref(ref) + +tensor = cache[cache_key] +``` + +对应参考代码: + +```text +../verl-SpeCo/verl_speco/workers/speco_worker.py:98 +../verl-SpeCo/verl_speco/workers/speco_worker.py:854 +``` + +独立训练如果每个 sample 都是独立 TQ key,主要使用 `kv_batch_get(keys=[...])` 批量读取,不一定需要跨 sample chunk cache;但 prefetch 重试或共享 packed object 时仍应保留 key cache。 + +#### 3.0.11 PR #48 什么时候删除数据 + +PR #48 没有在 `get_sample()` 后删除。其注释明确说明:同一 drafter replica 的多个 TP/SP rank 可能读取同一 key,第一次读取后立刻删除会让剩余 rank 失败。 + +当前策略是任务结束时由 owner: + +```python +tq.close() +``` + +统一结束 TQ 生命周期。也就是说,PR #48 当前没有实现精细的逐 step `kv_clear`。 + +#### 3.0.12 PR #48 的完整时序 + +```text +SpecoTaskRunner + → tq.init(config) + → 启动 Ray workers + +SGLang/actor Producer process + → configure_transfer_queue() + → 第一次 put 时 tq.init() + → kv_put(key, tensor fields, tag) + → 把 key 塞回原 sample/ref + +Ray driver + → 只中转小 sample/key + +drafter worker Consumer process + → 第一次 get 时 tq.init() + → kv_batch_get([key]) + → 解包 TensorDict/NestedTensor + → 恢复 sample["hidden_states"] + → 原 drafter collect/train 逻辑 + +任务结束 + → TaskRunner owner tq.close() +``` + +### 3.1 已经实现的可复用能力 + +PR #48 新增 `verl_speco/integration/transferqueue_bridge.py`,锁定思路是把 TQ 当成独立传输库,不修改上游 verl。它提供: + +```python +configure_transfer_queue(training_cfg) +init_transfer_queue(config) +make_sample_key(global_step, replica_rank, request_id) +put_sample(key, tensor_dict, tag=...) +get_sample(key) +close_transfer_queue() +``` + +实际写入调用是: + +```python +tq.kv_put( + key=key, + partition_id="speco_drafter_features", + fields=payload, + tag=tag, +) +``` + +实际读取调用是: + +```python +result = tq.kv_batch_get( + keys=[key], + partition_id="speco_drafter_features", +) +``` + +另外,PR #48 已经处理了多项 standalone 方案也需要的问题: + +1. TQ 返回值可能是 direct value、`{key: value}` 或 list,需要统一解包; +2. TQ 返回的 tensor 可能是 NestedTensor,必须通过 `unbind + cat` 恢复为生产端写入的 dense tensor; +3. 多个样本引用同一 hidden chunk 时,需要 per-step cache,避免重复 `kv_batch_get` 同一个大对象; +4. TQ key 存在但 payload 缺失时 fail loud,不能静默丢样本; +5. `enable=false` 时保留原传输路径。 + +这些逻辑应直接作为本项目 TQ adapter 的参考。 + +### 3.2 PR #48 的数据流 + +PR #48 优化的是 RL online 路径: + +```text +SGLang/actor worker + → kv_put(hidden states) + → 把 hidden_states_tq_key 塞进原 drafter_sample + → 原 Ray driver 继续传递小 sample/key + → drafter worker collect_rollout_features() + → kv_batch_get(key) +``` + +它没有让 consumer 自己从 TQ 发现下一批 key;key 仍沿原来的 Ray 控制路径到达 drafter worker。 + +### 3.3 PR #48 没有提供的 standalone 能力 + +PR #48 当前没有实现: + +- 从预生成 response 文件读取数据的独立 Producer; +- Producer 并行请求外部 vLLM endpoint; +- standalone DSpark trainer 主动发现 ready key; +- global batch 到各 torchrun rank 的分片; +- 每个 optimizer step 后精确 `kv_clear`; +- EOS; +- standalone 无 Ray 的 TQ bootstrap; +- MooncakeStore 的实际运行验证。 + +PR #48 当前配置是: + +```yaml +transfer_queue: + enable: false + backend: + storage_backend: SimpleStorage + SimpleStorage: + total_storage_size: 100000 + num_data_storage_units: 8 +``` + +并且 bridge 注释针对 `TransferQueue==0.1.7`。旁边当前 verl 主线已使用 `0.1.8` 文案并包含 `MooncakeStore` 配置。当前机器没有安装 `transfer_queue` 包,因此正式实现前必须锁定版本并实机验证 API 签名,不能把 0.1.7 和 0.1.8 混用。 + +### 3.4 standalone 方案对 PR #48 的扩展 + +不再自行实现 `publish_ready/claim/lease/ack` HTTP 服务。第一版在 PR #48 KV 模式上补四个操作: + +```python +tq.kv_batch_put(...) # Producer 批量写 +tq.kv_list(...) # rank 0 列出 key + tag +tq.kv_batch_get(...) # 各 rank 并行读 +tq.kv_clear(...) # optimizer step 成功后删 +``` + +第一版不依赖 TQ `Sampler/StreamingDataLoader/get_meta`,因为 PR #48 并未使用或验证这些接口。等 KV 流程稳定后再升级为 TQ StreamingDataLoader。 + +### 3.5 PR #48 与 standalone 独立训练逐项映射 + +| PR #48 online RL | standalone drafter training | +|---|---| +| `SpecoTaskRunner` 调用 `tq.init(config)` | 新增独立 TQ owner/bootstrap 进程,或在确认安全后由 Producer owner 调用 `tq.init(config)` | +| SGLang/actor worker 是 Producer | `feature_producer.py` 是独立 Producer | +| rollout 过程中已经得到 hidden states | Producer 读取预生成 response,再请求外部 vLLM prefill | +| `make_sample_key(global_step, replica_rank, request_id)` | `make_replay_sample_key(dataset,row,tokens,target fingerprint)` | +| `put_sample()` / `kv_put()` | 继续复用同一写入模式,可扩展为 `kv_batch_put()` | +| key 通过 Ray driver/sample 传给 consumer | 没有 driver;rank 0 用 `kv_list()` 主动发现 ready keys | +| drafter worker `get_sample(key)` | 每个 torchrun rank `kv_batch_get(local_keys)` | +| `collect_rollout_features()` 恢复 sample | 转换为 `DraftFeatureSample` 后调用现有 `prepare_training_batch_from_samples()` | +| 多 TP/SP rank 可能读同一 key,所以不立即删除 | data-parallel rank 读取互不重叠的 key;全 rank step 成功后统一 `kv_clear(global_keys)` | +| TaskRunner 结束时 `tq.close()` | 每 step clear;输入 drain 完成后 owner 最后 `tq.close()` | + +standalone 需要新增的控制流是: + +```text +Producer DSpark rank 0 其他 ranks + │ │ │ + │ kv_put(sample key, fields, tag) │ │ + ├────────────────────────────────────▶│ │ + │ │ kv_list READY keys │ + │ │ │ + │ │ broadcast selected_keys ──▶│ + │ │ │ + │ │ kv_batch_get(local keys) │ kv_batch_get(local keys) + │ │ │ + │ ├──── DSpark synchronized step ────┤ + │ │ │ + │ │ kv_clear(global keys) │ +``` + +这个映射中,TQ 同时承担: + +- 大 tensor 存储/传输; +- key、tag 和 partition 的轻量索引。 + +但第一版 global batch 的选择仍由单个 DSpark job 的 rank 0 完成。这样最接近 PR #48 的 KV API,避免在同一次改造中再引入未经该 PR 验证的 Sampler/StreamingDataLoader。 + +### 3.6 PR #48 的 SGLang 路径:逐函数、逐对象完整流程 + +下面从一次生成请求开始,不省略中间层。 + +#### 阶段 1:SGLang完成生成并收集 hidden states + +执行进程:SGLang rollout server。 + +输入是一次 request 对应的 prompt 和生成配置。生成结束时,代码已经持有: + +```python +prompt_tensor: int64[prompt_len] +response_tensor: int64[response_len] +hidden_states: bf16[hidden_rows, hidden_dim] +hidden_positions: int64[hidden_rows] | None +target_logprobs: tensor | None +request_id: str +collection_global_steps: int +self.replica_rank: int +``` + +这些变量的语义: + +- `prompt_tensor`:输入 prompt token IDs; +- `response_tensor`:SGLang生成的 response token IDs; +- `hidden_states`:target model 指定层在部分 token positions 上的输出; +- `hidden_positions`:每个 hidden row 对应完整 `prompt+response` 序列中的哪个 token position; +- `target_logprobs`:可选的目标概率监督; +- `request_id`:当前 rollout request 标识; +- `replica_rank`:执行该 request 的 rollout replica。 + +SGLang 先构造完整 sample: + +```python +drafter_sample = { + "input_ids": torch.cat( + [prompt_tensor, response_tensor], dim=0 + ).unsqueeze(0), + "prompts": prompt_tensor.unsqueeze(0), + "responses": response_tensor.unsqueeze(0), + "hidden_states": hidden_states.unsqueeze(0).cpu(), + "hidden_positions": hidden_positions.unsqueeze(0).cpu(), + "target_logprobs": ( + target_logprobs.unsqueeze(0).cpu() + if target_logprobs is not None + else None + ), + "global_step": collection_global_steps, + "replica_rank": self.replica_rank, + # 还有 hidden window/alignment metadata +} +``` + +前导维 `1` 表示这是一个单样本 batch。`.cpu()` 表示跨进程传输前把大 tensor 放到 CPU 内存。 + +#### 阶段 2:PR #48 将大 fields 写入 TQ + +同一个 SGLang进程执行: + +```python +tq_key = make_sample_key( + collection_global_steps, + self.replica_rank, + request_id, +) + +tq_payload = { + "hidden_states": hidden_states.unsqueeze(0).cpu(), +} + +put_sample( + tq_key, + tq_payload, + tag={ + "global_step": collection_global_steps, + "replica_rank": self.replica_rank, + }, +) +``` + +调用展开后是: + +```python +tq.init() # 当前进程第一次使用时 +tq.kv_put( + key=tq_key, + partition_id="speco_drafter_features", + fields=tq_payload, + tag=tag, +) +``` + +效果是 TQ 中增加一行: + +```text +partition = speco_drafter_features +key = speco:42:1:req-007 +fields = {hidden_states: bf16[1, H, D], ...} +tag = {global_step: 42, replica_rank: 1} +``` + +`kv_put` 返回后,SGLang侧将旧 sample 改成: + +```python +drafter_sample["hidden_states_tq_key"] = tq_key +drafter_sample["hidden_states"] = None +``` + +此后该 Python 字典不再携带 hidden tensor,只携带定位它的 key。 + +#### 阶段 3:把 drafter_sample 放进 TokenOutput.extra_fields + +SGLang返回: + +```python +TokenOutput( + token_ids=token_ids, + log_probs=log_probs, + routed_experts=routed_experts, + extra_fields={ + "global_steps": collection_global_steps, + "drafter_sample": drafter_sample, + }, +) +``` + +此时 `TokenOutput` 中有两类输出: + +- 正常 rollout 输出:`token_ids/log_probs`; +- SpeCo 训练旁路输出:`extra_fields.drafter_sample`。 + +TQ payload 不在 `TokenOutput` 中,只有 `hidden_states_tq_key` 在其中。 + +#### 阶段 4:多个 TokenOutput 聚合成 gen_batch_output + +rollout/agent-loop 层将多个 request 的输出合并为批量 `DataProto`: + +```python +gen_batch_output.non_tensor_batch["drafter_sample"] +``` + +可能是 object array: + +```python +array([ + {"hidden_states_tq_key": "speco:42:0:req-A", ...}, + {"hidden_states_tq_key": "speco:42:1:req-B", ...}, +], dtype=object) +``` + +之所以进入 `non_tensor_batch`,是因为每条 sample 的 hidden window、Python metadata 和可选字段不一定具有统一 dense shape。 + +#### 阶段 5:driver 从 DataProto 取出 drafter samples + +`generate_sequences_with_speco()` 包装原 rollout 调用: + +```python +gen_batch_output = original_generate_sequences(...) +collected = self._speco_collect_generation_samples(gen_batch_output) +``` + +`_speco_collect_generation_samples()` 调用: + +```python +samples = pop_drafter_samples(gen_batch_output) +``` + +`pop_drafter_samples()` 实际执行: + +```python +non_tensor_batch = gen_batch_output.non_tensor_batch +samples_array = non_tensor_batch.pop("drafter_sample", None) +samples = normalize_drafter_samples(samples_array) +``` + +这里 `pop` 有两个作用: + +1. 取得 SpeCo drafter side-channel samples; +2. 从正常 PPO 的 `gen_batch_output` 中移除该旁路字段,避免后续 PPO batch 继续携带它。 + +`normalize_drafter_samples()` 将 dict、object array 或 list 统一成: + +```python +samples: list[dict] +``` + +#### 阶段 6:driver 按 replica_rank 分桶 + +假设有两个 rollout/drafter replicas,收到: + +```python +samples = [ + {"replica_rank": 1, "hidden_states_tq_key": "k1", ...}, + {"replica_rank": 0, "hidden_states_tq_key": "k2", ...}, + {"replica_rank": 1, "hidden_states_tq_key": "k3", ...}, +] +``` + +执行: + +```python +buckets = bucket_drafter_samples_by_replica( + samples, + num_replicas=2, +) +``` + +结果: + +```python +buckets = [ + [sample_k2], # bucket 0 + [sample_k1, sample_k3], # bucket 1 +] +``` + +分桶依据只有: + +```python +owner_rank = int(sample["replica_rank"]) +buckets[owner_rank].append(sample) +``` + +这一步没有读取 TQ,也没有处理 hidden tensor;只对小字典做路由。 + +#### 阶段 7:driver 通过 WorkerGroup RPC 分发 buckets + +driver 调用: + +```python +self._speco_set_drafter_global_step() +self._speco_collect_rollout_features_rpc( + "rollout", + buckets, +) +``` + +RPC 内部调用: + +```python +self.drafter_wg.collect_rollout_features(buckets) +``` + +因为 worker 方法注册了 `drafter_owner_route` dispatch,WorkerGroup 将 `buckets[0]` 发给 drafter DP replica 0,将 `buckets[1]` 发给 drafter DP replica 1。一个训练 replica 内的 SP ranks 根据 mesh dispatch 规则参与对应调用。 + +这里传输的对象仍是: + +```python +list[dict] +``` + +其中包含 tokens、position metadata 和 TQ key,不包含被置空的 hidden states。 + +#### 阶段 8:SpecoWorker 根据 key 从 TQ 恢复 tensor + +目标 worker 执行: + +```python +def collect_rollout_features(self, samples): + for sample in samples: + tq_key = sample.get("hidden_states_tq_key") + payload = get_sample(tq_key) + sample["hidden_states"] = payload["hidden_states"] +``` + +`get_sample()` 展开为: + +```python +tq.init() # 此 Consumer 进程第一次使用时 +result = tq.kv_batch_get( + keys=[tq_key], + partition_id="speco_drafter_features", +) +payload = _extract_value(result, tq_key) +payload = _tensordict_to_dict(payload) +``` + +现在 `sample` 再次包含: + +```python +{ + "input_ids": ..., + "hidden_positions": ..., + "hidden_states": bf16[1, hidden_rows, hidden_dim], + "hidden_states_tq_key": "...", +} +``` + +这与关闭 TQ 时 worker 收到的逻辑内容一致。 + +#### 阶段 9:worker 构造 DrafterBaseTrainer 所需 batch dict + +worker 先保留 token fields: + +```python +batch = { + "input_ids": sample["input_ids"], + "prompts": sample["prompts"], + "responses": sample["responses"], +} +``` + +再复制 hidden alignment metadata,例如: + +```python +batch["hidden_positions"] +batch["hidden_position_start"] +batch["hidden_position_end"] +batch["hidden_states_layout"] +batch["global_step"] +``` + +hidden tensor 单独作为参数: + +```python +self._store_rollout_sample( + batch=batch, + hidden_states=hidden, + target_logprobs=target_logprobs, +) +``` + +#### 阶段 10:样本进入在线 buffer 或落盘 + +`_store_rollout_sample()` 根据 training mode 分支: + +```python +if mode == "collect_only": + self._write_rollout_feature_sample( + batch, + hidden_states, + target_logprobs, + ) +else: + self.trainer.collect_online_data( + batch, + hidden_states, + target_logprobs, + ) +``` + +`collect_only` 会转换成 `DraftFeatureSample` 并写 `TorchShardFeatureStore`。online 模式则进入 `DrafterBaseTrainer.collect_online_data()`。 + +`collect_online_data()` 做: + +1. 将 `input_ids/hidden_states/positions/logprobs` 规范化到 CPU; +2. 按 batch 维拆成逐样本; +3. 根据 `hidden_positions` 校验 hidden row 与 token position; +4. 截取可训练窗口; +5. 构造内部 training item; +6. 保存到当前 step 的 `collected_data`,或在启用 data buffer 时保存到跨 step buffer。 + +因此 TQ get 完成不代表立即 optimizer step。它先恢复 online training sample,再进入现有数据准备逻辑。 + +#### 阶段 11:driver 在 actor update 周期触发 drafter 训练 + +driver 包装了 `update_actor()`: + +```python +should_train_drafter = ( + self._speco_should_attempt_drafter_train_this_step() +) + +actor_output = original_update_actor(...) + +if should_train_drafter: + drafter_trained, metrics = self._speco_train_drafter() +``` + +`_speco_train_drafter()` 再向 WorkerGroup 发: + +```python +self.drafter_wg.train_drafter() +``` + +每个 `SpecoWorker.train_drafter()`: + +1. 检查是否属于 drafter training group; +2. 检查 `training_interval_steps`; +3. 激活 drafter training model; +4. 循环 `train_steps_per_trigger` 次; +5. 每次调用 `self.trainer.training_step(global_step)`; +6. 成功时准备需要发布的 drafter state dict; +7. 清理训练期间临时状态。 + +`training_step()` 从刚才的 online `collected_data/DataBuffer` 组成 batch,执行 drafter forward、loss、backward 和 optimizer step。 + +所以 SGLang TQ 路径的最终效果是: + +```text +TQ 只替换 hidden tensor 跨进程传输 +→ sample 收集逻辑不变 +→ online buffer 不变 +→ drafter training trigger 不变 +→ loss/optimizer 不变 +``` + +### 3.7 PR #48 old-logprob 路径的完整差异 + +old-logprob 路径没有 `TokenOutput.extra_fields`。它从 PPO 的 actor old-logprob forward 开始。 + +#### 阶段 1:driver 构造 collect plan + +driver 根据 batch、collect interval 和 drafter owner 数量决定: + +```python +collect_mask: bool[batch] +hidden_positions: list/tensor per sample +owner_rank: int64[batch] +prompt_lens: int64[batch] +response_lens: int64[batch] +``` + +并把 `global_step` 等控制字段放入 old-logprob micro-batch。 + +#### 阶段 2:actor forward hook 选择 hidden rows + +actor worker 在 old-logprob forward 中捕获指定层 hidden states,根据 `collect_mask/position_mask` 只保留需要训练的样本和 token rows。 + +输出可以是 dense selected tensor,也可以是 sparse rows。随后 `_put_oldlogprob_hidden_refs()` 将同一个 owner 的多条 sample rows 拼成一个较大的 `hidden_chunk`。 + +#### 阶段 3:hidden chunk 写入 TQ + +改造前: + +```python +chunk_ref = ray.put(hidden_chunk) +``` + +PR #48: + +```python +tq_key = make_sample_key( + global_step, + owner, + f"chunk{chunk_index}", +) + +put_sample( + tq_key, + {"hidden": hidden_chunk}, + tag={"global_step": global_step, "owner": owner}, +) + +chunk_ref = tq_key +``` + +TQ fields: + +```python +{"hidden": bf16[total_owner_rows, hidden_dim]} +``` + +控制路径中的 chunk metadata: + +```python +chunk_info = { + "sample_indices": [0, 3, 5], + "starts": [0, 128, 384], + "lengths": [128, 256, 96], + "row_indices": [...], + "dtype": "bfloat16", + "shape": [480, hidden_dim], +} +``` + +`starts/lengths` 描述每条 sample 在共享 chunk 中对应的行区间。 + +#### 阶段 4:driver 将 chunk ref 还原为逐样本引用 + +driver 的 `_speco_collect_oldlogprob_features()` 读取: + +```python +chunk_refs = ["speco:42:0:chunk0", ...] +chunk_meta = [chunk_info, ...] +``` + +然后为每个 batch sample 构造: + +```python +sample["hidden_states_ref_chunks"] = [ + { + "ref": "speco:42:0:chunk0", + "chunk_start": 128, + "chunk_length": 256, + "chunk_row_indices": ..., + "dtype": "bfloat16", + "shape": [480, hidden_dim], + } +] +``` + +同时构造该 sample 的: + +```python +input_ids +prompts +responses +hidden_positions +hidden_states_layout +replica_rank=owner +``` + +再按 owner 放入 `buckets[owner]`,通过同一个 `collect_rollout_features()` RPC 发给 drafter worker。 + +#### 阶段 5:Consumer 获取共享 chunk 并切片 + +drafter worker 发现: + +```python +sample.get("hidden_states") is None +sample.get("hidden_states_ref_chunks") is not None +``` + +于是调用 `_resolve_hidden_state_chunks()`。对字符串 ref: + +```python +if ref.startswith("speco:"): + full_chunk = get_sample(ref)["hidden"] + full_chunk = _densify_tq_tensor(full_chunk) +``` + +然后按 sample metadata 取行: + +```python +sample_hidden = full_chunk[ + chunk_start : chunk_start + chunk_length +] +``` + +同一个 chunk 被多个 sample 复用,所以使用: + +```python +self._tq_chunk_cache[ref] = full_chunk +``` + +保证一次 `collect_rollout_features()` 中同一个 TQ key 只 get 一次。 + +得到逐样本 hidden 后,后续 `_store_rollout_sample → collect_online_data → train_drafter` 与 SGLang 路径相同。 + +### 3.8 PR #48 数据生命周期和清理 + +PR #48 的 TQ row 生命周期是: + +```text +TaskRunner tq.init(config) +→ Producer kv_put +→ key 经 Ray 控制路径传递 +→ 一个或多个 drafter TP/SP rank kv_batch_get +→ online drafter 收集/训练继续执行 +→ 整个 trainer.fit() 结束 +→ TaskRunner finally 调用 tq.close() +``` + +当前没有: + +```python +tq.kv_clear(key) +``` + +原因是同一 sample/key 可能被一个 drafter replica 的多个 rank 读取。若第一个 rank get 后删除,其余 rank 可能 get 失败。 + +因此 PR #48 采用任务级生命周期,而不是样本级确认和回收。这简化了并发正确性,但意味着长任务中的 TQ storage 占用可能持续增长;代码注释也把“leader 在 barrier 后精细 clear”留作后续工作。 + +### 3.9 PR #48 开启与关闭时的行为差异 + +`configure_transfer_queue()` 返回: + +```python +enabled_in_config and transfer_queue_importable +``` + +关闭时: + +```text +SGLang drafter_sample 继续内联 hidden_states +old-logprob 继续 ray.put(hidden_chunk) +Consumer 继续 ray.get/ref resolve +``` + +开启时: + +```text +SGLang hidden fields → TQ,sample 只带 key +old-logprob hidden chunk → TQ,ref 变成字符串 key +Consumer 根据 key 类型走 TQ get +``` + +如果配置开启但 `transfer_queue` 包未安装,bridge 会记录 warning,并让 `is_transfer_queue_enabled()` 返回 false,保留旧 Ray 路径。若已经进入 `put_sample/get_sample` 却发生 TQ 错误,则抛异常,不静默丢 hidden data。 + +## 第二部分:基于 PR #48 的 standalone drafter training 适配 + +从这一部分开始才讨论 `verl-SpeCo-ls` 的独立训练。下面的 `kv_list`、rank 0 选 global keys、逐 step `kv_clear`、独立 vLLM Producer 都是需要在本项目新增的逻辑,不是 PR #48 已有逻辑。 + +### 当前 standalone 基线 + +当前独立训练是: + +```text +draft_train_launcher +→ torch.distributed.run +→ 每个 rank 创建 DraftFeatureDataLoader +→ 每个 rank 在训练循环内同步 TargetFeatureReplayer.materialize() +→ vLLM/file hidden payload +→ prepare_training_batch_from_samples() +→ training_step_from_batch() +``` + +新方案要把 `TargetFeatureReplayer.materialize()` 从训练 rank 的同步取数路径移到独立 Producer,同时保持后两步训练接口不变。 + +## 4. TQ metadata 到底记录什么 + +### 4.1 Partition + +一次训练运行使用一个独立 partition: + +```python +partition_id = f"speco:{run_id}:dspark_train" +``` + +partition 用来隔离: + +- 不同训练 run; +- train 和 validation; +- 不同 target checkpoint 生成的 hidden states。 + +不能让两个 target 模型共用同一 partition,否则训练侧可能消费错误的 hidden states。 + +### 4.2 Sample key + +每条输入样本使用稳定 key: + +```python +sample_key = sha256( + dataset_id + + row_id + + prompt_token_ids + + response_token_ids + + tokenizer_fingerprint + + target_model_fingerprint + + target_layer_ids + + hidden_states_layout +).hexdigest() +``` + +稳定 key 用于: + +- vLLM HTTP 请求重试时不生成不同对象; +- Producer 重启后识别相同样本; +- 检查 hidden states 是否属于正确模型和正确层; +- TQ/Mooncake 清理时准确定位对象。 + +### 4.3 Fields 与 READY 约定 + +每个样本包含固定字段: + +```python +{ + "input_ids": int64[seq], + "loss_mask": float32[seq], + "position_ids": int64[seq], + "hidden_states": bf16[seq, aux_hidden_dim or aux_hidden_dim+hidden_size], +} +``` + +这里必须保持当前 `DraftFeatureSample` 的契约:`TargetFeatureReplayer._feature_from_vllm_payload()` 会把多个 aux layers flatten;当 layout 是 `dflash_aux_plus_last` 时,还会把 final hidden 拼到同一个 `hidden_states` tensor 尾部。`DSparkTrainerBackend.preprocess_individual_items()` 再根据 metadata 中的 `hidden_states_layout` 将 final hidden 切出来,生成训练 batch 的 `target_last_hidden_states`。 + +因此 TQ 不需要新增一个独立 `target_last_hidden_states` field。开启 DSpark L1 时,要求: + +```text +metadata.hidden_states_layout = dflash_aux_plus_last +hidden_states.shape[-1] = num_context_layers * hidden_size + hidden_size +``` + +完整 TQ native metadata API 可以做字段级 ready 判定,但 PR #48 当前走的是高层 KV API:一次 `kv_put` 把一个样本的多个 tensor fields 一起写入。因此第一版采用更直接的约定: + +```python +required_fields = [ + "input_ids", + "loss_mask", + "position_ids", + "hidden_states", +] +``` + +Producer 先在内存中验证所有必需字段,再进行一次 `kv_put`。只有 `kv_put` 成功返回,key 才会出现在 `kv_list` 结果中,并带有: + +```python +tag={ + "status": "ready", + "run_id": run_id, + "sample_id": sample_key, +} +``` + +Consumer 只选择 `status=ready` 且 `run_id` 匹配的 key。不要先 put `input_ids`、再单独 put `hidden_states`,否则 Consumer 可能观察到半成品。 + +### 4.4 Tags + +tags 是轻量 metadata,不放大 tensor: + +```python +tags = { + "sample_id": sample_key, + "source_row": row_id, + "seq_len": seq_len, + "payload_bytes": payload_bytes, + "target_model_fp": target_model_fingerprint, + "target_layers": "8,16,24", + "hidden_layout": "dflash_aux_plus_last", + "producer_status": "success", +} +``` + +tags 可用于过滤、监控、背压统计和错误排查。它不能替代 tensor shape/dtype 校验。 + +### 4.5 Run ID,而不是先依赖 task_name + +PR #48 的 `kv_put/kv_batch_get` 路径没有使用 `task_name` 或 Sampler consumption history。第一版按它的已验证接口,在 partition 和 tags 中放 `run_id`: + +```python +partition_id = "speco_drafter_features" +tag = { + "run_id": run_id, + "status": "ready", +} +``` + +不同 run 最好直接使用不同 partition: + +```python +partition_id = f"speco_drafter_features_{run_id}" +``` + +这样 job 重启和清理更简单。`task_name=dspark_train` 留到后续迁移 TQ native metadata/Sampler 时再使用。 + +### 4.6 standalone 中一条样本的完整对象形态 + +standalone 没有 PR #48 的轻量 `drafter_sample → Ray driver` 路径,因此需要让 TQ tag 承担“如何找到和解释 payload”的 metadata 作用。 + +#### Producer 读到的原始记录 + +```python +source_record = { + "dataset_id": "math-train", + "row_id": 12345, + "prompt": "...", + "response": "已经提前生成的 response", +} +``` + +#### Token replay 样本 + +分词和对齐后: + +```python +replay_sample = DraftReplaySample( + input_ids=int64[full_seq], + loss_mask=float32[full_seq], + position_ids=int64[full_seq], + feature_positions=int64[feature_rows], + draft_position_ids=int64[feature_rows], + metadata={ + "dataset_id": "math-train", + "row_id": 12345, + }, +) +``` + +这里的 `full_seq` 是 prompt 与预生成 response 拼接后的长度;`feature_positions` 指明哪些 token 位置最终进入 drafter 监督样本。 + +#### vLLM 返回的原始 hidden payload + +当前文件协议要求 safetensors 至少包含: + +```python +vllm_payload = { + "token_ids": int64[prefill_rows], + "hidden_states": bf16[prefill_rows, returned_layers, hidden_size], +} +``` + +这还不能直接给 DSpark。Producer 应复用当前: + +```python +TargetFeatureReplayer._feature_from_vllm_payload(...) +``` + +完成 token 校验、position 对齐、选层和 flatten。 + +#### Producer 最终得到的 DraftFeatureSample + +```python +feature = DraftFeatureSample( + algorithm="DSpark", + input_ids=int64[feature_rows], + loss_mask=float32[feature_rows], + position_ids=int64[feature_rows], + hidden_states=bf16[feature_rows, feature_hidden_dim], + metadata={ + "hidden_states_layout": "dflash_aux_plus_last", + "target_layer_ids": [8, 16, 24], + "target_model_path": "...", + "target_config_fingerprint": "...", + "feature_start": 128, + "feature_end": 640, + "sequence_length": 512, + }, +) +``` + +若有 3 个 aux layer、target hidden size 为 4096,并包含 final hidden: + +```text +feature_hidden_dim = 3 * 4096 + 4096 = 16384 +hidden_states.shape = [feature_rows, 16384] +``` + +前 `12288` 维是 aux/context hidden,最后 `4096` 维是 DSpark L1 所需 final hidden。现有 DSpark backend 会根据 `dflash_aux_plus_last` 自动拆分。 + +#### 写入 TQ 的 data fields + +第一版建议一个 key 对应一个已经规范化完成的 `DraftFeatureSample`: + +```python +tq_fields = { + "input_ids": feature.input_ids.cpu(), + "loss_mask": feature.loss_mask.cpu(), + "position_ids": feature.position_ids.cpu(), + "hidden_states": feature.hidden_states.cpu(), +} +``` + +这里只放 tensor,因为 PR #48 的 `put_sample()` 会过滤非 tensor: + +```python +payload = { + key: value + for key, value in tensor_dict.items() + if torch.is_tensor(value) +} +``` + +#### 写入 TQ 的 tag metadata + +```python +tq_tag = { + "run_id": "run-20260818-001", + "status": "ready", + "sample_id": sample_key, + "sequence_no": 12345, + "algorithm": "DSpark", + "hidden_states_layout": "dflash_aux_plus_last", + "target_model_fingerprint": "sha256:...", + "target_layer_ids": "8,16,24", + "feature_rows": 512, + "hidden_dim": 16384, + "payload_bytes": 16777216, +} +``` + +tag 中只放 TQ 版本支持序列化的小标量/字符串。列表等复杂对象可以编码为稳定字符串或 JSON。`payload_bytes` 用于背压统计。 + +#### TQ 中逻辑上保存的 row + +```text +partition: speco_drafter_features_run-20260818-001 +key: 86a4...ef2 + +fields: + input_ids → int64[512] + loss_mask → float32[512] + position_ids → int64[512] + hidden_states → bf16[512, 16384] + +tag: + status → ready + sequence_no → 12345 + hidden_layout → dflash_aux_plus_last + target_model_fp → sha256:... +``` + +#### Consumer 恢复出的对象 + +rank 0 通过 `kv_list` 同时取得 key 和 tag;每个 rank 用 local keys 调用 `kv_batch_get` 取得 fields,然后组合: + +```python +feature = DraftFeatureSample( + algorithm=tag["algorithm"], + input_ids=densify(fields["input_ids"]).reshape(-1), + loss_mask=densify(fields["loss_mask"]).reshape(-1), + position_ids=densify(fields["position_ids"]).reshape(-1), + hidden_states=densify(fields["hidden_states"]), + metadata={ + "hidden_states_layout": tag["hidden_states_layout"], + "target_model_fingerprint": tag["target_model_fingerprint"], + }, +) + +feature.validate(strict=True) +``` + +这样传给: + +```python +trainer.prepare_training_batch_from_samples([feature, ...]) +``` + +的数据结构,与现有文件 feature store 读出的 `DraftFeatureSample` 一致。也就是说,TQ 替换的是存取介质和调度方式,不改变 DSpark backend 的样本契约。 + +## 5. 新的整体架构 + +```text + 小 metadata + ┌──────────────────────────┐ + │ TransferQueueController │ + │ KV metadata / key / tags │ + │ partition / storage map │ + └────────────┬─────────────┘ + │ +JSONL/token replay │ + │ │ + ▼ │ +Feature Producer │ + ├─ tokenizer/window │ + ├─ asyncio bounded concurrency │ + ├─ vLLM endpoint pool │ + ├─ validate/pack │ + └─ TQ put ─────────────────────┤ + ▼ + TQ Mooncake backend + hidden-state tensors + │ + ┌───────────────────┼───────────────────┐ + ▼ ▼ ▼ + DSpark rank 0 DSpark rank 1 DSpark rank N + TQ get TQ get TQ get + └───────────────────┼───────────────────┘ + ▼ + synchronized optimizer step + │ + ▼ + TQ clear after success +``` + +大 tensor 的路径是: + +```text +vLLM/Producer memory → TQ Mooncake backend → each training rank +``` + +不会走: + +```text +Mooncake → rank 0 → rank 1/2/3 +``` + +rank 0 最多只广播 `kv_list` 得到的 key 字符串列表;tags 只在选 batch 时由 rank 0 使用。 + +### 5.1 standalone 每一步为什么能实现推理和训练异步 + +#### 步骤 A:Producer 独立推进输入 cursor + +Producer 自己维护: + +```python +reader_cursor = 12346 +``` + +它不等待 Trainer 请求样本。只要 TQ ready bytes 没超过背压上限,就继续读取文件并创建 vLLM task。 + +效果是 Producer 的执行进度与 `optimizer_step` 解耦: + +```text +Producer sequence_no: 1200,1201,1202,... +Trainer optimizer_step: 87 +``` + +两者通过 TQ 中的 ready rows 衔接,不互相直接调用。 + +#### 步骤 B:并发 vLLM task 完成顺序可以乱序 + +例如 Producer 同时提交: + +```text +sequence_no 100 → endpoint 0 +sequence_no 101 → endpoint 1 +sequence_no 102 → endpoint 0 +``` + +完成顺序可能是: + +```text +101 → 100 → 102 +``` + +每个 task 完成后独立执行 `kv_put`,所以慢请求不会阻塞已经完成的请求写入。tag 中的 `sequence_no` 保留原数据顺序。 + +#### 步骤 C:`kv_put` 成功是 READY 可见性的边界 + +Producer 在调用前已经得到完整 `DraftFeatureSample`。一次 `kv_put` 写入该 sample 的全部 tensor fields,并在 tag 中标记 `status=ready`。 + +因此 Consumer 的判断规则是: + +```text +kv_list 能列出该 key +且 tag.run_id 匹配 +且 tag.status == ready +→ 可以尝试 kv_batch_get +``` + +Consumer 仍需对取回 fields 做完整性校验;tag 是调度 metadata,不是正确性证明。 + +#### 步骤 D:rank 0 只负责选 key + +rank 0 执行: + +```python +entries = list_ready_keys() +selected = sorted(entries, key=sequence_no)[:global_batch_size] +``` + +这一步处理的数据只是: + +```python +[ + {"key": "k100", "sequence_no": 100, ...}, + {"key": "k101", "sequence_no": 101, ...}, +] +``` + +不包含 `[seq, hidden_dim]` hidden tensor,所以 rank 0 不成为大数据中转瓶颈。 + +#### 步骤 E:广播保证所有 rank 对同一个 global step 达成一致 + +所有 rank 调用同一次: + +```python +dist.broadcast_object_list(holder, src=0) +``` + +广播结束后,每个 rank 看到完全相同的 global key list。然后按确定性区间切分: + +```text +rank 0: keys[0:per_rank] +rank 1: keys[per_rank:2*per_rank] +... +``` + +这样不会出现 rank 0 训练 batch A、rank 1 训练 batch B 数量不同,或者某个 rank 没进入 backward 的情况。 + +#### 步骤 F:各 rank 直接读取 Mooncake 后端 + +每个 rank 执行: + +```python +tq.kv_batch_get(keys=local_keys, partition_id=partition_id) +``` + +TQ 根据 key 找到 storage backend 中的数据。使用 MooncakeStore 时,大 tensor 数据路径是 Mooncake → 本 rank;rank 0 不读取其他 rank 的 local payload。 + +因此: + +```text +控制面:rank 0 → broadcast small keys +数据面:Mooncake → each rank directly +``` + +#### 步骤 G:恢复现有 DraftFeatureSample 契约 + +每个 rank 将 TQ fields 与 tag 合并、densify、校验,得到 `list[DraftFeatureSample]`。从这里开始继续执行项目现有代码: + +```python +batch = trainer.prepare_training_batch_from_samples( + materialized_samples, + step=optimizer_step, +) + +ok = await trainer.training_step_from_batch( + batch, + optimizer_step, +) +``` + +所以 TQ 不进入 DSpark model/backend 内部,训练数学逻辑不变。 + +#### 步骤 H:全 rank 成功以后才能清理 + +每个 rank 的 `ok` 通过现有 `_all_ranks_true()` 聚合: + +```text +rank 0 ok = true +rank 1 ok = true +rank 2 ok = true +rank 3 ok = true +→ global_ok = true +``` + +只有此时 rank 0 执行: + +```python +tq.kv_clear(keys=global_keys, partition_id=partition_id) +``` + +这样保证被删除的数据已经参与完成的 optimizer step。若任何 rank get/OOM/backward 失败,不执行 clear,便于作业失败后的诊断或恢复。 + +#### 步骤 I:异步重叠如何形成 + +时间线上: + +```text +时间 ─────────────────────────────────────────▶ + +Producer: vLLM(batch N+1) ─ put ─ vLLM(batch N+2) ─ put +Trainer: get(batch N) ─ train(batch N) ─ get/train(batch N+1) +``` + +Producer 和 Trainer 是不同进程,互相不调用;TQ ready rows 是缓冲区。因此 vLLM prefill、网络传输和 DSpark GPU 训练可以重叠。背压只在缓冲区达到容量上限时暂停 Producer。 + +## 6. Producer:读取预生成 response 并并行请求 vLLM + +### 6.1 输入处理 + +Producer 从现有 JSONL/token replay 数据源读取: + +```python +sample = { + "row_id": "12345", + "prompt": "...", + "response": "提前生成好的文本", +} +``` + +构造: + +```python +prompt_ids = tokenizer.encode(sample["prompt"]) +response_ids = tokenizer.encode(sample["response"]) +input_ids = prompt_ids + response_ids +``` + +同时产生: + +```python +loss_mask +position_ids +feature_positions +sample_key +``` + +### 6.2 有界并发 + +不能按样本串行请求: + +```python +for sample in samples: + result = request_vllm(sample) +``` + +改成: + +```python +async def run_producer(samples): + semaphore = asyncio.Semaphore(max_inflight_requests) + + async def run_one(sample): + async with semaphore: + result = await vllm_pool.prefill(sample) + feature = validate_and_pack(sample, result) + await tq_transport.put(feature) + + async with asyncio.TaskGroup() as group: + for sample in samples: + group.create_task(run_one(sample)) +``` + +`max_inflight_requests` 是 Producer 同时未完成的请求数量,不是 vLLM batch size。vLLM 服务端仍会对同时到达的请求做自己的 continuous batching。 + +### 6.3 多 endpoint + +多个 endpoint 例如: + +```yaml +vllm_endpoints: + - http://node0:8000/v1 + - http://node1:8000/v1 + - http://node2:8000/v1 +``` + +调度器维护每个 endpoint 的 inflight 数: + +```python +endpoint = min( + endpoints, + key=lambda item: item.inflight, +) +``` + +请求前 `inflight += 1`,在 `finally` 中 `inflight -= 1`。失败只重试对应样本,不阻塞全部 Producer。 + +### 6.4 当前 vLLM 文件桥接 + +当前客户端协议期望: + +```python +response.kv_transfer_params["hidden_states_path"] +``` + +所以第一阶段仍然是: + +```text +vLLM 写临时 safetensors +→ Producer load_file +→ 校验 token_ids/hidden_states +→ TQ put 到 Mooncake backend +→ TQ put 成功后删除临时文件 +``` + +删除必须发生在 TQ put 成功之后: + +```python +path = request_vllm_hidden_file(sample) +try: + feature = load_and_validate(path) + await tq_transport.put(feature) +finally: + if put_succeeded: + Path(path).unlink(missing_ok=True) +``` + +### 6.5 目标版本:vLLM 直接写 TQ/Mooncake + +目标响应可改成: + +```json +{ + "kv_transfer_params": { + "backend": "transfer_queue", + "partition_id": "speco:run-1:dspark_train", + "sample_key": "abc123" + } +} +``` + +服务端顺序必须是: + +```text +prefill +→ 捕获指定层 hidden states +→ TQ/Mooncake put 完成 +→ 返回 HTTP success 和 sample key +``` + +这需要定制 vLLM exporter;当前 `verl-SpeCo-ls` 中没有服务端 writer 实现。 + +## 7. 按 PR #48 扩展 TQ bridge + +不要重新发明一套 transport。将 PR #48 的 bridge 设计移植到本项目并增加 standalone 所需方法: + +```python +class StandaloneTQTransport: + def put_sample(self, key, tensor_dict, tag): ... + def list_ready_keys(self, run_id): ... + def get_samples(self, keys, fields=None): ... + def clear_samples(self, keys): ... + def put_control(self, key, tag): ... + def close(self): ... +``` + +写入延续 PR #48 的真实形式: + +```python +tq.kv_put( + key=key, + partition_id=partition_id, + fields={ + "input_ids": feature.input_ids.cpu(), + "loss_mask": feature.loss_mask.cpu(), + "position_ids": feature.position_ids.cpu(), + "hidden_states": feature.hidden_states.cpu(), + }, + tag={ + "run_id": run_id, + "status": "ready", + "sequence_no": sequence_no, + "payload_bytes": payload_bytes, + }, +) +``` + +批量读取延续 PR #48 的 `kv_batch_get`: + +```python +result = tq.kv_batch_get( + keys=keys, + partition_id=partition_id, + fields=required_fields, # 0.1.7 是否支持该参数需实机确认 +) +``` + +新增发现和清理: + +```python +items = tq.kv_list(partition_id=partition_id) +tq.kv_clear(keys=keys, partition_id=partition_id) +``` + +这里的 `kv_list/kv_clear` 参数名必须根据锁定的 TQ 版本验证。当前参考仓库只实际调用了 `kv_put/kv_batch_get`,没有为这两个接口提供运行证据。 + +读取结果继续复用 PR #48 的两个适配函数: + +```python +value = _extract_value(result, key) +row = _tensordict_to_dict(value) +row["hidden_states"] = _densify_tq_tensor(row["hidden_states"]) +``` + +## 8. DSpark 多 rank 如何消费 + +### 8.1 第一版:rank 0 用 kv_list 发现 READY keys + +PR #48 中 key 由 Ray driver 传给 drafter worker;standalone 没有这条控制路径,所以 rank 0 需要主动列出 key: + +```python +rank = dist.get_rank() +world_size = dist.get_world_size() +global_batch_size = batch_size_per_gpu * world_size + +if rank == 0: + entries = tq_transport.list_ready_keys(run_id=run_id) + entries.sort(key=lambda x: (x.tag["sequence_no"], x.key)) + selected_keys = [x.key for x in entries[:global_batch_size]] +else: + selected_keys = None + +holder = [selected_keys] +dist.broadcast_object_list(holder, src=0) +selected_keys = holder[0] +``` + +`sequence_no` 是输入文件顺序。它让多个 Producer 并发完成顺序不同的情况下,Trainer 仍能确定性地组成 batch。 + +rank 0 广播的是字符串 key 列表,不是 hidden-state tensor。 + +### 8.2 各 rank 切自己的 keys + +例如 global batch keys: + +```text +[s0, s1, s2, s3, s4, s5, s6, s7] +``` + +world size 为 4、每卡 batch size 为 2: + +```text +rank 0 → [s0, s1] +rank 1 → [s2, s3] +rank 2 → [s4, s5] +rank 3 → [s6, s7] +``` + +代码: + +```python +def shard_keys(keys, rank, world_size): + assert len(keys) % world_size == 0 + per_rank = len(keys) // world_size + start = rank * per_rank + end = start + per_rank + return keys[start:end] +``` + +### 8.3 每个 rank 并行 get + +所有进程执行: + +```python +local_keys = shard_keys( + selected_keys, + rank=rank, + world_size=world_size, +) + +local_payloads = tq_transport.get_samples(local_keys) +``` + +数据路径: + +```text +rank 0 ← Mooncake(s0,s1) +rank 1 ← Mooncake(s2,s3) +rank 2 ← Mooncake(s4,s5) +rank 3 ← Mooncake(s6,s7) +``` + +不是 rank 0 get 全部后再 scatter。 + +### 8.4 转成当前训练格式 + +TQ 返回的数据需要先按 PR #48 的规则解包、densify,再转换成现有 `DraftFeatureSample`: + +```python +def tq_row_to_feature(row, tag): + return DraftFeatureSample( + algorithm="DSpark", + input_ids=row["input_ids"], + loss_mask=row["loss_mask"], + position_ids=row["position_ids"], + hidden_states=row["hidden_states"], + metadata={ + "hidden_states_layout": tag["hidden_states_layout"], + **row.get("metadata", {}), + }, + ) +``` + +`hidden_states_layout` 不能只放在无法恢复的临时 Python 对象里;应随 TQ tag 保存。rank 0 从 `kv_list` 取得 key/tag 后,需要把每个 local key 对应的 tag 一起交给本 rank 的转换逻辑。Consumer 必须把 layout 放回 `DraftFeatureSample.metadata`。DSpark L1 所需的 `target_last_hidden_states` 随后由现有 backend 从拼接的 `hidden_states` 中切出,transport 层不计算 loss,也不拆 layout。 + +## 9. 修改当前训练循环 + +在 `run_standalone_draft_training()` 中增加数据源分支: + +```python +feature_store_type = str(feature_store_cfg.get("type", "torch_shard")) + +if feature_store_type == "transfer_queue": + tq_stream = build_transfer_queue_stream( + config=config, + rank=rank, + world_size=world_size, + ) + store = None + loader = None + feature_replayer = None +else: + store = build_feature_store_from_config( + feature_store_cfg, + read_only=True, + ) + loader = DraftFeatureDataLoader(...) +``` + +流式训练循环: + +```python +while successful_steps < max_steps: + global_keys, materialized_samples = tq_stream.next_local_batch() + + batch = trainer.prepare_training_batch_from_samples( + materialized_samples, + step=optimizer_step, + ) + + has_batch = batch is not None + if not _all_ranks_true(has_batch, trainer.runtime_device): + raise RuntimeError("at least one rank failed to fetch its TQ batch") + + ok = await trainer.training_step_from_batch( + batch, + optimizer_step, + ) + + if not _all_ranks_true(ok, trainer.runtime_device): + raise RuntimeError("DSpark step failed on at least one rank") + + dist.barrier() + if rank == 0: + tq_stream.clear_global_batch(global_keys) + dist.barrier() +``` + +TQ 接在数据输入层,不放进 `DSparkTrainerBackend`。Backend 只负责模型、forward、loss、backward 和 optimizer。 + +## 10. READY key、inflight key 和训练提交 + +### 10.1 Ready + +在本方案中 ready 表示: + +> `dspark_train` 所需的所有 fields 已经成功写入 TQ storage,训练侧可以读取。 + +第一版不是通过 PR #48 尚未验证的 Sampler 自动判定,而是 Producer 完整校验 payload 后一次 `kv_put`,并设置 `tag.status=ready`。rank 0 的 `kv_list` 只选择这些 key。 + +### 10.2 Inflight key + +rank 0 选出一个 global batch 后,在本地保存: + +```python +inflight_global_keys = selected_keys +``` + +其他 step 不应再次选中这批 key。因为 standalone 只有一个同步 DSpark job,第一版由 rank 0 进程内 `set` 排除 inflight key 即可: + +```python +ready = [x for x in listed if x.key not in inflight_keys] +``` + +若以后允许多个独立 Trainer job 同时消费同一 partition,进程内 set 就不够,届时必须使用 TQ Controller/Sampler 的原子消费分配能力。 + +### 10.3 Optimizer committed + +optimizer committed 表示所有 DSpark rank 已经完成: + +```text +forward → backward → gradient synchronization → optimizer.step +``` + +它比 `kv_list` 发现 key、甚至 `kv_batch_get` 取出 tensor 都更晚。 + +第一版推荐简单语义: + +```text +TQ 负责 key/tag 和 tensor 传输 +rank 0 负责单 Trainer job 的 batch 选择和 inflight set +训练失败 → 整个作业 fail-fast +训练成功 → kv_clear payload,并从 inflight set 移除 +恢复 → 从最近 checkpoint + 输入 cursor 重新启动 +``` + +这样不需要重新实现 lease/ack 状态机,也没有夸大 PR #48 当前尚未使用的 Sampler 能力。 + +## 11. 为什么训练完一个 step 才清理 + +不能在 `kv_batch_get()` 后立即 clear: + +```text +get 成功 +→ clear +→ forward OOM +→ 数据已不存在,无法重试 +``` + +正确顺序: + +```text +rank 0..N get +→ 所有 rank 确认 batch 有效 +→ training_step_from_batch +→ _all_ranks_true(ok) +→ rank 0 kv_clear global keys +``` + +当前 `_all_ranks_true()` 已经是项目现有的跨 rank 同步工具,可以继续复用。 + +## 12. 背压 + +背压表示 Producer 生成速度高于 Trainer 消费速度时,Producer 必须暂停继续提交,以免 Mooncake 内存无限增长。 + +建议限制: + +```yaml +max_vllm_inflight_requests: 32 +max_pending_put_bytes: 8589934592 +max_tq_ready_samples: 256 +max_tq_ready_bytes: 68719476736 +``` + +Producer 在 tags 中写: + +```python +{"payload_bytes": payload_bytes} +``` + +周期性通过 TQ list/metadata 统计当前 partition 尚未消费的数据量: + +```python +while ready_bytes >= max_tq_ready_bytes: + await asyncio.sleep(backpressure_poll_interval) +``` + +如果锁定的 TQ 版本提供容量/ready 统计接口,应直接使用,避免全量扫描 keys。 + +## 13. Stable ID、幂等和孤儿数据 + +### 13.1 幂等 + +幂等表示同一操作重复执行,最终逻辑结果仍只有一份。 + +Producer 对同一样本重试时必须使用相同 `sample_key`: + +```python +await tq.put(key="abc123", ...) +await tq.put(key="abc123", ...) +``` + +不能每次生成随机 key: + +```text +abc123-retry-1 +abc123-retry-2 +``` + +否则一个输入可能训练多次并持续占用 Mooncake。 + +### 13.2 孤儿数据 + +孤儿数据表示 tensor 已经写进 storage,但由于 Producer 崩溃或 metadata 更新失败,没有进入正常消费路径。 + +使用 TQ 后,metadata 和 storage 由同一套系统管理,可以减少“Mooncake有对象、自研 Coordinator 没记录”的双系统窗口,但仍要配置 TTL/partition cleanup: + +```text +训练正常结束 → clear partition +训练异常退出 → 下次启动检查旧 partition +超过 TTL → 清理未消费数据 +``` + +## 14. EOS 和 drop-last + +EOS 表示 Producer 已经读完输入,并且所有 vLLM/TQ put 都已完成。 + +TQ 中需要一种结束条件,具体使用 Controller API、特殊 metadata 或单独的运行状态记录取决于锁定版本。不能把“当前暂时没有 ready sample”当成 EOS,因为 Producer 可能仍在请求 vLLM。 + +最后不足一个 global batch 时: + +```python +global_batch_size = batch_size_per_gpu * world_size +``` + +第一版使用 `drop_last=true`,避免某些 rank 有数据、某些 rank 没数据导致分布式训练不同步。 + +结束条件: + +```text +producer_done == true +and ready_samples < global_batch_size +and inflight_requests == 0 +and pending_puts == 0 +``` + +## 15. 双缓冲预取 + +训练 batch N 时,CPU 后台线程预取 batch N+1: + +```python +next_future = executor.submit(tq_stream.next_local_batch) + +current_batch = first_batch +while current_batch is not None: + next_batch = next_future.result() + next_future = executor.submit(tq_stream.next_local_batch) + + train(current_batch) + current_batch = next_batch +``` + +实际顺序应调整为避免等待 future 后才训练。推荐: + +```python +current = tq_stream.next_local_batch() + +while current is not None: + future = executor.submit(tq_stream.next_local_batch) + train_and_clear(current) + current = future.result() +``` + +第一版只预取一个 global batch,避免 Trainer 崩溃时大量数据已被采样但未训练。 + +如果 Mooncake/TQ Python get 是阻塞函数,使用专用 `ThreadPoolExecutor`,不要阻塞 Producer 的 asyncio event loop。 + +## 16. 建议代码结构 + +```text +verl_speco/ + trainer/ + tq_transport.py # TQ client、put/get/meta/clear 封装 + tq_feature_stream.py # kv_list、global keys、rank shard、decode/prefetch + feature_producer.py # JSONL → 并发 vLLM → TQ + draft_training_loop.py # 增加 transfer_queue 数据源分支 + target_feature_replay.py # 复用 vLLM payload 校验/转换逻辑 +``` + +不要新增: + +```text +coordinator.py +coordinator_client.py +``` + +建议抽象: + +```python +class StreamingFeatureSource(Protocol): + def next_local_batch(self) -> tuple[list[str], list[DraftFeatureSample]]: ... + def clear_global_batch(self, keys: list[str]) -> None: ... + def close(self) -> None: ... +``` + +这样训练循环不依赖 TQ 的具体类型。 + +## 17. 配置草案 + +```yaml +actor_rollout_ref: + rollout: + drafter: + training: + mode: offline + backend: dspark + batch_size_per_gpu: 2 + max_steps: 1000 + + feature_store: + type: transfer_queue + partition_id: speco_drafter_features_${run_id} + drop_last: true + prefetch_steps: 1 + + transfer_queue: + # 与 PR #48 的配置层级和 init 方式保持一致。 + enable: true + package_version: 0.1.8 # 最终以实测版本为准 + backend: + storage_backend: MooncakeStore + MooncakeStore: + auto_init: false + metadata_server: localhost:50123 + master_server_address: localhost:50124 + local_hostname: localhost + protocol: tcp + global_segment_size: 4294967296 + local_buffer_size: 1073741824 + device_name: "" + + required_fields: + - input_ids + - loss_mask + - position_ids + - hidden_states + + producer: + input_path: /path/to/generated_responses.jsonl + vllm_endpoints: + - http://node0:8000/v1 + - http://node1:8000/v1 + max_inflight_requests: 32 + max_pending_put_bytes: 8589934592 + max_ready_samples: 256 + max_ready_bytes: 68719476736 +``` + +当前 examples 中的: + +```bash +transfer_queue.enable=False +``` + +属于 verl RL 主入口配置,当前 standalone `draft_train_launcher` 不读取它。不能只改成 `True`;必须实现上述 `feature_store.type=transfer_queue` 分支。 + +## 18. 启动顺序 + +逻辑顺序: + +```text +1. 启动 Mooncake metadata/master 服务; +2. 启动 standalone TQ owner 进程,调用一次 `tq.init(tq_config)` 创建/连接 TQ Controller 和 MooncakeStore backend; +3. 启动一个或多个定制 vLLM server +4. 启动 Feature Producer +5. Producer 和所有 Trainer rank 调用无参 `tq.init()`,连接 owner 创建的同一套 TQ; +6. 启动 verl_speco.draft_train_launcher +7. torchrun 启动所有 DSpark rank +8. 各 rank 连接 TQ +9. rank 0 用 `kv_list` 选择 global keys,各 rank 并行 `kv_batch_get`/train; +10. 输入耗尽后 Producer 发布 done 状态 +11. Trainer drain 完整 global batches 后退出 +12. 清理 partition,停止 Producer、vLLM、TQ、Mooncake +``` + +PR #48 的 `init_transfer_queue()` 在 Ray `SpecoTaskRunner` 中调用 `tq.init(config)`,worker 的无参 `tq.init()`依靠同一个 Ray 集群发现 named Controller;这部分不能原样复制到 standalone torchrun。 + +本项目要求一开始不使用 Ray,因此 Phase 0 必须先证明锁定的 TQ 版本支持独立 owner/controller 进程,以及 Producer/torchrun rank 如何获得连接信息。若 TQ 0.1.7/0.1.8 实际只能通过 Ray named actor 完成发现,那么有两个选择: + +1. 接受仅用 Ray 承载 TQ 控制面的最小方案; +2. 给 TQ 增加或使用其已有的独立 ZMQ/server-info bootstrap。 + +在这项验证完成前,文档不能声称 PR #48 已经提供“无 Ray TQ standalone 启动”。 + +## 19. 故障处理 + +### vLLM 请求失败 + +- 对单个 sample 按稳定 key 重试; +- 指数退避; +- 超过次数记录失败,并根据配置 fail-fast 或跳过; +- 不写不完整 TQ fields。 + +### vLLM 文件读取成功,但 TQ put 失败 + +- 暂时保留临时文件; +- 重试 TQ put; +- put 成功后再删除; +- 不把样本视为 ready。 + +### 某个训练 rank get 失败 + +- 该 rank 报告 `local_ok=false`; +- `_all_ranks_true()` 使全部 rank 得到一致失败结果; +- 第一版整个训练 fail-fast; +- 不 clear global batch。 + +### OOM/optimizer step 失败 + +- 不 clear; +- 所有 rank 一致退出; +- 从最近训练 checkpoint 恢复; +- 根据 TQ 消费提交语义决定是否重放当前 batch。 + +### clear 失败 + +- optimizer 已成功,不能再次训练这批; +- 将 batch keys 写入本地小型 `gc_pending` 日志; +- 后台重试 clear; +- checkpoint 保存最近 committed sample IDs,避免恢复时重复消费。 + +## 20. 观测指标 + +Producer: + +```text +producer/vllm_inflight +producer/vllm_requests_per_sec +producer/vllm_prefill_tokens_per_sec +producer/vllm_p50_latency +producer/vllm_p95_latency +producer/tq_put_bytes_per_sec +producer/tq_put_failures +producer/pending_put_bytes +``` + +TQ/Mooncake: + +```text +tq/ready_samples +tq/ready_bytes +tq/consumed_samples +tq/storage_bytes +tq/clear_failures +mooncake/put_bandwidth +mooncake/get_bandwidth +``` + +Trainer: + +```text +trainer/tq_wait_seconds +trainer/tq_get_seconds +trainer/tq_get_bytes_per_sec +trainer/decode_seconds +trainer/h2d_seconds +trainer/step_seconds +trainer/data_stall_ratio +trainer/successful_steps +``` + +## 21. 实施阶段 + +### Phase 0:锁定依赖和契约 + +- 从 PR #48 的 `TransferQueue==0.1.7` 起验证,同时对比当前 verl 文档使用的 0.1.8; +- 实测 `tq.init(config)`、worker `tq.init()`、`kv_put`、`kv_batch_get`、`kv_list`、`kv_clear`; +- 验证该版本是否支持无 Ray owner/controller 部署以及连接信息传递; +- 先用 PR #48 已配置的 `SimpleStorage` 做最小闭环; +- 再把 backend 换成 `MooncakeStore`,验证 tcp,最后再验证 rdma; +- 固定 fields、tags、partition 和 `run_id/sequence_no/status`; +- 写 fake TQ 单元测试。 + +### Phase 1:文件桥接 + TQ KV 模式 + +- 新增独立 Producer; +- 32 个有界并发 vLLM 请求; +- 读取 vLLM 临时 safetensors; +- TQ put 成功后删除文件; +- standalone trainer 由 rank 0 `kv_list` 并广播 global keys; +- 各 rank 并行 `kv_batch_get`; +- 复用 PR #48 的返回值解包、NestedTensor densify 和 per-step cache; +- optimizer 成功后 `kv_clear`。 + +验收:连续训练 1000 step,临时文件数量、TQ ready bytes 和 Mooncake占用均保持有界。 + +### Phase 2:双缓冲与多 endpoint + +- 增加多 endpoint 最少 inflight 调度; +- 增加一个 global batch 预取; +- 动态背压; +- 注入单 rank get 失败,确认所有 rank 一致退出而非死锁。 + +### Phase 3:vLLM 直接写 TQ/Mooncake + +- 修改外部定制 vLLM exporter; +- 去掉 `hidden_states_path` 临时文件; +- HTTP 响应返回 partition/sample key; +- 验证 HTTP 重试的幂等性。 + +### Phase 4:可选升级到 TQ StreamingDataLoader + +- 在当前保守方案稳定后再引入 RankAwareSampler; +- 让每个 rank 自动取得 local micro-batch; +- 去掉 rank 0 手工 key-list 广播; +- 验证与 torchrun/DSpark 的 global step 对齐。 + +## 22. 最终推荐 + +针对当前 `verl-SpeCo-ls`,推荐的第一版不是 AngelSpec 的完整架构,也不是单独写 Coordinator,而是: + +```text +当前预生成 response 文件 +→ 独立 asyncio Producer +→ 并行访问多个 vLLM endpoint +→ 读取并校验临时 hidden-state 文件 +→ TransferQueue put +→ Mooncake storage backend +→ rank 0 kv_list 获取 READY global keys +→ broadcast key list +→ 各 DSpark rank 并行 kv_batch_get +→ 现有 prepare_training_batch_from_samples() +→ 现有 training_step_from_batch() +→ 全 rank 成功 +→ TQ clear +``` + +这套方案保留当前 standalone DSpark 训练主体,只替换 `DraftFeatureDataLoader + TargetFeatureReplayer.materialize()` 所在的数据输入路径。它真正建立在 PR #48 已实现的 KV transport 之上,而不是假设 PR #48 已经实现了 standalone Sampler/StreamingDataLoader。 + +## 23. 参考 + +- verl TransferQueue: +- TransferQueue: +- Mooncake Store: +- 本地只读参考:`../verl-SpeCo/verl_speco/integration/transferqueue_bridge.py`(PR #48) +- 本地只读参考:`../verl-SpeCo/docs/transferqueue_integration_plan.md`(PR #48) diff --git a/docs/draft_feature_sample_tq_protocol_refactor_plan.md b/docs/draft_feature_sample_tq_protocol_refactor_plan.md new file mode 100644 index 00000000..10398c85 --- /dev/null +++ b/docs/draft_feature_sample_tq_protocol_refactor_plan.md @@ -0,0 +1,549 @@ +# TQ `DraftFeatureSample` 通用传输协议重构方案 + +> Last updated: 08/27/2026 + +## 1. 目标与结论 + +当前 standalone TQ 流程已经在 Consumer 侧恢复为 `DraftFeatureSample`,然后调用既有的 +`DrafterTrainer.prepare_training_batch_from_samples()`。但是传输层仍然通过一套较重的 +`SampleMetadata` 重新描述 hidden-state 布局、shape 和训练字段,导致: + +- `DraftFeatureSample.metadata` 不能完整往返; +- TQ codec 了解过多 DSpark/hidden-state 布局细节; +- 新算法即使已经能构造 `DraftFeatureSample`,仍可能需要修改 TQ 协议; +- Producer 和 Consumer 分别实现 ready tag 过滤,容易产生统计口径不一致。 + +本次重构采用以下边界: + +1. 保留 `SampleMetadata`,但将它缩减为 **TQ 控制信封**; +2. 完整训练数据只由 `DraftFeatureSample` 表达; +3. TQ codec 对 `DraftFeatureSample` 做通用、无损、算法无关的编码和解码; +4. Consumer 解码后直接把 `DraftFeatureSample` 交给现有训练流程; +5. 算法差异只保留在 Producer 的样本构造和 Trainer backend 中; +6. Producer 和 Consumer 复用同一个 ready-tag 解析函数。 +7. Producer 写入资格与 Consumer 读取资格使用同一个共享判定,禁止两端分别实现近似校验; +8. hidden-state token/position/row 对齐失败的样本在 Producer 侧直接丢弃,不得部分截取后写入 TQ。 + +TQ 不能直接存放 Python dataclass 实例。TQ 0.1.7 的数据面接收 tensor fields,因此仍然 +需要 `encode_sample()` / `decode_sample()`;这里要删除的是自定义的训练数据定义,而不是 +必要的传输编码。 + +## 2. 重构后的职责划分 + +### 2.1 `SampleMetadata`:只负责队列控制 + +建议保留以下字段: + +```python +@dataclass(frozen=True) +class SampleMetadata: + protocol_schema_version: int + run_id: str + sample_id: str + sequence_no: int +``` + +字段含义: + +| 字段 | 用途 | +| --- | --- | +| `protocol_schema_version` | TQ key/tag/fields 编码格式的版本,不是模型算法版本 | +| `run_id` | 隔离不同 standalone 训练任务 | +| `sample_id` | 保留输入样本身份,便于定位错误 | +| `sequence_no` | 为并发完成的样本建立确定顺序,并生成唯一 key | + +从 `SampleMetadata` 删除以下字段: + +```text +algorithm +target_model_id +target_model_revision +tokenizer_fingerprint +target_layer_ids +hidden_states_layout +hidden_dtype +hidden_shape +feature_length +full_sequence_length +feature_start +feature_end +use_logits +``` + +这些字段如果训练需要,应保存在 `DraftFeatureSample.algorithm` 或 +`DraftFeatureSample.metadata` 中。TQ 控制层不再验证其算法语义。 + +### 2.2 TQ tag:控制面索引 + +tag 保留可在不加载 tensor payload 的情况下完成发现、排序和 run 隔离所需的信息: + +```python +tag = { + "record_type": "sample", + "status": "ready", + "protocol_schema_version": 2, + "run_id": "dspark-a1b2c3", + "sequence_no": 6500, + "sample_id": "train-006500", +} +``` + +tag 不再携带 `algorithm`、layer IDs、hidden shape 等训练信息。EOS 仍使用独立 control tag: + +```python +tag = { + "record_type": "control", + "status": "eos", + "protocol_schema_version": 2, + "run_id": "dspark-a1b2c3", + "total_samples": 15000, +} +``` + +### 2.3 TQ fields:完整 `DraftFeatureSample` + +fields 使用原生 tensor 字段加一个 JSON manifest: + +```python +fields = { + "sample__input_ids": Tensor, + "sample__loss_mask": Tensor, + "sample__hidden_states": Tensor, + "sample__position_ids": Tensor, # 可选 + "sample__last_hidden_states": Tensor, # 可选 + "sample__target": Tensor, # 可选 + "sample__target_logprobs": Tensor, # 可选 + "sample__metadata_tensor__000000": Tensor, + "sample__manifest_json": UInt8Tensor, +} +``` + +`sample__manifest_json` 的逻辑内容示例: + +```json +{ + "draft_feature_schema_version": 1, + "algorithm": "DSPARK", + "present_fields": [ + "input_ids", + "loss_mask", + "hidden_states", + "position_ids" + ], + "hidden_states_kind": "tensor", + "metadata": { + "hidden_states_layout": "dflash_aux_plus_last", + "target_layer_ids": [1, 12, 23, 34, 45], + "feature_start": 31, + "feature_end": 543, + "hidden_positions": { + "__tq_tensor_ref__": "sample__metadata_tensor__000000" + } + } +} +``` + +manifest 和 tensor fields 合起来必须能够完整恢复: + +```python +DraftFeatureSample.from_dict(payload, strict=True) +``` + +## 3. 通用 metadata codec + +`DraftFeatureSample.metadata` 不能简单 `json.dumps()`,因为当前代码会在其中保存 +`hidden_positions` 等 tensor。新 codec 采用递归 tree 编码。 + +直接写入 JSON 的类型: + +```text +None、bool、int、float、str +dict[str, value] +list[value] +tuple[value](manifest 记录 tuple 类型,解码后恢复 tuple) +``` + +tensor 的处理方式: + +```text +metadata中的Tensor +→ 转为CPU contiguous Tensor +→ 单独写入fields +→ manifest原位置写tensor field引用 +``` + +例如: + +```python +metadata = { + "feature_start": 31, + "hidden_positions": torch.tensor([31, 32, 33]), +} +``` + +编码为: + +```python +fields["sample__metadata_tensor__000000"] = tensor([31, 32, 33]) + +manifest["metadata"] = { + "feature_start": 31, + "hidden_positions": { + "__tq_tensor_ref__": "sample__metadata_tensor__000000" + }, +} +``` + +不支持的对象不能静默执行 `str(value)`,否则协议不是无损的。第一版应 fail closed,错误中打印 +metadata 路径和实际类型。后续如果确实存在 NumPy scalar/array,可显式增加稳定编码规则。 + +## 4. `hidden_states` 两种表示 + +`DraftFeatureSample.hidden_states` 支持: + +```python +torch.Tensor | list[torch.Tensor] +``` + +单 tensor: + +```python +fields["sample__hidden_states"] = hidden +manifest["hidden_states_kind"] = "tensor" +``` + +tensor list: + +```python +fields["sample__hidden_states__000000"] = hidden_0 +fields["sample__hidden_states__000001"] = hidden_1 +manifest["hidden_states_kind"] = "list" +manifest["hidden_states_fields"] = [ + "sample__hidden_states__000000", + "sample__hidden_states__000001", +] +``` + +这样协议不会再因某个算法使用 tensor list 而报错。 + +## 5. 需要修改的文件和函数 + +### 5.1 `verl_speco/transport/drafter_sample_protocol.py` + +这是主要重构文件。 + +修改内容: + +1. 将 `SampleMetadata` 缩减为控制信封; +2. 将 `PROTOCOL_SCHEMA_VERSION` 从 1 升到 2; +3. 修改 `make_sample_key()`,继续使用 protocol version、run、sequence 和 sample ID; +4. 修改 `make_ready_tag()`,只生成控制面字段; +5. 重写 `encode_sample(sample, meta)`: + - 调用 `DraftFeatureSample.to_dict()`; + - 编码所有 dataclass tensor 字段; + - 编码 hidden-state tensor list; + - 递归编码 metadata; + - 生成 manifest; +6. 重写 `decode_sample(key, tag, fields, expected_config)`: + - 解析并校验控制信封; + - 解析 manifest; + - 恢复所有 tensor 和 metadata; + - 调用 `DraftFeatureSample.from_dict(..., strict=True)`; +7. 将 `_validate_primary_tensors()` 中与具体 hidden layout/shape 的约束删除; +8. 新增并导出统一函数: + +```python +parse_ready_tag(tag) -> SampleMetadata | None +is_ready_sample_tag(tag, *, run_id, protocol_schema_version) -> bool +``` + +Producer backpressure 和 Consumer discovery 必须复用这两个函数,禁止再分别复制过滤条件。 + +此外增加统一的发布资格函数: + +```python +validate_publishable_sample(sample) -> None +``` + +`encode_sample()` 和 Consumer 的 `decode_sample()` 都调用同一组 sample 结构校验。Producer 只有 +通过该校验后才能生成 ready tag;这样不存在“Producer 写入成功,但 Consumer 按另一套规则过滤”的 +中间状态。协议错误必须在 `put_sample()` 前暴露。 + +建议新增内部函数: + +```python +_encode_metadata_tree(value, fields, path) -> JSONValue +_decode_metadata_tree(value, fields, path) -> Any +_encode_hidden_states(value, fields, manifest) -> None +_decode_hidden_states(fields, manifest) -> Tensor | list[Tensor] +_json_to_uint8_tensor(value) -> Tensor +_uint8_tensor_to_json(value) -> Any +``` + +### 5.2 `verl_speco/standalone_tq_producer.py` + +修改内容: + +1. 保留 `PreparedFeature.metadata: SampleMetadata`,但它现在只是 TQ 信封; +2. 简化 `_sample_metadata()`,只读取: + +```text +run_id +request.sample_id +request.sequence_no +protocol_schema_version +``` + +3. 删除 `_sample_metadata()` 对 feature shape、layout、target model 和 logits 的复制; +4. `publish_one()` 仍保持: + +```python +fields = encode_sample(result.sample, result.metadata) +tag = make_ready_tag(result.metadata) +transport.put_sample(key, fields, tag=tag) +``` + +5. `_wait_for_pending_capacity()` 使用协议模块的 + `is_ready_sample_tag()`,与 Consumer 使用完全相同的过滤规则; +6. 保持“只有 TQ put 成功后才删除 vLLM 临时文件”的生命周期不变。 + +Producer 还必须对 hidden-state 对齐失败做样本级丢弃: + +```text +token_ids 与请求 token IDs 不一致 +feature positions 超出 hidden-state rows +hidden-state rows 不能覆盖完整训练窗口 +hidden-state layer 数不足 +→ 记录 sample_id/sequence_no/原因 +→ 删除本次 vLLM 临时文件和 lock +→ dropped_count += 1 +→ 不进入 publish_queue +→ 不写 ready tag/fields +→ request worker 继续处理下一条样本 +``` + +不能沿用当前“只丢弃越界 positions、使用剩余 positions 继续训练”的行为。TQ Producer 应启用严格 +对齐模式:只要一个目标位置无法与 hidden-state row 对应,整条样本就无效。普通网络错误、TQ put +错误和协议编程错误仍然 fail fast,不能被误当作脏样本吞掉。 + +Producer 的算法相关职责仍然保留在: + +```python +feature_from_vllm_payload(raw, request, feature_contract) +``` + +也就是说,Producer 必须先构造正确且完整的 `DraftFeatureSample`,TQ codec 不负责推断算法布局。 + +### 5.3 `verl_speco/trainer/tq_feature_store.py` + +修改内容: + +1. `list_ready()` 使用统一 `parse_ready_tag()`; +2. 删除本文件中重复的 tag 字段校验; +3. `get_many()` 继续调用 `decode_sample()`,返回类型仍为 + `list[DraftFeatureSample]`; +4. `ExpectedFeatureConfig` 只检查: + +```text +run_id +protocol_schema_version +``` + +5. 不再在 TQ store 中检查 algorithm、target model、layer IDs、dtype 和 layout; +6. EOS 解析使用同一个 protocol version 字段命名。 + +### 5.4 `verl_speco/trainer/tq_sample_source.py` + +主流程不需要改变: + +```text +rank 0 list_ready +→ 按sequence_no排序 +→ 为各rank分配key +→ 各rank get_many +→ 得到DraftFeatureSample +``` + +只需确保诊断日志中的 ready 统计也调用统一 tag parser,避免日志口径和正式读取口径不同。 + +### 5.5 `verl_speco/trainer/feature_store.py` + +第一版不修改 `DraftFeatureSample` 公共字段,避免影响现有离线 store、PR #48 和非 TQ 路径。 + +可以新增一个小的公共字段列表,供 store 和 TQ codec 复用,例如: + +```python +DRAFT_FEATURE_OPTIONAL_TENSOR_FIELDS = ( + "last_hidden_states", + "target", + "target_logprobs", + "position_ids", +) +``` + +不要让 TQ codec 再维护一份不同的 optional-field 列表。 + +### 5.6 `verl_speco/trainer/target_feature_replay.py` + +不修改训练转换逻辑。`feature_from_vllm_payload()` 继续负责把不同算法的 vLLM 输出转换为 +`DraftFeatureSample`。 + +当前它支持: + +```text +EAGLE3、DFLASH、DSPARK +``` + +以后增加新算法时,只需要在这里或对应算法 converter 中实现: + +```text +RawVllmFeature + TokenizedRequest → DraftFeatureSample +``` + +如果新的 `DraftFeatureSample` 字段都能被通用 codec 表达,则无需再次修改 TQ 传输层。 + +### 5.7 `verl_speco/trainer/base_trainer.py` + +不需要修改。Consumer 解码结果继续走: + +```python +trainer.prepare_training_batch_from_samples(samples, step=optimizer_step) +``` + +其中每个元素已经是完整 `DraftFeatureSample`,随后调用现有: + +```python +sample.to_training_item() +``` + +算法 backend 仍由启动配置的 `speculative_algorithm` 选择。 + +## 6. 重构后的端到端流程 + +### Producer + +1. 读取 prompt/response; +2. 请求 vLLM generate/prefill; +3. `feature_from_vllm_payload()` 按算法生成完整 `DraftFeatureSample`; +4. 生成简化 `SampleMetadata` 控制信封; +5. `encode_sample()` 无损编码 sample; +6. 将 tensor fields 和 ready tag 写入 TQ; +7. TQ put 成功后删除 vLLM safetensors 临时文件; +8. 全部样本完成后写 EOS。 + +### Consumer + +1. rank 0 使用统一 tag parser 发现当前 run 的 ready keys; +2. 按 `sequence_no` 排序并切出一个 global batch; +3. 广播各 rank 的 key/tag 分配; +4. 各 rank 从 TQ 获取自己的 tensor fields; +5. `decode_sample()` 无损恢复 `DraftFeatureSample`; +6. 调用 `prepare_training_batch_from_samples()`; +7. 由已选择的算法 backend 完成训练; +8. 所有 rank 训练成功后,rank 0 清理这一 global batch 的 keys。 + +## 7. 测试修改 + +### 7.1 `tests/unit/test_drafter_sample_protocol.py` + +替换以 DSpark shape 为中心的测试,增加完整 round-trip: + +1. 单 tensor hidden states; +2. hidden-state tensor list; +3. 所有 optional tensor 字段; +4. metadata 嵌套 dict/list/tuple; +5. metadata 中包含 tensor; +6. EAGLE3/DFLASH/DSPARK/DOMINO 的 algorithm 字符串均能往返; +7. 不支持的 metadata 对象明确报错并包含字段路径; +8. key/tag run、sequence、sample ID 不匹配时 fail closed; +9. protocol version 不匹配时 fail closed。 + +核心断言不是只比较 shape,而是逐字段比较原始 sample 与恢复 sample。 + +### 7.2 `tests/unit/test_tq_producer.py` + +增加: + +1. Producer 写入 fields 后能完整恢复其原始 `DraftFeatureSample.metadata`; +2. pending capacity 使用统一 ready parser; +3. 其他 run、错误 protocol version 和非 sample control tag 不计入 pending; +4. TQ put 失败时仍不删除 vLLM 临时文件。 +5. token IDs 不匹配时删除临时文件、增加 dropped 计数且不调用 TQ put; +6. 任一 feature position 越界时整条样本丢弃,不生成部分长度 sample; +7. 丢弃无效样本后 worker 能继续发布后续有效样本并最终写 EOS; +8. EOS 的 `total_samples` 使用实际成功发布数,不包含 dropped 样本。 + +### 7.3 `tests/unit/test_tq_consumer.py` + +增加: + +1. TQ store 恢复完整 `DraftFeatureSample`; +2. metadata tensor 完整恢复; +3. hidden-state list 完整恢复; +4. ready discovery 与 Producer pending 统计对同一组 tags 给出相同结果; +5. 多算法 sample 都能进入 `prepare_training_batch_from_samples()` 的现有入口。 + +### 7.4 回归测试 + +必须继续运行: + +```bash +pytest -q tests/unit/test_drafter_sample_protocol.py +pytest -q tests/unit/test_tq_producer.py +pytest -q tests/unit/test_tq_consumer.py +pytest -q tests/unit/test_target_feature_replay.py +pytest -q tests/unit/test_draft_feature_store.py +``` + +然后运行一组 Producer/TQ/Consumer smoke test,至少确认: + +```text +put → list → get → decode → train one step → clear → EOS +``` + +## 8. 版本与兼容策略 + +建议将协议版本提升到 2,不实现 v1/v2 混读。 + +理由: + +- standalone TQ 是在线临时队列,不是长期离线数据集; +- Producer 和 Consumer 本来就应作为同一版本部署; +- 同一 run 中混用两种 fields 格式会增加错误恢复复杂度; +- fail closed 比错误地恢复训练数据更安全。 + +启动时 Producer、Owner 和 Consumer 必须使用同一个 protocol version。旧 run 的 TQ 数据不能被新 +Consumer 接续;重启完整 pipeline 时使用新的 `run_id`。 + +## 9. 实现顺序 + +建议按以下顺序实施,每一步都能单独测试: + +1. 简化 `SampleMetadata`,确定 v2 tag/key/EOS 格式; +2. 实现 metadata tree codec; +3. 实现完整 `DraftFeatureSample` encode/decode; +4. 完成 protocol round-trip 单测; +5. 修改 Producer 的控制信封构造和统一 ready 统计; +6. 为 TQ Producer 增加严格 hidden-state 对齐和样本级丢弃; +7. 修改 Consumer store 的统一 tag 解析; +8. 修改 Producer/Consumer 单测; +9. 运行单进程 TQ round-trip smoke; +10. 运行多 rank Consumer 一步训练; +11. 分别用 EAGLE3、DFLASH、DSPARK 的合成 sample 验证通用传输; +12. 最后再为尚未支持的算法增加 vLLM-output-to-sample converter。 + +## 10. 完成标准 + +满足以下条件才算重构完成: + +- TQ codec 中不存在 `if algorithm == "DSPARK"` 一类分支; +- `decode_sample(encode_sample(sample))` 能无损恢复所有公共字段和 metadata; +- Producer 和 Consumer 使用同一个 ready tag parser; +- Consumer 解码后直接得到 `DraftFeatureSample`; +- `base_trainer.py` 的后续训练入口无需为 TQ 添加算法分支; +- 新算法只要能构造现有 `DraftFeatureSample`,传输层无需修改; +- ready backpressure 与 Consumer discovery 对同一组 tags 的计数完全一致; +- TQ put 成功前不删除 vLLM 临时文件,训练成功前不清理 TQ sample。 +- Producer 与 Consumer 对 sample fields 调用同一组结构校验; +- hidden-state 对齐不完整的样本不会以缩短后的 feature window 进入 TQ; +- dropped 样本有明确计数和限频日志,且不会阻止后续有效样本和 EOS。 diff --git a/docs/standalone_tq_consumer_implementation.md b/docs/standalone_tq_consumer_implementation.md new file mode 100644 index 00000000..75eec9a8 --- /dev/null +++ b/docs/standalone_tq_consumer_implementation.md @@ -0,0 +1,713 @@ +# 独立 DSpark 训练 TQ Consumer 实现说明 + +Last updated: 08/21/2026 + +## 1. 文档范围和当前结论 + +本文只说明当前仓库中已经实现的独立训练 Consumer。这里的 Consumer 是由 `torchrun` 启动的 DSpark 草稿模型训练任务:它持续从 TransferQueue(下文简称 TQ)发现样本,各训练 rank 分别取得自己负责的 Tensor,复用原有 DSpark 训练逻辑完成一次 optimizer step,然后由 rank 0 删除这一整个 global batch 对应的 TQ 记录。 + +当前已完成的能力是: + +1. `feature_store.type=tq` 可以作为独立训练的数据源,不要求磁盘 `path`。 +2. 每个训练 rank 都连接同一个 Ray 集群、同一个 TQ Controller 和同一个 partition。 +3. 只有 rank 0 调用 `kv_list` 发现 ready key,并把 key/tag 分配给各 rank。 +4. key 和 tag 通过 `torch.distributed.broadcast_object_list` 传输;hidden states 等 Tensor 不经过该广播。 +5. 每个 rank 根据分配到的 key,直接调用 TQ `kv_batch_get` 获取自己的 Tensor。 +6. TQ Tensor 被解码成原训练代码已经认识的 `DraftFeatureSample`,然后复用 `DrafterBaseTrainer` 的 batch 构造、DSpark loss、反向传播和 optimizer step。 +7. 只有当所有 rank 都成功完成该 step 后,rank 0 才调用 `kv_clear` 删除整个 global batch。 +8. Producer 发布 EOS 后,如果剩余样本不足一个 global batch,当前第一版会丢弃并清理这部分尾样本,然后正常结束训练迭代。 + +本文不会把尚未实现的 Producer 写成现有能力。Producer 后续需要复用本文第 6 节所述的公共协议,调用 `encode_sample()` 生成 fields,再使用 bridge 写入相同 TQ。 + +## 2. 本次涉及的文件 + +### 2.1 本次新增的 Consumer 核心文件 + +| 文件 | 实现的组件 | 作用 | +|---|---|---| +| `verl_speco/trainer/tq_feature_store.py` | `TQFeatureStore`、`ReadyEntry`、`EosMetadata` | 将公共 TQ bridge 包装成 Consumer 数据访问层,负责连接、发现、批量读取、解码、删除和读取 EOS | +| `verl_speco/trainer/tq_sample_source.py` | `TQFeatureDataLoader`、`TQLocalBatch`、`build_assignments()` | 实现多 rank 流式取数:rank 0 发现样本并分配 key,各 rank 自己从 TQ 取 Tensor | +| `tests/unit/test_tq_consumer.py` | Consumer 单元测试 | 覆盖 store 构建、ready 过滤排序、协议解码、rank 分配、EOS、清理和非 rank 0 行为 | + +### 2.2 本次修改的既有文件 + +| 文件 | 修改内容 | 为什么要改 | +|---|---|---| +| `verl_speco/trainer/feature_store.py` | factory 新增 `type=tq` 分支 | 让既有独立训练入口能够像选择磁盘 feature store 一样选择流式 TQ 数据源 | +| `verl_speco/trainer/draft_training_loop.py` | 接入 TQ store/loader、跨 rank 连接检查、训练成功后清理 | 将流式取数接入原训练循环,同时保留原 DSpark trainer、loss、optimizer、metric 和 checkpoint 逻辑 | +| `verl_speco/draft_train_launcher.py` | 增加 TQ 启动参数的 fail-fast 检查 | 在启动多个 torchrun 子进程前检查 `enable`、Ray address 和 `run_id`,避免各 rank 启动后才失败 | +| `verl_speco/config/speco_base.yaml` | 标注 `feature_store.type=tq` 为无路径流式数据源 | 保留统一 Hydra 配置入口;TQ 的公共配置仍位于 sibling `training.transfer_queue` | +| `tests/unit/test_draft_train_launcher.py` | 增加 TQ 参数检查测试 | 验证必要配置缺失时 launcher 直接拒绝启动 | +| `tests/unit/test_draft_training_loop.py` | 增加连接和 clear 时序测试 | 验证只由 rank 0 清理、clear 失败会报告、连接失败会传播 | + +### 2.3 直接复用的公共基础 + +以下文件不是这次 Consumer 才创造的概念,但 Consumer 直接使用它们: + +| 文件 | 被复用的能力 | +|---|---| +| `verl_speco/integration/transferqueue_bridge.py` | 屏蔽 TQ 0.1.7 API 细节,提供连接、`kv_list`、`kv_batch_get`、`kv_clear` 和本地关闭接口 | +| `verl_speco/transport/drafter_sample_protocol.py` | 定义 key、tag、fields、metadata 格式,以及 `encode_sample()` / `decode_sample()` | +| `verl_speco/trainer/feature_store.py` | 复用 `DraftFeatureSample`,使 TQ 数据进入训练侧后与磁盘 feature sample 类型一致 | +| `verl_speco/trainer/base_trainer.py` 及既有 backend | 复用 `DrafterBaseTrainer.prepare_training_batch_from_samples()` 和 `training_step_from_batch()` 等训练实现 | + +## 3. 运行时角色 + +### 3.1 TQ Owner + +TQ Owner 是单独的普通 Python 进程。它连接指定 Ray 集群,并以带配置的 `tq.init(config)` 创建任务级 named Controller 和 storage actors。Owner 持有全局 TQ 生命周期;Consumer 结束时不能关闭它。 + +Owner 不是训练 rank,也不执行 DSpark 模型。它的主要作用是让 Producer 和 Consumer 能通过同一个 Ray actor registry 找到同一个 TQ Controller。 + +### 3.2 Producer + +Producer 是后续需要实现的独立推理进程。它应并行调用 vLLM hidden-state 接口,构造一条条 `DraftFeatureSample` 和 `SampleMetadata`,再写入 TQ。 + +Producer 与 Consumer 不通过 Ray RPC 互相调用,也不通过 HTTP 直接传 Tensor。二者只需满足: + +- 连接同一个 Ray address; +- 使用同一个 Ray namespace; +- 使用同一个 TQ partition; +- 使用同一个 `run_id` 和协议版本。 + +### 3.3 Consumer launcher + +`python -m verl_speco.draft_train_launcher` 是父进程。它检查命令行 override,构造 `python -m torch.distributed.run ...` 命令,然后启动训练子进程。 + +launcher 自己不连接 TQ、不取样本、也不持有 GPU 模型。 + +### 3.4 Consumer training rank + +`torchrun --nproc_per_node=N` 会启动 N 个训练 OS 进程。每个进程有独立的: + +- global rank; +- local rank; +- GPU; +- `DrafterBaseTrainer`; +- `TQFeatureStore` 和本地 TQ client; +- DSpark 模型分片及 optimizer 状态。 + +这些 rank 共同执行一个分布式草稿模型训练任务。rank 0 额外负责发现和删除 TQ key;但所有 rank 都会取得各自的训练 Tensor,并参加模型 collective、梯度同步和 optimizer step。 + +### 3.5 Ray 和 torch.distributed 的职责不同 + +本方案仍然使用 Ray,但只因为 TQ 0.1.7 通过 Ray named actor 找 Controller。Consumer 不创建用于训练的 Ray actor,训练本身仍由 `torchrun` 和 `torch.distributed` 执行。 + +两种通信分别是: + +- Ray/TQ:Owner、Producer、每个 Consumer rank 连接共享 TQ;大 Tensor 通过 TQ backend 传输。 +- `torch.distributed`:训练 rank 之间广播小型 key/tag 命令、同步成功状态、训练模型 collective。 + +## 4. 共同配置以及“连接同一个 TQ”的实现 + +关键配置位于: + +```yaml +actor_rollout_ref: + rollout: + drafter: + training: + mode: offline + feature_store: + type: tq + path: null + transfer_queue: + enable: true + ray: + address: 127.0.0.1:6379 + namespace: speco-drafter + partition_id: speco_drafter_features + run_id: dspark-standalone-run + schema_version: 1 + poll_interval_seconds: 0.5 + drop_last: true +``` + +这些字段的含义如下: + +| 字段 | 使用者 | 含义 | +|---|---|---| +| `feature_store.type=tq` | Consumer | 选择流式 TQ source,而不是磁盘 shard/replay source | +| `feature_store.path=null` | Consumer | TQ 不从本地路径读文件,因此无需 path | +| `transfer_queue.enable` | Owner、Producer、Consumer | 开启 bridge 的 TQ 路径 | +| `ray.address` | 三端 | 连接同一个 Ray 集群 | +| `ray.namespace` | 三端 | 在同一 actor namespace 查找 named Controller | +| `partition_id` | 三端 | 对同一个 TQ KV 分区执行 put/list/get/clear | +| `run_id` | Producer、Consumer | 在共享 partition 中区分本次训练数据;Consumer 只接收匹配的样本 | +| `schema_version` | Producer、Consumer | 共同使用的数据协议版本 | +| `poll_interval_seconds` | Consumer rank 0 | ready 数量不足时的轮询间隔 | +| `drop_last` | Consumer | 第一版必须为 true;EOS 后不足 global batch 的尾样本被清理 | + +三端并不是通过共享 Python 对象得到这些配置。每个进程都各自读取相同取值,然后执行: + +```python +configure_transfer_queue(config) +connect_ray_cluster(ray_address, ray_namespace) +connect_transfer_queue_client() +``` + +`connect_ray_cluster()` 内部调用 `ray.init(address=..., namespace=...)`。`connect_transfer_queue_client()` 再使用与 Owner 相同的 native 配置调用 `tq.init(config)`;TQ 会优先在当前 Ray namespace 查找 Owner 创建的 named Controller,找到时忽略本次配置并只创建本地 Client。即使 Client 意外先于 Owner 初始化,也会使用同一份 backend/controller 配置,而不会按默认配置创建服务。之后所有 KV 操作都显式携带相同的 `partition_id`。 + +因此,“连接同一个 TQ”实际由三层身份共同决定:同一 Ray 集群、同一 namespace 下的同一 named Controller、同一 `partition_id`。 + +## 5. 一条样本在 TQ 中的实际格式 + +### 5.1 一条 key 对应一个 sample + +本协议没有把一个训练 batch 存成一个 TQ key。一条 key 对应一条独立训练样本。假设: + +```text +run_id = dspark-run-001 +sequence_no = 17 +sample_id = prompt-000017 +``` + +则 key 为: + +```text +drafter:v1:dspark-run-001:000000000017:prompt-000017 +``` + +`sequence_no` 是本次 run 内的样本顺序号,不是 batch 编号,也不是 optimizer step。Consumer 用它稳定排序,之后每次从有序 ready 列表前部取一个 global batch。 + +### 5.2 tag:用于轻量发现和过滤 + +该 key 的 tag 是普通小字典: + +```python +{ + "record_type": "sample", + "status": "ready", + "schema_version": 1, + "run_id": "dspark-run-001", + "sequence_no": 17, + "sample_id": "prompt-000017", + "algorithm": "DSPARK", +} +``` + +tag 存在 TQ 的 KV 元信息中。`kv_list(partition_id=...)` 返回 `key -> tag`,不需要先加载 hidden states。rank 0 正是依靠 tag 筛选当前 run、当前 schema、DSPARK 且状态为 ready 的记录。 + +### 5.3 fields:真正的 Tensor payload + +同一 key 的 fields 是一个 Tensor 字典: + +```python +{ + "input_ids": Tensor[int64, shape=[L]], + "loss_mask": Tensor[float32, shape=[L]], + "position_ids": Tensor[int64, shape=[L]], + "hidden_states": Tensor[dtype, shape=[L, D]], + "metadata_json": Tensor[uint8, shape=[M]], + # 以下是可选字段: + "last_hidden_states": Tensor[..., ...], + "target": Tensor[..., ...], + "target_logprobs": Tensor[..., ...], +} +``` + +这里 `L` 是 feature window 的 token 数,`D` 是目标模型 hidden size,`M` 是 metadata JSON 序列化后的 UTF-8 字节数。 + +`hidden_states` 等大 Tensor 只存于 fields,通过 TQ `kv_batch_get` 传输;不会放进 tag,也不会通过训练 rank 的 object broadcast。 + +### 5.4 metadata_json:内容丰富但仍随 fields 读取 + +TQ fields 只能承载 Tensor,因此结构化 metadata 被编码为 `uint8` Tensor。解码后的字典格式是: + +```python +{ + "schema_version": 1, + "run_id": "dspark-run-001", + "sample_id": "prompt-000017", + "sequence_no": 17, + "algorithm": "DSPARK", + "target_model_id": "/models/Qwen3-8B", + "target_model_revision": "main", + "tokenizer_fingerprint": "...", + "target_layer_ids": [35], + "hidden_states_layout": "token_major", + "hidden_dtype": "bfloat16", + "hidden_shape": [L, D], + "feature_length": L, + "full_sequence_length": 256, + "feature_start": 64, + "feature_end": 64 + L, + "use_logits": False, +} +``` + +字段分工是: + +- tag:只放发现、过滤、排序所需的小字段;`kv_list` 可直接得到。 +- fields:放训练 Tensor 和完整 metadata;只有被某个 rank 选中后才 `kv_batch_get`。 +- key:把 tag 和 fields 重新关联起来,也是清理记录时传给 `kv_clear` 的标识。 + +### 5.5 控制记录 + +控制记录与 sample 放在同一 partition,但通过 tag 的 `record_type=control` 区分。 + +Owner readiness key: + +```text +control:v1::owner-ready +``` + +EOS key: + +```text +control:v1::eos +``` + +EOS tag 包含 `status=eos` 和 `total_samples`。EOS 表示 Producer 不会再为本次 run 增加新样本;它不是一条训练样本。 + +## 6. 公共协议如何把 Producer 输出还原为训练对象 + +Producer 应调用: + +```python +fields = encode_sample(sample, metadata) +key = make_sample_key(metadata) +tag = make_ready_tag(metadata) +put_sample(key, fields, tag=tag) +``` + +`encode_sample()` 会将所有 Tensor detach、转到 CPU、整理为 contiguous,并统一 `input_ids/position_ids` 为 int64、`loss_mask` 为 float32。随后校验 token 长度、hidden shape 和 metadata 一致,再把 metadata JSON 编成 uint8 Tensor。 + +Consumer 的逆过程位于 `TQFeatureStore.get_many()`: + +```python +records = get_samples([entry.key for entry in entries]) +sample = decode_sample( + key=key, + tag=entry.tag, + fields=fields, + expected_config=self.expected_config, +) +``` + +`get_samples()` 最终调用一次 TQ `kv_batch_get(keys=[...], partition_id=...)`。bridge 将 TQ 返回的 batched TensorDict 或 mapping 拆成与请求 key 顺序一致的普通 fields 字典。 + +`decode_sample()` 随后: + +1. 检查必需 fields 是否存在。 +2. 将 `metadata_json` 从 uint8 Tensor 还原为字典和 `SampleMetadata`。 +3. 根据 metadata 重新计算 key,并与实际 key 比较。 +4. 比较 tag 与 metadata 的公共身份字段。 +5. 检查 Consumer 的 expected config。 +6. 将 Tensor detach 到 CPU,统一基础 dtype/shape。 +7. 检查所有主 Tensor 第一维等于 `feature_length`,hidden shape/dtype 与 metadata 一致。 +8. 构造 `DraftFeatureSample`。 + +输出不再是 TQ 专用对象,而是既有训练代码使用的: + +```python +DraftFeatureSample( + input_ids=..., + loss_mask=..., + position_ids=..., + hidden_states=..., + metadata=..., + ..., +) +``` + +这是能够复用原训练逻辑的关键边界:TQ 只负责上游存储和传输,`decode_sample()` 后的数据类型与磁盘 feature store 读取结果一致。 + +## 7. Consumer 从启动到结束的完整执行流程 + +### 阶段 1:launcher 检查配置并启动 torchrun + +执行者是 launcher 父进程。入口是 `verl_speco.draft_train_launcher.main()`。 + +当 override 中出现 `feature_store.type=tq`,`validate_tq_launch_config()` 会要求: + +- `training.transfer_queue.enable=true`; +- `training.transfer_queue.ray.address` 非空; +- `training.transfer_queue.run_id` 非空。 + +检查成功后构造: + +```text +python -m torch.distributed.run + --nnodes=... + --nproc_per_node=... + -m verl_speco.draft_train + <全部 Hydra overrides> +``` + +配置参数是普通子进程命令行参数。此阶段没有 TQ Tensor 传输。 + +### 阶段 2:每个 rank 初始化训练运行时 + +每个 torchrun 子进程进入 `run_standalone_draft_training()`,调用 `_init_distributed()` 得到 `rank/local_rank/world_size`,绑定本 rank GPU,然后构造原有 `DrafterBaseTrainer` 和 DSpark backend。 + +`speculative_algorithm=DSPARK` 决定 backend 和 DSpark 模型训练实现;`feature_store.type=tq` 只改变数据来源,不替换 trainer。 + +### 阶段 3:factory 创建 TQFeatureStore + +训练循环调用: + +```python +store = build_feature_store_from_config( + feature_store_cfg, + read_only=True, + transfer_queue_cfg=training_cfg.get("transfer_queue"), +) +``` + +factory 在 `type=tq` 时不读取 `feature_store.path`,而是把 sibling `training.transfer_queue` 交给 `TQFeatureStore.from_config()`。 + +TQ store 被限定为 `read_only=True`,意思是它是训练 Consumer source。这里的“read only”不表示永不修改 TQ;成功消费后仍可通过明确的 `clear_many()` 删除记录,但不会把它当作通用 feature writer。 + +### 阶段 4:所有 rank 分别连接同一个 TQ + +训练循环调用 `_connect_tq_store_across_ranks()`。每个 rank 都独立执行 `store.connect()`: + +```text +configure_transfer_queue +→ ray.init(address, namespace) +→ tq.init(same native config) 连接 named Controller +→ 本 rank 设置 _connected=True +``` + +之后 `_all_ranks_true()` 使用 `dist.all_reduce(MIN)` 汇总连接结果。只要一个 rank 连接失败,所有 rank 都停止,不允许部分 rank 进入后续 broadcast 或 FSDP collective。 + +这里没有“rank 0 建一个 client 给其他 rank 共用”。TQ client 是进程本地对象,N 个 rank 有 N 个 client,但它们指向同一 Controller/partition。 + +### 阶段 5:创建 TQFeatureDataLoader + +每个 rank 构造自己的 loader,参数包括相同的 `batch_size_per_gpu`、`world_size`、轮询间隔和 drop-last,以及不同的 `rank`。 + +假设: + +```text +world_size = 2 +batch_size_per_gpu = 2 +global_batch_size = 4 +``` + +那么只有 ready 数量至少为 4,rank 0 才发布一个 batch 命令。 + +### 阶段 6:rank 0 发现 ready key + +rank 0 首先检查 `owner_ready()`。Owner 尚未发布 readiness marker 时,rank 0 sleep 后继续轮询,不会让其他 rank 开始取数。 + +Owner ready 后,rank 0 调用 `list_ready()`,其底层是: + +```text +tq.kv_list(partition_id) +→ key -> tag +→ 按 record_type/status/run/schema/algorithm 过滤 +→ 按 (sequence_no, key) 排序 +``` + +此阶段没有读取 fields,因此 hidden states 尚未传到训练进程。 + +### 阶段 7:rank 0 切分 global batch + +若排序后的前四条是 `k0、k1、k2、k3`,`build_assignments()` 产生: + +```python +assignments = [ + [ReadyEntry(k0, tag0), ReadyEntry(k1, tag1)], # rank 0 + [ReadyEntry(k2, tag2), ReadyEntry(k3, tag3)], # rank 1 +] +``` + +每条样本只出现在一个 rank 的 assignment 中,因此各 rank 不会取得同一训练样本。这里采用连续、不重叠的切片。 + +rank 0 随后构造普通 Python 命令字典: + +```python +{ + "kind": "batch", + "global_keys": [k0, k1, k2, k3], + "assignments": [ + [{"key": k0, "tag": tag0}, {"key": k1, "tag": tag1}], + [{"key": k2, "tag": tag2}, {"key": k3, "tag": tag3}], + ], +} +``` + +`global_keys` 只用于 rank 0 在训练完成后一次清理整个 batch;`assignments` 用于每个 rank 知道自己应该 get 哪些 key。 + +### 阶段 8:小型命令通过 torch.distributed 广播 + +各 rank 同时进入: + +```python +dist.broadcast_object_list(payload, src=0) +``` + +rank 0 的 payload 中是上述字典,其他 rank 的初始值是 `None`。PyTorch 会序列化这个普通 Python 对象并广播给所有 rank。 + +这条边界只传输字符串、整数和小字典 tag。`hidden_states`、`input_ids` 等 fields 不在命令中,所以不会经 rank 0 中转,也不会随 broadcast 复制完整 global batch Tensor。 + +### 阶段 9:每个 rank 直接从 TQ 取本地 payload + +每个 rank 从 `assignments[self.rank]` 还原自己的 `ReadyEntry`: + +```python +local_entries = [_entry_from_wire(item) for item in assignments[self.rank]] +samples = self.store.get_many(local_entries) +``` + +在上述例子中: + +- rank 0 调用 `kv_batch_get(keys=[k0, k1], partition_id=...)`; +- rank 1 调用 `kv_batch_get(keys=[k2, k3], partition_id=...)`。 + +大 Tensor 的数据面因此是 TQ storage 到目标训练 rank,不经过训练 rank 0 的 Python 内存。每个 rank 得到两个 CPU `DraftFeatureSample`。 + +loader yield: + +```python +TQLocalBatch( + local_keys=[本 rank 的 key], + local_samples=[本 rank 的 DraftFeatureSample], + global_keys=[完整 global batch key] if rank == 0 else None, +) +``` + +非 rank 0 不保存 `global_keys`,避免多个 rank 都尝试 clear。 + +### 阶段 10:复用已有训练 batch 构造 + +训练循环识别 `TQLocalBatch` 后,只取: + +```python +samples = tq_local_batch.local_samples +``` + +然后调用原有接口: + +```python +batch = trainer.prepare_training_batch_from_samples( + materialized_samples, + step=optimizer_step, +) +``` + +TQ 路径禁止同时开启 `target_feature_pipeline`,因为样本已经包含目标模型 hidden states,不需要训练侧再访问 vLLM materialize 一次。 + +此时 TQ 专用的 key/tag 已不参与 DSpark 数学计算;训练接口看到的是普通 `DraftFeatureSample`,并按原逻辑整理 input ids、hidden states、mask、position ids 和 DSpark 训练所需输入。 + +### 阶段 11:所有 rank 同步 batch 是否可训练 + +每个 rank 判断 `batch is not None`,再通过 `_all_ranks_true()` 做 `all_reduce(MIN)`。 + +只有所有 rank 都成功构造 batch,才能进入训练。如果任一 rank 解码或 batch 构造失败,TQ 路径直接报错,而且这些 key 不会被删除。 + +### 阶段 12:执行原有 DSpark training step + +每个 rank 调用: + +```python +ok = await trainer.training_step_from_batch(batch, optimizer_step) +``` + +该调用复用既有模型 forward、DSpark loss(包括配置开启时的 L1 loss)、backward、梯度同步和 optimizer step。TQ 新代码没有重新实现 loss 或 optimizer。 + +之后再次以 `_all_ranks_true(ok)` 同步。只有所有 rank 都返回成功,才认为这一个 global batch 已经安全消费。 + +### 阶段 13:训练成功后由 rank 0 删除 global batch + +训练循环调用 `_clear_tq_batch_across_ranks()`: + +1. rank 0 使用 `tq_local_batch.global_keys` 调用 `loader.clear_completed_batch()`。 +2. loader 调用 `store.clear_many(global_keys)`。 +3. bridge 最终调用 `tq.kv_clear(keys=[k0,k1,k2,k3], partition_id=...)`。 +4. 所有 rank 通过 `all_reduce(MAX)` 同步 clear 是否失败。 + +删除发生在 optimizer step 全 rank 成功之后。不是“某个 rank get 完就删除”,因为 get 完只代表 Tensor 已读取,不能代表训练 step 已成功。 + +clear 成功后才增加 `successful_steps`,然后复用原有 metrics 和 checkpoint 调度。 + +### 阶段 14:EOS 和尾 batch + +当 ready 样本少于一个 global batch时,rank 0 查询 EOS: + +- 没有 EOS:说明 Producer 以后仍可能写入更多样本,sleep 后继续轮询。 +- 已有 EOS 且 ready 为空:广播 `{"kind": "stop"}`,所有 rank 结束迭代。 +- 已有 EOS 且存在不足一个 global batch 的尾样本:rank 0 先 clear 这些尾 key,再广播 stop。 + +第一版强制 `drop_last=true`,因此不会构造各 rank batch size 不一致的最后一步。 + +### 阶段 15:checkpoint 和退出清理 + +正常 step 完成后仍按原 `save_interval_steps` 保存 checkpoint;循环结束后按 `save_final_checkpoint` 决定是否保存最终 checkpoint。 + +`finally` 中每个 rank 调用 `store.close()`。对 `TQFeatureStore` 而言,这只是: + +```text +关闭本进程 TQ client +→ 如果本进程自行 ray.init,则 ray.shutdown() +``` + +它不会调用全局 `tq.close()`,不会杀死 Owner 创建的 Controller,也不会影响仍在运行的 Producer 或其他 rank。 + +## 8. 控制面和数据面的完整边界 + +| 数据 | 从哪里到哪里 | 传输机制 | 是否经过 rank 0 | +|---|---|---|---| +| 启动配置 | launcher 到 torchrun 子进程 | 命令行 Hydra overrides | 每个 rank 都收到 | +| ready key/tag | TQ Controller 到 rank 0 | `tq.kv_list` | 是,只有 rank 0 list | +| batch assignment | rank 0 到全部 rank | `dist.broadcast_object_list` | 由 rank 0 发出 | +| hidden states 等 fields | TQ storage 到被分配的 rank | `tq.kv_batch_get` | rank 1 的 Tensor 不经过 rank 0 | +| batch 准备/训练成功状态 | 全部 rank 之间 | Tensor `all_reduce` | collective,无单点 payload relay | +| clear 请求 | rank 0 到 TQ | `tq.kv_clear(global_keys)` | 只有 rank 0 发起 | +| 梯度和模型 collective | 训练 rank 之间 | 既有 PyTorch distributed/FSDP 路径 | 与 TQ 无关 | + +## 9. 当前“最简单校验”具体简单在哪里 + +`TQFeatureStore` 构造的 expected config 只固定: + +```python +ExpectedFeatureConfig( + run_id=<当前训练 run_id>, + schema_version=<当前 schema>, +) +``` + +TQ 会保留 Producer 写入的 `SampleMetadata.algorithm`,但不使用它选择训练 backend, +也不额外与启动配置比较。与原有离线 feature-store 训练一致,实际 trainer/backend 只由 +`rollout.drafter.speculative_algorithm` 和既有 backend factory 决定。 + +因此当前不会拿 Consumer 配置额外比较: + +- target model ID/revision; +- tokenizer fingerprint; +- target layer IDs; +- hidden layout; +- hidden dtype 的外部预期值。 + +但这不等于完全不校验。`decode_sample()` 仍然强制检查: + +- 必需 fields 存在; +- key、tag、metadata 三者身份一致; +- schema/run/algorithm 符合 Consumer; +- Tensor 类型正确; +- input/mask/position/hidden 长度一致; +- hidden 实际 shape/dtype 与该样本 metadata 一致; +- feature window 合法。 + +这满足“第一版少做外部模型身份检查”,同时避免把结构损坏或错 run 的数据送入训练。 + +## 10. 失败、删除和重复消费语义 + +当前实现遵循以下规则: + +1. 连接失败:所有 rank 同步停止。 +2. rank 0 list/EOS 失败:rank 0 广播 error 命令,其他 rank 不会永久等待 batch broadcast。 +3. 某 rank get/decode 失败:`_next_batch_across_ranks()` 将失败同步给全部 rank,不进入模型训练 collective。 +4. 某 rank 无法构造 batch:报错,global keys 保留在 TQ。 +5. 某 rank training step 失败:报错,global keys 保留在 TQ。 +6. 全 rank training step 成功:rank 0 clear 整个 global batch。 +7. clear 失败:错误传播到全部 rank,训练停止;不会把该 step 继续当成已正常完成。 +8. 达到 `max_steps`:循环停止;尚未选择的 ready 样本保留在 TQ。 + +第一版尚未实现完整的崩溃恢复协议。尤其是“optimizer step 已成功,但进程在 clear 前崩溃”时,key 仍存在;重新启动 Consumer 可能再次读取它。要实现严格 exactly-once,需要把 checkpoint step、已消费 sequence 或事务状态纳入协议。该能力应作为后续增强,而不是当前已实现能力。 + +## 11. 如何启动和检查 + +正式运行统一使用端到端 launcher;它负责 Ray、TQ Owner、Producer 和 Consumer 的启动与清理: + +```bash +bash examples/run_qwen3-8b_drafter_separate_training.sh +``` + +## 12. 已完成的测试 + +### 12.1 Consumer/factory/协议/launcher 单元测试 + +已执行: + +```text +python -m pytest \ + tests/unit/test_tq_consumer.py \ + tests/unit/test_draft_train_launcher.py \ + tests/unit/test_transferqueue_bridge.py \ + tests/unit/test_drafter_sample_protocol.py \ + tests/unit/test_draft_feature_store.py \ + -q +``` + +结果:`44 passed`。 + +### 12.2 真实 TQ 0.1.7 跨进程 smoke + +早期临时跨进程工具验证过以下行为,当前回归由协议、bridge、Consumer和launcher单元测试承担: + +- Owner 发布 owner-ready; +- 两条 sample 写入 TQ; +- Consumer 经 `TQFeatureStore` 和 `TQFeatureDataLoader` 读到两条样本; +- hidden shape 正确; +- Consumer clear 已完成 batch; +- EOS 后迭代停止; +- Consumer 只关闭本地 client,Owner 仍能继续观察完成标记并正常关闭。 + +实际 smoke 输出包含: + +```text +CLIENT_OK samples=2 shape=(3,4) +CLIENT_CLOSED_LOCAL_ONLY +OWNER_OBSERVED_SAMPLES_CLEARED +OWNER_CLOSED +``` + +### 12.3 当前环境未覆盖的部分 + +完整 `tests/unit/test_draft_training_loop.py` 在当前 Windows 环境无法完整收集,因为上游 `verl/ray` 依赖不齐;新增训练循环测试代码已通过 Python 编译检查,连接/clear helper 也通过针对性单元逻辑验证。真实多 GPU DSpark 训练仍需要在目标 Linux GPU 环境执行集成测试。 + +## 13. 当前限制和后续建议 + +当前第一版有意不实现以下复杂能力: + +1. Producer 本身尚未在本次 Consumer 改动中实现。 +2. TQ 公共协议和 Consumer 已不再写死 `DSPARK`;当前测试 Producer、启动脚本和已验证的 + feature 语义仍是 DSPARK。其他算法若能复用当前公共 dense fields,只需由对应 Producer + 生成正确的 `DraftFeatureSample`;若字段结构不同,则在协议模块增加对应 codec,不需要改 + TQ 的 key/tag 发现、rank 分配和 clear 流程。 +3. 只支持 `drop_last=true`。 +4. 不支持 TQ 与 `target_feature_pipeline.enabled=true` 同时开启。 +5. 不提供严格的 crash exactly-once 或 checkpoint/queue 联合恢复。 +6. rank 0 仍通过 `kv_list` 轮询整个 partition;数据量很大时可考虑 cursor/ready queue 优化。 +7. 当前外部 expected config 校验较简化,后续可把 model revision、tokenizer fingerprint、layer/layout/dtype 预期接入 Hydra 配置。 +8. 尚需在真实多机、多 GPU、Mooncake backend 环境验证吞吐、背压、Owner 生命周期和网络故障行为。 + +建议下一阶段优先完成 Producer,并严格复用 `drafter_sample_protocol.py`,不要在 Producer 另造一套 key/tag/fields 格式。完成 Producer 后,首先跑 world size 1 的端到端训练,再跑多 rank 验证每条 key 只分配给一个 rank、训练成功后只由 rank 0 clear。 + +## 14. 最终路径摘要 + +```text +Producer(待实现) + vLLM 并行 prefill + → DraftFeatureSample + SampleMetadata + → encode_sample 得到 Tensor fields + → TQ kv_put(key, fields, tag) + +Consumer rank 0 + kv_list 只取 key/tag + → 过滤并按 sequence_no 排序 + → 切出 global batch + → broadcast 每个 rank 的 key/tag assignment + +每个 Consumer rank + 取 assignments[rank] + → kv_batch_get 本 rank keys + → decode_sample 得到 DraftFeatureSample + → 原 prepare_training_batch_from_samples + → 原 DSpark training_step_from_batch + +全部 rank + 同步确认 optimizer step 成功 + → rank 0 kv_clear(global_keys) + → 原 metrics/checkpoint + → 下一批 + +Producer 发布 EOS + → rank 0 确认没有完整 global batch + → 清理不足一批的尾样本 + → broadcast stop + → 各 rank 关闭本地 TQ client 并退出 +``` diff --git a/docs/standalone_tq_drafter_resume_plan.md b/docs/standalone_tq_drafter_resume_plan.md new file mode 100644 index 00000000..17ad0586 --- /dev/null +++ b/docs/standalone_tq_drafter_resume_plan.md @@ -0,0 +1,448 @@ +# Standalone TQ 独立训练断点续训方案 + +## 1. 第一版目标 + +当前 DSpark 独立训练已经能够从 checkpoint 恢复模型权重、optimizer、LR scheduler、`optimizer_steps_total` 和 `training_steps`,但没有恢复数据进度。重启后 producer 会从文件开头重新生产,导致 checkpoint 之前已经训练的数据再次进入训练。 + +第一版采用最直接的方案: + +```text +consumer 记录已经成功训练的 sequence_no +→ checkpoint 保存这些 sequence_no +→ 重启时 producer 加载这个集合 +→ 读文件时跳过已消费编号 +→ 其他样本仍按原逻辑并行请求 vLLM +``` + +不修改 producer 的并发请求和完成顺序,不恢复旧 TQ,也不引入顺序发布。 + +## 2. 当前流程已经具备的条件 + +### 2.1 输入已经有编号 + +`standalone_tq_producer.py::read_inputs()` 当前执行: + +```python +record = replace(source_record, sequence_no=stats.input_count) +``` + +编号进入现有 TQ tag: + +```python +tag = { + "record_type": "sample", + "status": "ready", + "schema_version": 2, + "run_id": "...", + "sample_id": "...", + "sequence_no": 1234, +} +``` + +所以不需要新增样本身份协议,直接使用现有 `sequence_no`。前提是续训使用相同输入文件且行顺序不变。 + +### 2.2 consumer rank 0 已经知道 batch 编号 + +`TQFeatureDataLoader.__iter__()` 中 rank 0 执行: + +```python +ready = self.store.list_ready() +selected = ready[:global_batch_size] +``` + +`selected` 中每个 `ReadyEntry` 都有 `entry.tag["sequence_no"]`。因此 rank 0 已经知道本 global batch 实际使用了哪些输入编号。 + +### 2.3 消费成功边界已经明确 + +训练循环当前顺序为: + +```text +training_step_from_batch() +→ 所有 rank 确认成功 +→ rank 0 clear_completed_batch(global_keys) +→ successful_steps += 1 +→ 按间隔保存 checkpoint +``` + +规定: + +```text +训练成功 + TQ clear 成功 = 本 batch 的 sequence_no 已消费 +``` + +训练或 clear 失败时不能更新已消费集合。 + +## 3. checkpoint 增加什么 + +当前 checkpoint: + +```text +draft_step_60000/ +├── config.json +├── model.safetensors 或模型 shards +├── metadata.json +└── optimizer/ +``` + +新增: + +```text +draft_step_60000/consumed_sequence_nos.pt +``` + +内容为排序、去重的 CPU int64 tensor: + +```python +tensor([0, 1, 2, 4, 5, 8, ...], dtype=torch.int64) +``` + +允许存在间隔:alignment 失败的数据、尚未训练的 TQ 数据和仍在推理的数据都不在集合中。 + +空间开销约为: + +```text +100 万个编号:7.6 MiB +1000 万个编号:76 MiB +``` + +相比模型和 optimizer checkpoint 很小。 + +`metadata.json` 增加: + +```python +"standalone_data_progress": { + "version": 1, + "consumed_sequence_file": "consumed_sequence_nos.pt", + "consumed_sequence_count": int, + "input_fingerprint": {...}, +} +``` + +## 4. 为什么乱序推理不影响这个方案 + +假设 vLLM 完成顺序为: + +```text +5, 1, 8, 2, 4, 0, 3 +``` + +consumer 实际训练: + +```text +batch 1: [1, 5] +batch 2: [0, 2] +``` + +checkpoint 保存: + +```python +consumed_sequence_nos = tensor([0, 1, 2, 5]) +``` + +重启后 producer 只跳过 0、1、2、5,其余编号重新请求。因此不要求已消费数据连续,也不需要 producer 按顺序写 TQ。 + +## 5. consumer 修改 + +### 5.1 `TQLocalBatch` 携带 global 编号 + +文件:`verl_speco/trainer/tq_sample_source.py` + +改为: + +```python +@dataclass(frozen=True) +class TQLocalBatch: + local_keys: list[str] + local_samples: list[DraftFeatureSample] + global_keys: list[str] | None + global_sequence_nos: list[int] | None +``` + +rank 0 构造 command 时增加: + +```python +"global_sequence_nos": [ + int(entry.tag["sequence_no"]) + for entry in selected +] +``` + +只有 rank 0 需要保存 `global_sequence_nos`,其他 rank 保持 `None`。 + +### 5.2 clear 成功后更新集合 + +文件:`verl_speco/trainer/draft_training_loop.py` + +启动时: + +```python +consumed_sequence_nos = load_consumed_sequence_nos(drafter_cfg.model_path) +``` + +在 `_clear_tq_batch_across_ranks()` 成功返回以后: + +```python +if rank == 0: + consumed_sequence_nos.update( + tq_local_batch.global_sequence_nos or [] + ) +``` + +运行时使用 `set[int]` 便于去重;保存前转换为不可变 snapshot: + +```python +consumed_snapshot = torch.tensor( + sorted(consumed_sequence_nos), + dtype=torch.int64, +) +``` + +## 6. checkpoint 修改 + +不修改 `verl_speco/trainer/base_trainer.py`。该类同时被 co-train 使用,把独立训练的数据进度塞进它的公共 checkpoint 接口,会扩大影响范围。 + +独立训练仍先调用现有的: + +```python +trainer.save_checkpoint(step=step, wait=wait) +``` + +然后只在 `draft_training_loop.py` 的 `_save_standalone_checkpoint()` 中追加独立训练 sidecar: + +```text +draft_step_N/ +├── 原有模型、optimizer、scheduler 和 metadata +├── consumed_sequence_nos.pt +└── standalone_resume.json +``` + +其中 `standalone_resume.json` 记录 consumed 数量、输入 fingerprint 和 sidecar 版本。只有原 checkpoint 已成功保存后,才发布这两个文件。 + +使用临时文件原子写入: + +```python +temporary = checkpoint_path / "consumed_sequence_nos.pt.incomplete" +final = checkpoint_path / "consumed_sequence_nos.pt" +torch.save(consumed_snapshot, temporary) +os.replace(temporary, final) +``` + +异步保存时,`_save_standalone_checkpoint()` 先复制不可变 snapshot,再给现有 checkpoint future 注册 callback。callback 只在模型 checkpoint future 成功后原子写 sidecar,不能让后台 callback 直接读取仍在变化的 Python set。 + +如果进程恰好在模型 checkpoint 完成、sidecar 尚未完成时崩溃,这个目录不能用于“精确数据续训”,应回退到上一个同时具备完整模型 checkpoint 和完整 sidecar 的目录。 + +加载时校验: + +- 文件存在; +- 一维 `torch.int64`; +- 所有值非负; +- 已排序、无重复; +- tensor数量与 metadata一致。 + +旧 checkpoint没有该文件时,默认提示它只能恢复训练状态、不能精确恢复数据;严格模式下拒绝续训。 + +## 7. producer 修改 + +### 7.1 新配置 + +```yaml +speco: + standalone_tq_producer: + consumed_sequence_path: null +``` + +launcher 在续训时传入: + +```text +/consumed_sequence_nos.pt +``` + +### 7.2 启动时加载 + +文件:`verl_speco/standalone_tq_producer.py`。 + +```python +consumed_sequence_nos = load_consumed_sequence_nos( + producer_cfg.get("consumed_sequence_path") +) +``` + +### 7.3 扫描时跳过 + +需要拆开两个计数: + +```python +source_sequence_no # 所有扫描过的输入,跳过也递增 +queued_count # 本次真正送入 input_queue 的数量 +``` + +第一版保持 `iter_input_records()` 接口不变,在它产出记录后、tokenizer 和 vLLM 之前,根据 epoch 内扫描顺序算出全局 `sequence_no`: + +```python +sequence_no = source_sequence_no +source_sequence_no += 1 + +if sequence_no in consumed_sequence_nos: + continue + +record = replace(source_record, sequence_no=sequence_no) +await input_queue.put(request) +queued_count += 1 +``` + +因此恢复时仍需从输入文件开头顺序读取并解析一次,以重建稳定编号,但已消费行不会进入 tokenizer、input queue 或 vLLM。集合查询是平均 O(1),后续仍使用原来的多个 `request_worker()`,不会降低 producer 并发。 + +这里不是每个训练 step 都重新读取 checkpoint 文件。`consumed_sequence_nos.pt` 只在 producer 启动时加载一次;之后每扫描到一个源记录,只做一次内存集合查询。对百万级样本,主要额外成本是一次顺序读文件和 JSON/Parquet 行解析,通常远小于 tokenizer 和 vLLM 推理。只有实际测量发现启动扫描成为瓶颈后,才考虑给 input reader 增加解析前跳过、文件 offset 索引或 bitmap,第一版不做。 + +## 8. 多 epoch 编号 + +编号必须跨 epoch 连续: + +```text +数据集 10000 行 +epoch 0:0~9999 +epoch 1:10000~19999 +epoch 2:20000~29999 +``` + +加入 skip 后不能继续用一个 `stats.input_count` 同时表示扫描位置和排队数量,否则跳过数据后编号会改变。必须使用独立的 `source_sequence_no`。 + +## 9. 本次需要生产多少数据 + +独立训练的 `MAX_STEPS` 改为目标总 optimizer step,与首次训练配置保持一致: + +```text +MAX_STEPS = 训练完成时的总步数 +``` + +例如 checkpoint为 60000,目标总步数 930000: + +```bash +DRAFTER_PATH=/path/to/draft_step_60000 +MAX_STEPS=930000 +LR_WARMUP_STEPS=<仍使用首次训练时的原值> +``` + +checkpoint 已恢复 optimizer 和 scheduler 状态,所以 warmup 不会从头开始,当前 learning rate 和 scheduler 计数从 checkpoint 继续。用户不需要手算剩余步数,也不需要改其他训练超参;正常情况下只新增/替换 `DRAFTER_PATH`。 + +训练循环不再用本次进程的 `successful_steps < max_steps` 判断结束,而使用恢复后的总步数: + +```python +while max_steps <= 0 or trainer.optimizer_steps_total < max_steps: + ... +``` + +launcher计算 producer 配额时使用: + +```python +remaining_steps = max(max_steps - resumed_optimizer_step, 0) +max_samples = remaining_steps * batch_size_per_gpu * world_size +``` + +但 producer应使用 `queued_count` 判断本次配额。跳过的已消费编号不计入 `queued_count`。 + +该行为要求传入的是完整训练 checkpoint,里面具有 optimizer 和 scheduler 状态。只有模型权重的目录只能作为初始化权重,不能保证 learning rate、warmup 和 optimizer 状态精确续接。 + +## 10. 输入文件校验 + +编号只有在输入文件内容和顺序不变时才稳定。checkpoint至少保存: + +```python +"input_fingerprint": { + "path": str, + "size_bytes": int, + "mtime_ns": int, +} +``` + +推荐增加 SHA-256。续训时 fingerprint不一致则默认报错,避免旧 `sequence_no` 对应到新数据。 + +## 11. 崩溃语义 + +- vLLM已完成但未训练:不在 checkpoint集合,重启后重新推理。 +- 已写 TQ但未训练:不在集合,重启后重新推理。 +- optimizer成功但新 checkpoint未完成:恢复上一个模型和集合,重新训练上个 checkpoint之后的数据。 +- checkpoint完整:模型状态和 consumed集合对应同一步,已消费数据会被跳过。 +- checkpoint写到一半:恢复时忽略不完整目录,回退上一个 `complete=true` checkpoint。 + +## 12. 文件修改清单 + +### 新增 `verl_speco/trainer/standalone_resume.py` + +```python +load_consumed_sequence_nos(path) -> set[int] +save_consumed_sequence_nos(path, values) -> dict +build_input_fingerprint(path) -> dict +validate_input_fingerprint(saved, current) -> None +``` + +### 修改 `verl_speco/trainer/tq_sample_source.py` + +- `TQLocalBatch.global_sequence_nos`; +- rank 0从 selected tags提取编号。 + +### 修改 `verl_speco/trainer/draft_training_loop.py` + +- 启动时加载集合; +- clear成功后更新集合; +- checkpoint完成后写排序 int64 snapshot和 standalone sidecar; +- 使用 `optimizer_steps_total < max_steps` 作为独立训练终止条件。 + +`base_trainer.py` 及 co-train 入口不修改。standalone sidecar 的校验和加载全部收敛在 `standalone_resume.py` 与独立训练入口中。 + +### 修改 `verl_speco/standalone_tq_producer.py` + +- 加载 consumed集合; +- 拆分 source sequence和 queued count; +- tokenizer和vLLM前排除已消费编号; +- 其他并发逻辑保持不变。 + +### 修改 `verl_speco/standalone_tq_training_launcher.py` + +- 从 resume checkpoint获取 consumed文件; +- 把路径传给 producer; +- 校验输入 fingerprint。 +- 从 checkpoint读取已恢复 optimizer step,并仅生产剩余总步数需要的样本。 + +### 修改配置和 example + +- 增加 `consumed_sequence_path`; +- 说明只需设置 `DRAFTER_PATH=draft_step_N`;`MAX_STEPS`、warmup等保持首次训练配置。 + +## 13. 测试 + +1. consumer训练乱序编号 batch,clear成功后全部加入集合。 +2. 训练失败或 clear失败时不更新集合。 +3. checkpoint文件排序、去重,metadata数量一致。 +4. 输入 `0~9`,集合 `{0,2,5}`,producer只请求 `1,3,4,6,7,8,9`。 +5. 跳过的数据不占本次 `max_samples`配额。 +6. 多 epoch编号连续且重启后稳定。 +7. vLLM worker数量和请求并发与修改前一致。 +8. 端到端保存、重启后,producer不再请求 checkpoint集合中的编号。 + +## 14. 最终流程 + +首次训练: + +```text +producer并行请求 vLLM +→ TQ +→ consumer训练 batch +→ clear成功 +→ rank 0记录 sequence_no +→ checkpoint保存模型、optimizer和 consumed_sequence_nos.pt +``` + +续训: + +```text +加载 draft_step_N +→ 恢复模型、optimizer、LR和 step +→ 加载 consumed_sequence_nos.pt +→ 校验输入文件 +→ 创建新 Ray/TQ run +→ producer跳过已消费编号 +→ 其余样本继续并行请求 vLLM +``` + +该方案不改变 producer并发或 TQ消费顺序,改动集中在“consumer记录已成功数据”和“producer排除已消费数据”,适合作为第一版实现。 diff --git a/docs/standalone_tq_foundation_implementation.md b/docs/standalone_tq_foundation_implementation.md new file mode 100644 index 00000000..c7d39fd9 --- /dev/null +++ b/docs/standalone_tq_foundation_implementation.md @@ -0,0 +1,1030 @@ +# Standalone TQ 公共基础层实现说明 + +Last updated: 08/21/2026 + +## 1. 文档范围和已验证结论 + +本文解释当前仓库已经实现并测试通过的 TQ 公共基础层: + +```text +verl_speco/transport/drafter_sample_protocol.py +verl_speco/integration/transferqueue_bridge.py +verl_speco/config/speco_base.yaml +verl_speco/tq_owner.py +tests/unit/test_drafter_sample_protocol.py +tests/unit/test_transferqueue_bridge.py +pyproject.toml +``` + +当前已经实现: + +1. Producer 和 Consumer 共用的样本 key、tag、fields 和 metadata 协议; +2. 普通进程连接 Ray 集群; +3. TQ Owner 创建 named `TransferQueueController`; +4. 独立 Client 发现并连接同一个 Controller; +5. 单样本 put、元数据 list、批量 get 和批量 clear; +6. Owner 与 Client 不同的关闭边界; +7. 独立 Owner 入口和共享 Hydra 配置; +8. mock 单元测试和真实双进程 TQ 0.1.7 smoke test。 + +当前还没有实现: + +1. Producer 读取 JSONL、并发访问多个 vLLM endpoint 的完整 pipeline; +2. `feature_store.type=tq` 工厂分支; +3. `TQFeatureStore` 和 `TQFeatureDataLoader`; +4. rank 0 选择 global keys、各 rank 读取 local keys; +5. TQ batch 接入 DSpark optimizer step; +6. optimizer step 成功后的 rank 0 clear。 + +因此当前代码已经证明“两个独立进程能通过 Ray 连接同一个 TQ,并按共享协议批量传输多个样本”,但尚未接到正式 Producer 和 DSpark Consumer 主循环。 + +## 2. 运行时角色和术语 + +### 2.1 Ray head + +Ray head 是 Ray 集群的控制节点。TQ 0.1.7 使用 Ray 管理 named Controller 和 storage actors。 + +Ray 不传输 standalone hidden-state payload。当前代码没有对这些样本调用: + +```python +ray.put(hidden_states) +``` + +### 2.2 TQ Owner + +TQ Owner 是普通 Python OS 进程,入口为: + +```text +python -m verl_speco.tq_owner +``` + +它不是 Ray actor。它先调用 `ray.init(address=...)` 加入 Ray,再调用带完整配置的 `tq.init(config)` 创建 TQ Controller 和 storage。 + +Owner 是唯一允许调用全局 `tq.close()` 的进程。 + +### 2.3 Named TransferQueueController + +TQ 0.1.7 内部创建: + +```python +TransferQueueController.options( + name="TransferQueueController" +).remote(...) +``` + +`name` 将 Controller actor 注册到 Ray actor registry。其他进程加入同一 Ray address 和 namespace 后,通过: + +```python +ray.get_actor("TransferQueueController") +``` + +取得 actor handle,再读取 TQ backend 配置。 + +Controller 保存控制信息和数据位置;使用 MooncakeStore 时,大 tensor 本身保存在 MooncakeStore。 + +### 2.4 TQ Client + +Owner、Producer、每个 Consumer rank 都在各自 OS 进程内拥有独立 TQ Client。 + +普通 Client 也传入相同的 native 配置: + +```python +tq.init(native_config) +``` + +发现 named Controller并初始化本进程的 storage manager。Client 不是 Ray actor,Producer 和 torchrun rank 也不需要改成 Ray actor。 + +### 2.5 Partition、key、tag 和 fields + +当前固定 partition: + +```text +speco_drafter_features +``` + +TQ 中一条记录逻辑上是: + +```text +partition_id +└── key + ├── tag:轻量 dict,由 kv_list 发现 + └── fields:Tensor payload,由 kv_batch_get 读取 +``` + +## 3. 共享配置如何工作 + +共享配置定义在 `verl_speco/config/speco_base.yaml`: + +```yaml +transfer_queue: + enable: false + package_version: "0.1.7" + ray: + address: null + namespace: speco-drafter + partition_id: speco_drafter_features + run_id: null + schema_version: 1 + connect_timeout_seconds: 120 + poll_interval_seconds: 0.5 + drop_last: true + controller: + polling_mode: true + backend: + storage_backend: SimpleStorage + SimpleStorage: + total_storage_size: 100000 + num_data_storage_units: 8 + MooncakeStore: + auto_init: false + metadata_server: localhost:50050 + master_server_address: localhost:50051 + local_hostname: "" + protocol: tcp + global_segment_size: 4294967296 + local_buffer_size: 1073741824 + device_name: "" +``` + +### 3.1 Ray 连接字段 + +```yaml +ray: + address: 10.0.0.1:6379 + namespace: speco-drafter +``` + +它们决定当前进程连接哪个 Ray 集群,以及在哪个 namespace 查找 `TransferQueueController`。Owner、Producer 和所有 Consumer ranks 必须使用相同值。 + +### 3.2 SPECO 协议字段 + +```yaml +partition_id: speco_drafter_features +run_id: dspark-20260819-a +schema_version: 1 +``` + +这些字段用于路由和校验,不属于 TQ 原生配置。`run_id` 用于隔离不同 pipeline 的记录。 + +### 3.3 TQ 原生字段 + +```yaml +controller: ... +backend: ... +``` + +只有这些字段应传给 `tq.init(full_config)`。bridge 的 `_native_tq_config()` 会删除: + +```text +enable +package_version +ray +partition_id +run_id +schema_version +connect_timeout_seconds +poll_interval_seconds +drop_last +``` + +对象变化为: + +```text +完整 SPECO transfer_queue dict +→ _native_tq_config() +→ controller/backend等TQ字段 +→ OmegaConf DictConfig +→ tq.init(same native config) +``` + +## 4. Bridge 的进程内状态 + +`verl_speco/integration/transferqueue_bridge.py` 在每个 OS 进程内分别维护: + +```python +_state = { + "enabled": False, + "configured": False, + "initialized": False, + "config": None, + "owner": False, + "ray_initialized_here": False, + "ray_address": None, + "ray_namespace": None, +} +``` + +该 dict 不跨进程共享。Owner、Producer、每个 rank 分别拥有自己的 `_state`。 + +| 字段 | 含义 | +|---|---| +| `enabled` | 当前进程配置是否开启 TQ | +| `configured` | 是否调用过 `configure_transfer_queue()` | +| `initialized` | 当前进程是否执行过 `tq.init()` | +| `config` | 当前进程保存的普通 dict 配置 | +| `owner` | 当前进程是否创建了全局 Controller/Storage | +| `ray_initialized_here` | bridge 是否负责调用了本进程的 `ray.init()` | +| `ray_address/namespace` | 本进程的 Ray 连接信息 | + +`_state_lock` 只保护同一进程内多个线程同时初始化,不是分布式锁。 + +## 5. Owner 的完整启动数据流 + +Owner 入口是 `verl_speco/tq_owner.py`,直接加载 Hydra 主配置 +`verl_speco/config/speco_base.yaml`。`run_owner()` 会复制其中的 +`transfer_queue` 子配置,并仅在 Owner 进程的副本中强制设置 +`enable=true`;普通训练任务看到的共享默认值仍是 `false`。 + +### 阶段 1:读取配置 + +执行者:Owner OS 进程。 + +入口: + +```python +run_owner(config) +``` + +取得: + +```python +training_cfg = config.actor_rollout_ref.rollout.drafter.training +tq_cfg = training_cfg.transfer_queue +``` + +然后调用: + +```python +configure_transfer_queue(training_cfg) +``` + +该函数只把 OmegaConf 转成普通 dict并更新当前进程 `_state`,不会连接 Ray,也不会创建 TQ。 + +### 阶段 2:连接 Ray + +Owner 调用: + +```python +connect_ray_cluster(ray_address, namespace) +``` + +内部执行: + +```python +if not ray.is_initialized(): + ray.init(address=ray_address, namespace=namespace) +``` + +边界类型是 Ray control-plane connection。此时还没有传输训练 tensor。 + +### 阶段 3:创建 Controller 和 Storage + +Owner 调用: + +```python +start_transfer_queue_owner(tq_cfg) +``` + +执行顺序: + +1. `_extract_tq_config()` 得到普通 dict; +2. 检查 `enable=true`; +3. 检查 `TransferQueue` 包可用; +4. 防止本进程重复初始化; +5. `_native_tq_config()` 删除 SPECO 字段; +6. `_as_tq_config()` 转 OmegaConf; +7. `tq.init(native_config)` 创建 Controller、Storage 和 Owner Client; +8. 设置 `_state.owner=True`、`initialized=True`。 + +Ray 中形成: + +```text +Ray cluster / namespace +├── named actor: TransferQueueController +└── storage backend + ├── SimpleStorage actors + └── 或 MooncakeStore connection/process +``` + +### 阶段 4:发布 owner-ready + +调用: + +```python +publish_owner_ready(run_id, schema_version) +``` + +生成: + +```python +key = "control:v1::owner-ready" +fields = {"marker": torch.tensor([1], dtype=torch.uint8)} +tag = { + "record_type": "control", + "status": "owner_ready", + "schema_version": 1, + "run_id": run_id, +} +``` + +这是一条控制记录,不进入训练 batch。 + +### 阶段 5:常驻和关闭 + +Owner 安装 `SIGINT/SIGTERM` handler,并等待: + +```python +stop_event.wait() +``` + +收到信号后调用: + +```python +close_transfer_queue_owner() +``` + +执行: + +```text +验证owner身份 +→ tq.close() +→ 清理Controller/Storage +→ ray.shutdown() +``` + +Owner 必须在 Producer 和 Consumer 退出后才能关闭。 + +## 6. 普通 Client 如何连接同一个 TQ + +Producer 和每个 Consumer rank 后续使用相同顺序: + +```python +configure_transfer_queue(training_cfg) +connect_ray_cluster(ray_address, namespace) +connect_transfer_queue_client() +``` + +`connect_transfer_queue_client()` 最终调用: + +```python +tq.init(same_native_config) +``` + +TQ 0.1.7 内部通过: + +```python +ray.get_actor("TransferQueueController") +``` + +找到 Owner 创建的 Controller,读取 backend 配置,然后创建当前进程的 TQ Client。 + +对象和边界变化: + +```text +actor名称字符串 +→ Ray actor registry +→ Controller actor handle +→ Controller.get_config.remote() +→ TQ DictConfig +→ 当前进程TransferQueueClient +→ 同一个SimpleStorage/MooncakeStore +``` + +## 7. 一条具体样本的初始对象 + +真实 smoke test使用: + +```python +sample = DraftFeatureSample( + algorithm="DSPARK", + input_ids=torch.tensor([1, 2, 3]), # int64[3], CPU + loss_mask=torch.tensor([0.0, 1.0, 1.0]), # float32[3], CPU + position_ids=torch.tensor([0, 1, 2]), # int64[3], CPU + hidden_states=torch.arange( + 12, dtype=torch.float32 + ).reshape(3, 4), # float32[3,4], CPU +) +``` + +同时构造: + +```python +meta = SampleMetadata( + schema_version=1, + run_id="codex-batch-smoke", + sample_id="smoke-0000", + sequence_no=0, + algorithm="DSPARK", + target_model_id="smoke-target", + target_model_revision="smoke-revision", + tokenizer_fingerprint="smoke-tokenizer", + target_layer_ids=[0], + hidden_states_layout="dflash_aux", + hidden_dtype="float32", + hidden_shape=[3, 4], + feature_length=3, + full_sequence_length=3, + feature_start=0, + feature_end=3, + use_logits=False, +) +``` + +`DraftFeatureSample` 和 `SampleMetadata` 都是进程内 Python 对象,不直接经过 TQ。 + +## 8. Key 的生成和两个同名函数 + +共享协议调用: + +```python +make_sample_key(meta) +``` + +输出: + +```text +drafter:v1:codex-batch-smoke:000000000000:smoke-0000 +``` + +字段顺序: + +```text +drafter / schema version / run_id / 12位sequence_no / sample_id +``` + +`sequence_no` 是输入文件 record 序号,不是训练 batch 编号。 + +bridge 为兼容 PR #48 还保留另一个: + +```python +transferqueue_bridge.make_sample_key( + global_step, + replica_rank, + request_id, +) +``` + +它生成: + +```text +speco::: +``` + +standalone Producer 必须从 `verl_speco.transport.drafter_sample_protocol` import `make_sample_key`,不能使用 bridge 中的 PR #48 旧函数。 + +## 9. Tag 如何生成 + +```python +tag = make_ready_tag(meta) +``` + +输出: + +```python +{ + "record_type": "sample", + "status": "ready", + "schema_version": 1, + "run_id": "codex-batch-smoke", + "sequence_no": 0, + "sample_id": "smoke-0000", + "algorithm": "DSPARK", +} +``` + +tag 只包含发现、过滤和排序需要的标量/字符串。`list_samples()` 只返回 key/tag,不读取 hidden states。 + +## 10. `encode_sample()` 如何生成 fields + +调用: + +```python +fields = encode_sample(sample, meta) +``` + +### 10.1 校验 + +执行: + +```text +SampleMetadata.validate() +DraftFeatureSample.validate(strict=True) +``` + +随后检查: + +1. hidden states 是一个 dense tensor; +2. ids/mask/position 长度等于 `feature_length`; +3. hidden 第一维等于 `feature_length`; +4. hidden shape 等于 metadata; +5. hidden dtype 等于 metadata; +6. feature window 长度正确。 + +### 10.2 Tensor 规范化 + +```text +input_ids → CPU contiguous int64[L] +loss_mask → CPU contiguous float32[L] +position_ids → CPU contiguous int64[L] +hidden_states → CPU contiguous,保持模型dtype +``` + +没有 `position_ids` 时生成 `torch.arange(L, dtype=int64)`。 + +### 10.3 Metadata JSON 编码 + +```text +SampleMetadata dataclass +→ dict +→ JSON UTF-8 bytes +→ torch.uint8[M] +``` + +实现等价于: + +```python +raw = json.dumps(metadata).encode("utf-8") +metadata_json = torch.tensor(list(raw), dtype=torch.uint8) +``` + +### 10.4 最终 fields + +```python +fields = { + "input_ids": int64[3], + "loss_mask": float32[3], + "position_ids": int64[3], + "hidden_states": float32[3,4], + "metadata_json": uint8[M], +} +``` + +如果存在,协议也保留 `last_hidden_states`、`target` 和 `target_logprobs`。 + +## 11. Bridge 如何写入 TQ + +调用: + +```python +put_sample(key, fields, tag=tag) +``` + +bridge 执行: + +1. 检查 TQ 已启用; +2. 丢弃 fields 中非 tensor 值; +3. 确保本进程已经使用相同 native 配置执行 `tq.init(config)`; +4. 取得配置中的 partition; +5. 调用: + +```python +tq.kv_put( + key=key, + partition_id="speco_drafter_features", + fields=fields, + tag=tag, +) +``` + +使用 MooncakeStore 时,大 tensor 路径是: + +```text +Producer CPU tensor +→ Producer TQ Client +→ MooncakeStore +``` + +不是 Ray `ObjectRef`。 + +## 12. Consumer 如何发现 key + +调用: + +```python +records = list_samples() +``` + +内部调用: + +```python +tq.kv_list(partition_id="speco_drafter_features") +``` + +标准化返回类型: + +```python +dict[str, dict[str, Any]] +``` + +示例: + +```python +{ + "drafter:v1:...:smoke-0000": { + "record_type": "sample", + "status": "ready", + "run_id": "codex-batch-smoke", + "sequence_no": 0, + ... + } +} +``` + +bridge 兼容 `key → tag` 和 `partition → key → tag` 两种 wrapper。这一步不读取 fields。 + +## 13. Consumer 如何批量取样本 + +输入: + +```python +keys = [key0, key1] +``` + +调用: + +```python +records = get_samples(keys) +``` + +bridge 只调用一次: + +```python +result = tq.kv_batch_get( + keys=keys, + partition_id="speco_drafter_features", +) +``` + +TQ 0.1.7 返回带 batch 维的 TensorDict。bridge 检查 `result.batch_size`,再执行: + +```python +rows = [result[index] for index in range(len(keys))] +``` + +每行转成普通 dict,最终返回: + +```python +[ + (key0, fields0), + (key1, fields1), +] +``` + +返回顺序与输入 keys 一致。重复 key 会提前报错。 + +## 14. `decode_sample()` 如何恢复训练对象 + +调用: + +```python +sample = decode_sample( + key, + tag, + fields, + expected_config, +) +``` + +### 14.1 Metadata 解码 + +```text +metadata_json uint8[M] +→ bytes +→ UTF-8 +→ json.loads +→ dict +→ SampleMetadata.from_dict +``` + +### 14.2 身份一致性 + +代码根据 metadata 重新生成 key,要求输入 key 完全相等;然后逐项校验 tag 的: + +```text +record_type/status/schema_version/run_id/sequence_no/sample_id/algorithm +``` + +所以 key、tag 和 payload metadata 不能来自不同样本。 + +### 14.3 Consumer 合同 + +Consumer 提供: + +```python +ExpectedFeatureConfig( + run_id="codex-batch-smoke", + schema_version=1, + algorithm="DSPARK", + target_model_id="smoke-target", + target_model_revision=None, + tokenizer_fingerprint=None, + target_layer_ids=None, + hidden_states_layout="dflash_aux", + hidden_dtype="float32", +) +``` + +值为 `None` 的字段不检查;其他字段必须完全一致。正式 Consumer 应填写 target checkpoint、tokenizer、layers、layout 和 dtype,避免使用错误 target 特征。 + +### 14.4 输出 + +完成 tensor 类型、长度、shape、dtype 校验后,构造: + +```python +DraftFeatureSample.from_dict(payload, strict=True) +``` + +输出可以交给现有 `trainer.prepare_training_batch_from_samples()`。当前 smoke test验证到这里,正式 `TQFeatureDataLoader` 尚未实现。 + +## 15. EOS 控制记录 + +调用: + +```python +key, fields, tag = make_eos_record(run_id, total_samples) +``` + +输出: + +```python +key = "control:v1::eos" +fields = {"marker": torch.tensor([1], dtype=torch.uint8)} +tag = { + "record_type": "control", + "status": "eos", + "schema_version": 1, + "run_id": run_id, + "total_samples": total_samples, +} +``` + +EOS 表示不会再发布新样本。Consumer 应先 drain ready samples 再退出。协议函数已实现,正式 Producer/Consumer 尚未调用。 + +## 16. Clear 和数据生命周期 + +bridge 提供: + +```python +clear_samples(keys) +``` + +内部调用: + +```python +tq.kv_clear(keys=keys, partition_id="speco_drafter_features") +``` + +基础层只执行删除,不决定删除时机。正式 Consumer 必须遵守: + +```text +rank 0选择global keys +→ 各rank读取local keys +→ 所有rank完成同一optimizer step +→ 汇总global success +→ rank 0 clear global keys +``` + +不能在 `get_samples()` 后立即 clear,因为 get 成功不代表训练 step 成功。 + +## 17. Client close 和 Owner close + +### 17.1 Client close + +Producer/rank 调用: + +```python +close_transfer_queue_client() +``` + +执行: + +```text +tq.get_client() +→ 当前进程client.close() +→ 如果bridge负责ray.init,则ray.shutdown() +``` + +它不调用全局 `tq.close()`,不会 kill Controller。调用后本进程不能继续使用 TQ。 + +### 17.2 Owner close + +Owner 调用: + +```python +close_transfer_queue_owner() +``` + +执行: + +```text +tq.close() +→ Controller/Storage全局清理 +→ ray.shutdown() +``` + +Owner 若误用 Client close,bridge 会抛 `RuntimeError`。 + +## 18. PR #48 兼容边界 + +bridge 继续保留: + +```python +init_transfer_queue(config) +get_sample(key) +close_transfer_queue() +``` + +PR #48 的 `SpecoTaskRunner` 已是 Ray actor,所以不调用 standalone 的 `connect_ray_cluster()`。旧 worker 继续使用单 key `get_sample()`。 + +standalone 后续使用新增的: + +```python +list_samples() +get_samples(keys) +clear_samples(keys) +``` + +因此没有修改 PR #48 现有调用点的函数签名。 + +## 19. 依赖和命令入口 + +`pyproject.toml` 新增: + +```toml +[project.optional-dependencies] +transfer-queue = ["TransferQueue==0.1.7"] +``` + +安装: + +```bash +pip install -e ".[transfer-queue]" +``` + +TQ 0.1.7 会安装 Ray,并要求 `numpy<2.0.0`。若现有包要求 NumPy 2,需要隔离环境或重新解决依赖。 + +Owner 命令: + +```text +verl-speco-tq-owner +``` + +## 20. 单元测试 + +协议测试 `tests/unit/test_drafter_sample_protocol.py` 覆盖: + +1. encode/decode round trip; +2. key 格式; +3. tag 身份冲突; +4. Consumer contract 冲突; +5. hidden shape 冲突; +6. EOS 格式。 + +bridge 测试 `tests/unit/test_transferqueue_bridge.py` 覆盖: + +1. Ray address/namespace 参数; +2. Owner 只向 TQ 传原生配置; +3. Client 使用相同 native 配置调用 `tq.init(config)`; +4. put/list/get-many/clear; +5. batch 返回顺序; +6. Client close 不调用全局 close; +7. Owner 不能误用 Client close。 + +运行: + +```bash +python -m pytest \ + tests/unit/test_drafter_sample_protocol.py \ + tests/unit/test_transferqueue_bridge.py \ + -q +``` + +## 21. 真实双进程 smoke test + +早期用于该验证的临时双进程工具已经移除;正式入口统一由 +`verl_speco.standalone_tq_training_launcher` 管理 Owner、Producer 和 Consumer 生命周期。 + +Owner 路径: + +```text +连接Ray +→ tq.init(full config) +→ 写sample 0和sample 1 +→ 等待client-done +→ clear done marker +→ 全局关闭 +``` + +Client 路径: + +```text +连接同一个Ray +→ tq.init(same native config) +→ kv_list发现两个key +→ 一次kv_batch_get([k0,k1]) +→ 拆成两个fields dict +→ 分别decode_sample +→ clear两个sample keys +→ 写client-done +→ 只关闭本地client +``` + +已验证输出: + +```text +OWNER_READY keys=[k0, k1] +CLIENT_OK samples=2 shape=(3, 4) +CLIENT_CLOSED_LOCAL_ONLY +OWNER_OBSERVED_SAMPLES_CLEARED +OWNER_CLOSED +``` + +这证明: + +1. 两个普通进程能连接同一个 TQ; +2. named Controller 发现有效; +3. 0.1.7 的 `kv_list/kv_batch_get/kv_clear` 参数有效; +4. TensorDict batch 能按 key 顺序拆开; +5. 共享协议能恢复 `DraftFeatureSample`; +6. Client close 不会杀掉 Owner; +7. Owner 能最终统一关闭。 + +## 22. 当前完整路径总结 + +```text +Owner +→ ray.init(address, namespace) +→ tq.init(native config) +→ named TransferQueueController + +普通Client +→ ray.init(same address, same namespace) +→ tq.init(same native config) +→ 找到同一个Controller + +DraftFeatureSample + SampleMetadata +→ make_sample_key +→ make_ready_tag +→ encode_sample +→ fields + metadata_json tensor +→ bridge.put_sample +→ tq.kv_put +→ SimpleStorage/MooncakeStore + +Consumer/测试Client +→ bridge.list_samples +→ key + tag +→ bridge.get_samples(keys) +→ tq.kv_batch_get +→ TensorDict batch +→ 每个key对应一个fields dict +→ decode_sample +→ DraftFeatureSample + +正式训练成功后(待实现) +→ bridge.clear_samples(global_keys) + +Client退出 +→ close_transfer_queue_client + +所有业务进程退出 +→ Owner close_transfer_queue_owner +→ tq.close +→ ray.shutdown +``` + +## 23. 下一阶段接入约束 + +后续代码不能重新定义协议或直接访问 TQ 私有对象。 + +Producer 应复用: + +```text +SampleMetadata +make_sample_key +make_ready_tag +encode_sample +bridge.put_sample +make_eos_record +``` + +Consumer 应复用: + +```text +bridge.list_samples +bridge.get_samples +decode_sample +bridge.clear_samples +``` + +下一阶段需要新增: + +```text +verl_speco/trainer/tq_feature_store.py +verl_speco/trainer/tq_sample_source.py +feature_store.py 的 type=tq 分支 +draft_training_loop.py 的流式训练分支 +Producer入口、输入读取和并发vLLM文件 +``` + +这些文件应建立在本文已经实现和真实验证过的连接、协议、KV 和关闭接口之上。 diff --git a/docs/standalone_tq_producer.md b/docs/standalone_tq_producer.md new file mode 100644 index 00000000..bb766907 --- /dev/null +++ b/docs/standalone_tq_producer.md @@ -0,0 +1,212 @@ +# Standalone vLLM → TransferQueue Producer + +Last updated: 08/21/2026 + +本文解释 standalone Producer:它可以直接读取 verl 的 prompt-only Parquet(包括 +DAPO-Math-17k 的 chat-message `prompt`),也兼容已有 `prompt`/`response` 的 JSONL +或 Parquet。缺少 response 时由 target vLLM 生成,并在同一请求中提取 prompt 与 +output hidden states,之后把样本写到已存在的 TransferQueue(TQ)。 + +这条路径面向第一版 DSpark standalone 训练:Producer、TQ owner 和 Consumer +是三个独立 OS 进程;Ray 只用于让它们找到同一个 TQ Controller,hidden states +不通过 Ray object store 传输。 + +## 为什么需要这个 Producer + +此前仓库已经有两块基础能力: + +- `drafter_sample_protocol.py`:规定一条 TQ sample 的 key、tag、Tensor 字段和 + EOS record; +- `transferqueue_bridge.py` 与 `tq_owner.py`:负责连接 Ray/TQ、写读清理样本和 + owner 生命周期。 + +缺少的是把预先生成的文本变成 DSpark 训练特征并发布到 TQ 的独立进程。新增的 +Producer 补上这一段,不引入第二套协议或 feature store。 + +## 数据流 + +```text +verl prompt Parquet 或 prompt/response JSONL/Parquet + │ + │ 按文件顺序分配 sequence_no 和 sample_id + ▼ +Tokenizer + │ input_ids / loss_mask / feature window + ▼ +多个 vLLM endpoint(有界并发) + │ OpenAI completions 请求 → 临时 safetensors 文件 + ▼ +公共 hidden-state 转换函数 + │ DSpark DraftFeatureSample + SampleMetadata + ▼ +TransferQueue kv_put(一条输入记录对应一条 sample) + │ + ├─ put 成功:删除该请求的临时文件 + └─ 全部成功:写一个 EOS control record +``` + +Producer 在开始请求前会等到对应 `run_id` 的 `owner_ready` 控制记录。它不会创建 +Ray head、TQ Controller 或 storage backend;这些由 `verl-speco-tq-owner` 管理。 + +## 输入文件 + +输入可以是 JSONL 或 Parquet。`prompt` 可以是字符串,也可以是 verl 常用的 +`[{"role": ..., "content": ...}]` chat-message 列表。`response` 是可选字符串: +存在时直接 replay;不存在时由 target vLLM 生成。Parquet 通过 +`data.train_files` 直接传入,不需要转换。 + +```json +{"sample_id":"train-000017","prompt":"Question: 1 + 1 = ","response":"2"} +{"prompt":"Translate hello: ","response":"你好"} +``` + +- `sequence_no` 按非空行的文件顺序从 0 分配;并发完成顺序不会影响它。 +- `sample_id` 可选;省略时生成 `train-000000`、`train-000001` 等稳定值。 +- verl 数据的 `extra_info.index` 存在时会优先作为稳定 `sample_id`。 +- chat-message prompt 通过 target tokenizer 的 `apply_chat_template()` 编码,并加上 + generation prompt;不能把 `reward_model.ground_truth` 当作模型 response。 +- Producer tokenize `prompt` 和 `prompt + response`。后者必须以 prompt 的 token IDs + 为前缀;否则会报错,而不会猜测 response 的 loss-mask 边界。 +- `loss_mask` 中 prompt token 为 0,response token 为 1。 +- feature window 从 response 前一个 token 开始,长度由 + `max_feature_length` 限制;传给 vLLM 的 token IDs 截止于该 window 末端。 +- 其他 JSON 字段目前只作为 Producer 进程内来源元数据;第一版协议不会把它们写入 + TQ,所以 Consumer 不能读取这些字段。 + +## vLLM 与 hidden states + +对已有 response,Producer 使用 OpenAI-compatible completions API 做 prefill。 +对 prompt-only 数据,Producer 在一次请求中生成 response 并要求保存输出 hidden: + +```text +prompt= +max_tokens=<内部有界长度> +extra_body={ + "return_token_ids": true, + "kv_transfer_params": {"include_output_tokens": true} +} +``` + +响应必须同时满足: + +1. 若返回 `choices[0].prompt_token_ids`,它必须等于请求的 token IDs; +2. `kv_transfer_params.hidden_states_path` 必须存在; +3. 该文件必须含 `token_ids` 和形状为 `[seq, layers, hidden]` 的 `hidden_states`。 + +vLLM 0.23 已内置满足这个合同的 `ExampleHiddenStatesConnector`。不需要 SpeCo +Mooncake connector。在线服务必须关闭 chunked prefill,并显式配置一个 Producer +可见的临时目录。例如: + +```bash +export MODEL_PATH=/path/to/target-model +export HIDDEN_STATES_DIR=/dev/shm/speco-hidden-states +mkdir -p "${HIDDEN_STATES_DIR}" + +vllm serve "${MODEL_PATH}" \ + --host 0.0.0.0 \ + --port 8000 \ + --speculative-config \ + '{"method":"extract_hidden_states","num_speculative_tokens":1,"draft_model_config":{"hf_config":{"eagle_aux_hidden_state_layer_ids":[1,9,17,25,33,36]}}}' \ + --kv-transfer-config \ + "{\"kv_connector\":\"ExampleHiddenStatesConnector\",\"kv_role\":\"kv_producer\",\"kv_connector_extra_config\":{\"shared_storage_path\":\"${HIDDEN_STATES_DIR}\",\"use_synchronization_lock\":true}}" \ + --no-enable-chunked-prefill +``` + +上面的 layer IDs 只是 Qwen3-4B 示例。实际值必须按 target 模型和训练配置确定; +DSpark L1 开启时,vLLM 列表是 auxiliary layer IDs 加 final layer,而 Producer 的 +`TARGET_LAYER_IDS` 只填写 auxiliary 部分。 + +官方 connector 使用持久存在的 `.lock` 文件和 `flock` 协调异步落盘。Producer +读取前等待文件锁释放;TQ `put_sample` 成功后同时删除 safetensors 和 `.lock`。 + +`feature_from_vllm_payload()` 是从旧 replay 路径提取出的公共纯函数。它校验 token +对齐、选择 feature rows、拼接 auxiliary layers;DSpark L1 开启时额外拼接 final +hidden state。旧 replay 路径仍通过薄封装调用此函数,避免两套转换规则。 + +## TQ 写入和失败语义 + +每个输入 record 只写一个协议 key: + +```text +drafter:v1::<12位sequence_no>: +``` + +写入顺序是严格的: + +```text +加载临时 safetensors +→ 校验并转换 +→ TQ kv_put +→ 删除临时文件 +``` + +因此: + +- `kv_put` 失败时临时文件保留,且 Producer 不写 EOS; +- 任一请求、转换或写入失败会停止整条 Producer,不做自动重试或 endpoint 熔断; +- 只有所有 sample 都发布完成,才写 `control:v1::eos`; +- 进程退出时只调用 `close_transfer_queue_client()`,不会调用全局 `tq.close()`, + 不会销毁共享 Controller。只有 owner 可以关闭 TQ。 + +`max_pending_samples` 是简单背压:当前 run 的 ready sample 数达到该阈值时,新的 +vLLM 请求会暂停,等待 Consumer 清理已成功训练的 key。 + +## 配置与启动 + +默认 Producer 配置位于 +`speco.standalone_tq_producer`,TQ 连接配置仍位于 +`actor_rollout_ref.rollout.drafter.training.transfer_queue`。 + +必须设置的 Producer 字段: + +| 字段 | 含义 | +| --- | --- | +| `input_path` | 上述 JSONL 或 Parquet 文件 | +| `tokenizer_path` / `tokenizer_fingerprint` | 用于 tokenization 和 Consumer 合同校验 | +| `target_model_id` / `target_model_revision` | target checkpoint 身份 | +| `target_layer_ids` | auxiliary target layer IDs;DSpark L1 时 wire metadata 会额外写 `-1` 表示 final layer | +| `vllm_endpoints` / `vllm_model` | 一个或多个 OpenAI-compatible vLLM endpoint 与模型名 | + +必须与 owner/Consumer 一致的 TQ 字段: + +| 字段 | 固定要求 | +| --- | --- | +| `package_version` | `0.1.7` | +| `partition_id` | `speco_drafter_features` | +| `schema_version` | `1` | +| `run_id`、Ray address、Ray namespace | 三个进程必须相同 | + +单独调试时可以通过安装后的命令入口运行;正式训练由统一launcher启动Producer: + +```bash +verl-speco-tq-producer \ + actor_rollout_ref.rollout.drafter.training.transfer_queue.enable=true \ + actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.address= \ + actor_rollout_ref.rollout.drafter.training.transfer_queue.run_id= \ + speco.standalone_tq_producer.input_path= \ + speco.standalone_tq_producer.tokenizer_path= \ + speco.standalone_tq_producer.tokenizer_fingerprint= \ + speco.standalone_tq_producer.target_model_id= \ + speco.standalone_tq_producer.target_model_revision= \ + speco.standalone_tq_producer.target_layer_ids='[2,8,14,20,26]' \ + speco.standalone_tq_producer.vllm_endpoints='[http://node0:8000/v1]' \ + speco.standalone_tq_producer.vllm_model= +``` + +完整生命周期顺序仍是:Ray/TQ backend → TQ owner → Consumer → Producer → Consumer +drain → owner shutdown。Producer 完成不代表训练完成,EOS 只表示不会再有新样本。 +正式独立训练入口 +`examples/run_qwen3-8b_drafter_separate_training.sh` 会通过 +`verl_speco.standalone_tq_training_launcher` 自动管理这套生命周期;上面的 Producer +脚本仅用于单独调试 Producer。 + +## 测试覆盖与未验证项 + +新增测试覆盖:JSONL/真实 Parquet 解析、DAPO chat prompt、target response generation、 +token 边界、多个 endpoint 的并发限制、ready 队列背压、 +成功时 sample 后 EOS 与临时文件删除、失败时无 EOS 且保留临时文件,以及旧 EAGLE3 +转换路径仍可复用公共函数。 + +这些测试使用 fake vLLM/TQ。真实 Ray + TransferQueue + vLLM 的多进程 +联调没有在当前环境执行;运行前仍需确认 vLLM 版本能返回上述 +`hidden_states_path` 以及 TQ 0.1.7 依赖环境可用。 diff --git a/docs/standalone_tq_training_parameters.md b/docs/standalone_tq_training_parameters.md new file mode 100644 index 00000000..a7a48988 --- /dev/null +++ b/docs/standalone_tq_training_parameters.md @@ -0,0 +1,169 @@ +# Standalone TQ 独立训练参数说明 + +> Last updated: 08/27/2026 + +本文说明下面两个脚本暴露的参数: + +- `tools/run_qwen3-8b_drafter_hidden_state_vllm.sh`:按可见设备与 TP 自动启动一个或多个 target vLLM 服务。 +- `examples/run_qwen3-8b_drafter_separate_training.sh`:启动 Producer、TQ 和 DSpark Consumer 训练。 + +参数可以通过环境变量设置,例如: + +```bash +MAX_STEPS=1000 \ +BATCH_SIZE_PER_GPU=2 \ +DSPARK_CE_LOSS_ALPHA=0.1 \ +DSPARK_L1_LOSS_ALPHA=0.9 \ +bash examples/run_qwen3-8b_drafter_separate_training.sh +``` + +## 1. vLLM 服务参数 + +以下参数由 `run_qwen3-8b_drafter_hidden_state_vllm.sh` 使用。 + +| 参数 | 默认值 | 作用 | +| --- | --- | --- | +| `MODEL_PATH` | `/path/to/Qwen3-8B` | 所有 target vLLM 服务加载的模型路径。 | +| `DEVICE_ENV` | `ASCEND_RT_VISIBLE_DEVICES` | 控制设备可见性的环境变量。GPU 环境可设为 `CUDA_VISIBLE_DEVICES`。 | +| `VLLM_DEVICES` | `0,1,2,3,4,5` | 分配给vLLM的完整设备列表,脚本按连续的 `VLLM_TP` 张设备切成多个实例。 | +| `VLLM_TP` | `1` | 每个vLLM实例的 tensor parallel 大小;设备总数必须能被该值整除。 | +| `VLLM_HOST` | `127.0.0.1` | 所有vLLM服务监听的主机地址。 | +| `VLLM_BASE_PORT` | `8000` | 第一个实例的端口;后续实例依次使用 `8001`、`8002` 等。 | +| `VLLM_GPU_MEMORY_UTILIZATION` | `0.8` | 单个 vLLM 服务允许使用的设备显存比例。 | +| `VLLM_MAX_NUM_SEQS` | `256` | 单个 vLLM 服务最多同时调度的 sequence 数。 | +| `VLLM_HIDDEN_STATE_LAYER_IDS` | `[1,9,17,25,33,36]` | vLLM 导出的 hidden-state 层;前面的辅助层必须与训练侧 `DSPARK_TARGET_LAYER_IDS` 相同,最后一层用于构造 L1 loss 所需的 target 概率分布。 | +| `HIDDEN_STATES_DIR` | `/tmp/speco-vllm-hidden-states` | vLLM connector 临时写 hidden-state 文件的根目录,每个实例使用独立的 `service-N` 子目录。 | + +如果每个服务使用两张卡: + +```bash +MODEL_PATH=/nas/disk1/Qwen3-4B \ +VLLM_DEVICES=0,1,2,3,4,5 \ +VLLM_TP=2 \ +bash tools/run_qwen3-8b_drafter_hidden_state_vllm.sh +``` + +## 2. 基础训练参数 + +以下参数由 `run_qwen3-8b_drafter_separate_training.sh` 使用。 + +| 参数 | 默认值 | 作用 | +| --- | --- | --- | +| `MODEL_PATH` | `/path/to/Qwen3-8B` | Target model 路径,同时用于 tokenizer、模型配置和 target embedding/LM head。 | +| `TRAIN_FILE` | `/path/to/train_file.parquet` | Producer 读取的 JSONL 或 Parquet 数据文件。 | +| `DRAFTER_PATH` | 空 | 可选的已有 drafter/checkpoint 路径;为空时根据 target 配置从头初始化 DSpark。 | +| `DRAFT_CKPTS_DIR` | `/path/to/dspark_draft_checkpoints` | 保存 drafter checkpoint 的目录。 | +| `TRAIN_DEVICES` | `2,3` | Consumer 训练使用的设备。不能与任何 vLLM 实例占用的设备重叠。 | +| `TRAIN_GPUS` | `2` | 本节点启动的训练 rank 数,通常等于 `TRAIN_DEVICES` 中的设备数量。 | +| `DEVICE_ENV` | `ASCEND_RT_VISIBLE_DEVICES` | 训练侧设备可见性环境变量;GPU 环境可改为 `CUDA_VISIBLE_DEVICES`。 | +| `SPECO_VLLM_ENDPOINTS` | `[http://127.0.0.1:8000/v1,http://127.0.0.1:8001/v1]` | Producer 并行访问的 vLLM endpoint 列表。 | +| `VLLM_READY_TIMEOUT_SECONDS` | `120` | 启动训练前等待所有 vLLM endpoint 就绪的最长时间。 | +| `PYTHON_BIN` | `python3` | 启动 Python 模块所用的解释器。 | +| `PROJECT_NAME` | `verl_dspark_drafter` | 实验项目名。 | +| `EXP_NAME` | `qwen3_8b_dspark_separate_training` | 本次实验名称。 | + +## 3. Producer和请求并发参数 + +| 参数 | 默认值 | 作用 | +| --- | --- | --- | +| `VLLM_REQUEST_TIMEOUT` | `120` | 单次 generate 或 prefill HTTP 请求的超时时间,单位为秒。 | +| `VLLM_MAX_INFLIGHT_REQUESTS` | `16` | Producer 整体最多同时存在的 vLLM 请求数。 | +| `VLLM_PER_ENDPOINT_CONCURRENCY` | `4` | 每个 vLLM endpoint 独立的并发请求上限。 | +| `PRODUCER_INPUT_QUEUE_SIZE` | `32` | 已读取、等待 vLLM worker 处理的请求队列容量。 | +| `PRODUCER_PUBLISH_QUEUE_SIZE` | `16` | 已完成推理、等待写入 TQ 的样本队列容量。 | +| `PRODUCER_MAX_PENDING_SAMPLES` | `1024` | TQ 中尚未被 Consumer 训练并删除的样本数量上限,用于限制积压。 | +| `PRODUCER_PENDING_POLL_INTERVAL` | `0.5` | TQ 积压达到上限后,Producer 重新检查容量的间隔,单位为秒。 | +| `PRODUCER_MAX_SEQUENCE_LENGTH` | `8192` | prompt 和 response 处理前允许的最大总 token 长度。 | +| `PRODUCER_MAX_FEATURE_LENGTH` | `512` | 每个样本最终保留用于训练的最大 token 窗口。 | +| `PRODUCER_GENERATION_MAX_TOKENS` | `512` | 输入没有 response 时,vLLM 最多生成的 completion token 数。 | + +两个 endpoint 下的有效客户端并发近似为: + +```text +min(VLLM_MAX_INFLIGHT_REQUESTS, + endpoint数量 × VLLM_PER_ENDPOINT_CONCURRENCY) +``` + +`VLLM_MAX_NUM_SEQS` 是 vLLM 服务端调度上限;`VLLM_MAX_INFLIGHT_REQUESTS` 和 `VLLM_PER_ENDPOINT_CONCURRENCY` 是 Producer 客户端请求上限。 + +## 4. 训练过程参数 + +| 参数 | 默认值 | 作用 | +| --- | --- | --- | +| `MAX_STEPS` | `10` | 本次运行最多完成的训练 step 数。Producer 会根据它和全局 batch size 计算需要生成多少样本。 | +| `BATCH_SIZE_PER_GPU` | `2` | 每个训练 rank 每个 step 使用的样本数。全局 batch size 为该值乘训练 rank 总数。 | +| `SAVE_INTERVAL_STEPS` | `5` | 每隔多少 optimizer step 保存一次 checkpoint;设为0表示不做周期保存。 | +| `SAVE_FINAL_CHECKPOINT` | `true` | 训练结束时是否保存最终 checkpoint。 | +| `LEARNING_RATE` | `1e-6` | Drafter optimizer 的基础学习率。 | +| `LR_WARMUP_STEPS` | `0` | 学习率 warmup 的 step 数。 | +| `LR_SCHEDULER_TYPE` | `constant` | 学习率调度类型,支持 `constant`、`cosine`、`linear`、`global_cosine`。 | +| `LR_DECAY_STEPS` | `100` | 需要衰减的 scheduler 使用的衰减 step 数。 | +| `MIN_LR_RATIO` | `0.1` | 学习率衰减后的最小值与基础学习率的比例。 | +| `PARAM_OFFLOAD` | `true` | FSDP 是否把模型参数 offload 到 CPU。 | +| `OPTIMIZER_OFFLOAD` | `true` | FSDP 是否把 optimizer state offload 到 CPU。 | + +Producer 需要发布的样本数为: + +```text +MAX_STEPS × BATCH_SIZE_PER_GPU × TRAIN_GPUS × 节点数 +``` + +如果文件样本不足,Producer 会重新从文件开头读取并再次请求 vLLM,达到所需样本数后才发布 EOS。 + +## 5. DSpark模型和采样参数 + +| 参数 | 默认值 | 作用 | +| --- | --- | --- | +| `DSPARK_BLOCK_SIZE` | `7` | 每个 anchor 并行预测的 token 数。 | +| `DSPARK_NUM_ANCHORS` | `32` | 每个样本选取的 anchor 数;越大训练计算量和显存占用越高。 | +| `DSPARK_MAX_WINDOW` | `512` | DSpark 从输入样本中取出的最大训练窗口长度。 | +| `DSPARK_NUM_TARGET_LAYERS` | `5` | 输入 DSpark 的 target 辅助 hidden-state 层数量。 | +| `DSPARK_NUM_HIDDEN_LAYERS` | `5` | DSpark drafter 自身 transformer 层数。 | +| `DSPARK_TARGET_LAYER_IDS` | `[1,9,17,25,33]` | Target model 中采集的辅助 hidden-state 层编号。必须与 vLLM 服务侧配置的辅助层一致。 | +| `DSPARK_MARKOV_RANK` | `256` | Markov head 的低秩维度。 | +| `DSPARK_MARKOV_HEAD_TYPE` | `vanilla` | Markov head 类型。当前独立训练路径建议使用 `vanilla`。 | + +## 6. DSpark损失参数 + +| 参数 | 默认值 | 作用 | +| --- | --- | --- | +| `DSPARK_LOSS_MODE` | `full_vocab` | CE 计算方式,可使用 `full_vocab`、`restricted_ce` 或 `sampled_ce`。 | +| `DSPARK_SAMPLED_CE_NEGATIVES` | `0` | `sampled_ce` 模式下采样的负类数量。 | +| `DSPARK_LOSS_DECAY_GAMMA` | `7` | block 内不同预测位置的指数衰减系数。 | +| `DSPARK_CE_LOSS_ALPHA` | `0.1` | Token CE loss 在总 loss 中的权重。 | +| `DSPARK_L1_LOSS_ALPHA` | `0.45` | Draft 与 target token 概率分布之间的 L1 loss 权重。Target 概率由 final hidden state 和 LM head 计算;设为0可关闭 L1 loss。 | +| `DSPARK_L1_CHUNK_SIZE` | `0` | L1 loss 分块计算大小;0表示不主动分块。显存不足时可设置正整数。 | +| `DSPARK_CONFIDENCE_LOSS_ALPHA` | `0.0` | Confidence loss 权重。当前 standalone 协议没有 acceptance target,必须保持0。 | + +当前 DSpark 总损失为: + +```text +loss = DSPARK_CE_LOSS_ALPHA × ce_loss + + DSPARK_L1_LOSS_ALPHA × l1_loss +``` + +只训练 CE 的配置: + +```bash +DSPARK_CE_LOSS_ALPHA=1.0 +DSPARK_L1_LOSS_ALPHA=0.0 +``` + +## 7. 调试参数 + +| 参数 | 默认值 | 作用 | +| --- | --- | --- | +| `DSPARK_DEBUG_LOG` | `false` | 是否输出 DSpark forward 的详细调试日志。 | +| `DSPARK_DEBUG_LOG_FIRST_N` | `2` | 开启调试日志后,前多少次 forward 必定打印。 | +| `DSPARK_DEBUG_LOG_INTERVAL` | `100` | 前几次之后,每隔多少次 forward 打印一次调试信息。 | + +## 8. Layer ID一致性 + +默认配置为: + +```text +训练侧 DSPARK_TARGET_LAYER_IDS = [1,9,17,25,33] +服务侧 VLLM_HIDDEN_STATE_LAYER_IDS = [1,9,17,25,33,36] +``` + +服务侧前五项是辅助层,必须与训练侧完全相同。最后的 `36` 是 Qwen3-4B 对应的 final hidden-state layer,供 DSpark L1 loss 使用。更换 target model 或辅助层配置时,两边需要一起修改。 diff --git a/docs/standalone_vllm_tq_dspark_training_plan.md b/docs/standalone_vllm_tq_dspark_training_plan.md new file mode 100644 index 00000000..7a865b8b --- /dev/null +++ b/docs/standalone_vllm_tq_dspark_training_plan.md @@ -0,0 +1,1130 @@ +# 独立 vLLM Producer + TQ + DSpark Consumer 第一版方案 + +Last updated: 08/21/2026 + +## 1. 第一版要实现什么 + +只实现下面这条主链路: + +```text +verl prompt-only 数据或包含 prompt + response 的输入文件 +→ Producer 并发请求 vLLM prefill +→ Producer 将每条训练样本写入 TQ +→ Consumer 从同一个 TQ 取样本 +→ 独立 torchrun/FSDP DSpark 训练 +→ 一个 optimizer step 成功后删除该 step 的 TQ 样本 +``` + +第一版允许 **TQ 基础设施内部使用 Ray**,因为实测 `TransferQueue==0.1.7` 通过 Ray named actor 发现 `TransferQueueController`。Producer 和 Consumer 仍是普通 OS 进程,不改成 Ray actor;二者不使用 Ray RPC、Ray `ObjectRef` 或 Ray object store 传输训练样本。hidden-state payload 仍通过 TQ 的 `kv_put/kv_batch_get` 和配置的 MooncakeStore backend 传输。第一版不使用 DataProto/WorkerGroup,不生成长期 hidden-state feature store,也暂不实现复杂重试、自动重启和严格 checkpoint 数据恢复。 + +需要运行的组件: + +| 组件 | 数量 | 作用 | +|---|---:|---| +| vLLM server | 一个或多个 | 加载 target model,执行并行 prefill | +| Ray head | 1 个集群 | 保存 TQ named Controller/Storage actors;不承载 Producer/Consumer 业务 RPC | +| TQ owner | 1 个普通进程 | 连接 Ray,创建并保持 TQ controller/storage,最后统一关闭 TQ | +| Producer | 1 个进程 | 读文件、并发请求 vLLM、写 TQ | +| Consumer | 1 个 torchrun 任务 | 多个 rank 从 TQ 取数并训练 DSpark | + +Producer 和 Consumer 不互相调用,也不通过 HTTP 或 Ray 传样本。二者先连接同一个 Ray 集群,再通过携带相同 native 配置的 `tq.init(config)` 找到同一个 TQ,最后使用 TQ KV API 读写样本。已有 Controller 时 TransferQueue 0.1.7 会忽略后续配置并只连接;若 Client 意外先初始化,同一配置可避免默认 backend 抢先生效。 + +## 2. 共同的数据约定 + +这部分由两位开发者共同完成并先合入。建议文件: + +```text +verl_speco/transport/drafter_sample_protocol.py +tests/unit/test_drafter_sample_protocol.py +``` + +### 2.1 一个 key 对应一条样本 + +第一版固定: + +```text +一个输入文件 record +→ 一个 sequence_no +→ 一个 sample_id +→ 一个 TQ sample_key +→ 一个单样本 payload +``` + +`sequence_no` 是输入文件中的样本序号,不是 batch 编号。多个 sample keys 到 Consumer 后才组成训练 batch。 + +例如: + +```python +run_id = "dspark-20260818-a" +sequence_no = 17 +sample_id = "train-000017" + +partition_id = "speco_drafter_features" # 第一版沿用 PR 固定值 +sample_key = ( + "drafter:v1:dspark-20260818-a:" + "000000000017:train-000017" +) +``` + +### 2.2 Partition、key、tag 和 payload 的关系 + +TQ 中逻辑上是: + +```text +TQ 实例 +└── partition_id + └── sample_key + ├── tag + └── fields/payload +``` + +- TQ 实例:Producer 和 Consumer 共同连接的 controller/storage; +- `partition_id`:第一版沿用 PR 固定的 `speco_drafter_features`;每次启动新的 TQ owner,保证该 TQ 实例初始为空; +- `sample_key`:该分区中一条训练样本的地址; +- `tag`:轻量索引,Consumer 用 `kv_list()` 看见; +- `fields/payload`:真正的 Tensor 数据,Consumer 用 `kv_batch_get()` 读取。 + +Producer 写入: + +```python +tq.kv_put( + partition_id=partition_id, + key=sample_key, + fields=fields, + tag=tag, +) +``` + +Consumer 先发现 key: + +```python +all_records = tq.kv_list() +tags_by_key = all_records[partition_id] +``` + +这一步只拿 key 和 tag,不搬运 hidden states。 + +Consumer 再取数据: + +```python +result = tq.kv_batch_get( + partition_id=partition_id, + keys=selected_keys, +) +``` + +`kv_batch_get` 中的 batch 表示“一次读取多个独立 sample keys”,不是这些样本在 Producer 写入时就属于同一个对象。 + +### 2.3 Payload 字段 + +一个 sample key 对应的 `fields`: + +```python +fields = { + "input_ids": input_ids, # CPU int64[L] + "loss_mask": loss_mask, # CPU float32[L] + "position_ids": position_ids, # CPU int64[L] + "hidden_states": hidden_states, # CPU bf16[L,D] + "metadata_json": metadata_bytes, # CPU uint8[M] +} +``` + +| field | 含义 | Consumer 中的用途 | +|---|---|---| +| `input_ids` | 经过 feature window 选择后的 token IDs | DSpark token 输入 | +| `loss_mask` | 每个 token 是否参与 loss | 排除 prompt/padding/无效位置 | +| `position_ids` | 每个 row 对应的序列位置 | 位置编码和对齐校验 | +| `hidden_states` | target 指定层 hidden 沿最后一维拼接 | DSpark target/context 特征 | +| `metadata_json` | 模型、layers、layout、shape、样本身份 | Consumer 校验并恢复 metadata | + +符号: + +- `L`:这条训练 feature 保留的 token row 数; +- `H`:target model hidden size; +- `C`:DSpark context layer 数; +- L1 关闭:`D=C*H`,layout=`dflash_aux`; +- L1 开启:`D=C*H+H`,layout=`dflash_aux_plus_last`。 + +示例:`H=4096,C=5,L=1536`,开启 L1: + +```python +input_ids.shape == [1536] +loss_mask.shape == [1536] +position_ids.shape == [1536] +hidden_states.shape == [1536, 24576] +``` + +`metadata_json` 使用 JSON UTF-8 编码为 Tensor,因为参考 PR 的 bridge 只把 Tensor 放入 TQ fields: + +```python +raw = json.dumps(metadata, sort_keys=True).encode("utf-8") +metadata_bytes = torch.tensor(list(raw), dtype=torch.uint8) +``` + +### 2.4 Tag 字段 + +```python +tag = { + "record_type": "sample", + "status": "ready", + "schema_version": 1, + "run_id": "dspark-20260818-a", + "sequence_no": 17, + "sample_id": "train-000017", + "algorithm": "DSPARK", +} +``` + +tag 只存用于发现和筛选的 flat scalar/string。Consumer 只选择: + +```text +record_type=sample +status=ready +schema_version=1 +run_id=当前 run +algorithm=DSPARK +``` + +### 2.5 Metadata 字段 + +`metadata_json` 解码后至少包含: + +```python +metadata = { + "schema_version": 1, + "run_id": "dspark-20260818-a", + "sample_id": "train-000017", + "sequence_no": 17, + "algorithm": "DSPARK", + "target_model_id": "/models/Qwen3-8B", + "target_model_revision": "revision-or-checksum", + "tokenizer_fingerprint": "sha256:...", + "target_layer_ids": [2, 8, 14, 20, 26, -1], + "hidden_states_layout": "dflash_aux_plus_last", + "hidden_dtype": "bfloat16", + "hidden_shape": [1536, 24576], + "feature_length": 1536, + "full_sequence_length": 1800, + "feature_start": 264, + "feature_end": 1800, + "use_logits": False, +} +``` + +其中: + +- `target_model_id/revision`:hidden states 来自哪个 target checkpoint; +- `tokenizer_fingerprint`:Producer 使用的 tokenizer/template 版本; +- `target_layer_ids`:vLLM 返回和参与拼接的层; +- `hidden_states_layout`:Consumer 应如何拆 hidden 最后一维; +- `feature_length`:payload 中四个主要 Tensor 的第一维; +- `full_sequence_length`:完整 prompt+response 的 token 数; +- `[feature_start,feature_end)`:feature 在完整序列中的范围。 + +### 2.6 共享协议接口 + +Producer 和 Consumer 都 import 同一个模块,不各自手写字段名: + +```python +@dataclass(frozen=True) +class SampleMetadata: + schema_version: int + run_id: str + sample_id: str + sequence_no: int + algorithm: str + target_model_id: str + target_model_revision: str + tokenizer_fingerprint: str + target_layer_ids: list[int] + hidden_states_layout: str + hidden_dtype: str + hidden_shape: list[int] + feature_length: int + full_sequence_length: int + feature_start: int + feature_end: int + use_logits: bool + +def make_sample_key(meta: SampleMetadata) -> str: ... +def make_ready_tag(meta: SampleMetadata) -> dict: ... +def encode_sample(sample, meta: SampleMetadata) -> dict[str, Tensor]: ... +def decode_sample(key, tag, fields, expected_config) -> DraftFeatureSample: ... +def make_eos_record(run_id: str, total_samples: int): ... +``` + +Producer 使用 `SampleMetadata/make_sample_key/make_ready_tag/encode_sample`;Consumer 使用 `decode_sample`。 + +`SampleMetadata` Python 对象本身不经过 TQ: + +```text +Producer SampleMetadata +→ JSON +→ uint8 Tensor +→ TQ metadata_json +→ uint8 Tensor +→ JSON +→ Consumer metadata dict +``` + +`decode_sample()` 负责: + +1. 解码 `metadata_json`; +2. 校验 key、tag、metadata 中的 sample 身份一致; +3. 校验模型、tokenizer、layers 和 layout 与 Consumer 配置一致; +4. 校验 Tensor 必需字段、dtype 和 shape; +5. 返回现有 `DraftFeatureSample`。 + +### 2.7 EOS + +Producer 完成全部输入后写一个控制 record: + +```python +eos_key = f"control:v1:{run_id}:eos" +eos_fields = {"marker": torch.tensor([1], dtype=torch.uint8)} +eos_tag = { + "record_type": "control", + "status": "eos", + "schema_version": 1, + "run_id": run_id, + "total_samples": total_samples, +} +``` + +EOS 不进入训练 batch。Consumer 看到 EOS 后继续处理剩余 ready samples;`EOS 已出现且 ready 为空` 时结束。 + +## 3. TQ 怎么启动,Producer 和 Consumer 怎么连接同一个 TQ + +### 3.1 已验证的 TQ 0.1.7 连接机制 + +`TransferQueue==0.1.7` 没有提供“把 Controller 地址直接传给第二个进程”的高层连接接口。`tq.init(config)` 会先尝试以下已有 Controller 连接逻辑;存在时忽略传入配置,不存在时才用配置创建服务: + +```python +_TQ_CONTROLLER = ray.get_actor("TransferQueueController") +conf = ray.get(_TQ_CONTROLLER.get_config.remote()) +_maybe_create_tq_client(conf) +``` + +因此所有进程必须先加入同一个 Ray 集群。Ray 在本方案中只承担 TQ 控制面:保存 named `TransferQueueController`、返回 TQ 配置、管理 TQ 创建的 actor。业务进程不通过 Ray 发送 sample key 或 hidden-state tensor。 + +实际连接链路是: + +```text +TQ owner:ray.init(address) → tq.init(full_tq_config) → 创建 named Controller +Producer:ray.init(address) → tq.init(same config) → ray.get_actor() → 创建本地 TQ client +Consumer rank 0..N:ray.init(address) → tq.init(same config) → ray.get_actor() → 创建各自 TQ client +``` + +### 3.2 直接移植并扩展 PR #48 的 bridge + +参考文件: + +```text +C:/Users/xxyyrr/Desktop/上班/verl-SpeCo/ + verl_speco/integration/transferqueue_bridge.py +``` + +第一版不新增 `standalone_tq.py`,直接把该文件移植到目标项目并扩展。保留 PR 已有的 import 检查、进程内幂等状态、`kv_put`、`kv_batch_get`、TensorDict 解包和 owner-only shutdown;新增 Ray 显式连接、批量 list/get/clear 和本地 client 关闭接口。 + +目标文件必须实现以下函数,而不只是提供一个笼统的 transport class: + +```python +def configure_transfer_queue(config: Mapping[str, Any]) -> bool: ... +def connect_ray_cluster(ray_address: str, namespace: str | None = None) -> None: ... +def start_transfer_queue_owner(tq_config: Mapping[str, Any]) -> None: ... +def connect_transfer_queue_client() -> None: ... +def make_sample_key(run_id: str, sequence_no: int, sample_id: str) -> str: ... +def put_sample(key: str, fields: dict[str, Tensor], tag: dict[str, Any]) -> None: ... +def list_samples() -> dict[str, dict[str, Any]]: ... +def get_samples(keys: list[str]) -> list[tuple[str, dict[str, Tensor]]]: ... +def clear_samples(keys: list[str]) -> None: ... +def close_transfer_queue_client() -> None: ... +def close_transfer_queue_owner() -> None: ... +``` + +逐个函数的责任如下。 + +#### `configure_transfer_queue(config)` + +- 从 Hydra/OmegaConf 中读取 Ray address、可选 namespace、固定 partition、run ID、schema version 和 TQ native backend 配置; +- 转成普通 Python dict,保存在进程内 `_state`; +- 校验 `TransferQueue==0.1.7` 可 import; +- 不连接 Ray,不创建 TQ,不产生跨进程副作用; +- 返回该进程是否启用了 TQ。 + +#### `connect_ray_cluster(ray_address, namespace)` + +- 若 `ray.is_initialized()` 为 false,调用 `ray.init(address=ray_address, namespace=namespace)`; +- 若已经初始化,校验现有 Ray context 指向期望集群,不能悄悄连到另一个本地 Ray; +- Owner、Producer 和所有 torchrun ranks 都调用它; +- 此函数只建立当前 OS 进程到 Ray control plane 的连接。 + +#### `start_transfer_queue_owner(tq_config)` + +- 仅由 `tq_owner.py` 调用; +- 前置条件是 `connect_ray_cluster()` 已成功; +- 调用一次 `tq.init(OmegaConf.create(tq_config))`; +- 将 `_state.owner=True`、`_state.initialized=True`; +- TQ 0.1.7 会创建名为 `TransferQueueController` 的 Ray actor,并创建所选 storage backend; +- 重复调用必须报错,不能启动第二套同名 Controller。 + +#### `connect_transfer_queue_client()` + +- 由 Producer 和每个 Consumer rank 调用; +- 前置条件是当前进程已经连接 Ray; +- 调用 `tq.init(same native config)`,通过 `ray.get_actor("TransferQueueController")` 发现 owner;已有 Controller 时配置会被忽略,意外抢先时则以相同配置创建; +- 只创建当前进程的 TQ client,不创建新的 Controller; +- 成功后设置 `_state.initialized=True`;重复调用直接返回。 + +#### `put_sample/list_samples/get_samples/clear_samples` + +- 全部固定使用 `_SPECO_TQ_PARTITION = "speco_drafter_features"`; +- `put_sample()` 调用单样本 `tq.kv_put()`; +- `list_samples()` 调用 `tq.kv_list(partition_id=...)`,只返回 key/tag 元数据; +- `get_samples()` 一次调用 `tq.kv_batch_get(keys=...)`,再按输入 key 顺序解包为普通 dict; +- `clear_samples()` 调用 `tq.kv_clear(keys=..., partition_id=...)`; +- 这些函数不调用 Ray RPC,不把 tensor 放入 Ray object store。 + +#### `close_transfer_queue_client()` 与 `close_transfer_queue_owner()` + +TQ 0.1.7 的公共 `tq.close()` 会 kill 共享 Controller,所以两者不能写成同一个实现: + +- `close_transfer_queue_client()`:通过 0.1.7 已公开的 `tq.get_client()` 取得当前进程 client并调用它的 `close()`,然后 `ray.shutdown()`;绝不能调用会 kill Controller 的全局 `tq.close()`。这个 client close 只用于进程退出阶段,关闭后本进程不能再次调用 TQ; +- `close_transfer_queue_owner()`:仅当 `_state.owner=True` 时调用 `tq.close()`,清理 Controller/Storage,最后 `ray.shutdown()`; +- Producer 或任一训练 rank 提前退出都不能关闭全局 TQ。 + +### 3.3 共享配置 + +Owner、Producer 和 Consumer 必须使用相同的 Ray address/namespace,并通过同一个 named Controller 获得 backend 配置。第一版 partition 固定,不按任务动态创建: + +```yaml +transfer_queue: + enable: true + package_version: "0.1.7" + ray: + address: "ray-head-node:6379" + namespace: "speco-drafter" + partition_id: "speco_drafter_features" + run_id: "dspark-20260819-a" + schema_version: 1 + backend: + storage_backend: MooncakeStore + MooncakeStore: + auto_init: false + metadata_server: "node0:50050" + master_server_address: "node0:50051" + local_hostname: "" + protocol: tcp + global_segment_size: 4294967296 + local_buffer_size: 1073741824 + device_name: "" +``` + +`partition_id/run_id/schema_version` 用于过滤和校验样本;它们不能帮助进程发现 TQ。真正让三个任务连接到同一 TQ 的是“连接同一个 Ray 集群和 namespace,然后找到同名 Controller”。 + +依赖也必须锁定并单独验证。实机安装 `TransferQueue==0.1.7` 会安装 Ray,并要求 `numpy<2.0.0`;当前测试环境中的 `twinkle-kit` 要求 `numpy>=2.0.0`,两者冲突。开发时应使用专门的 TQ/训练环境或重新确认整套依赖约束,不能直接把 0.1.7 安装进已有生产环境后忽略 resolver warning。 + +### 3.4 `tq_owner.py` 要实现的入口和函数 + +新增: + +```text +verl_speco/tq_owner.py +``` + +`tq_owner.py` 建议明确实现: + +```python +def install_signal_handlers(stop_event: threading.Event) -> None: ... +def publish_owner_ready(run_id: str, schema_version: int) -> None: ... +def wait_until_stopped(stop_event: threading.Event) -> None: ... +def run_owner(config: DictConfig) -> int: ... +def main() -> None: ... +``` + +`run_owner()` 的执行顺序必须是: + +```text +configure_transfer_queue(config) +→ connect_ray_cluster(ray.address, ray.namespace) +→ start_transfer_queue_owner(full TQ native config) +→ put owner_ready 控制 record +→ 安装 SIGINT/SIGTERM handler +→ 保持 owner 进程存活 +→ 收到停止信号 +→ close_transfer_queue_owner() +``` + +Owner 必须常驻。TQ 0.1.7 创建 Controller 时没有设置 `lifetime="detached"`,不能在初始化后立即退出。 + +### 3.5 启动和关闭顺序 + +第一版由外部脚本管理全生命周期: + +```text +1. ray start --head,记录 Ray address +2. 启动 Mooncake metadata/master(若 auto_init=false) +3. 启动 TQ owner;owner 连接 Ray并调用 tq.init(full config) +4. 等待 owner_ready +5. 启动一个或多个 vLLM servers +6. 启动 Consumer;每个 torchrun rank 连接 Ray,然后 tq.init(same native config) +7. 启动 Producer;连接 Ray,然后 tq.init(same native config) +8. Producer 写 EOS,关闭本地 client并退出 +9. Consumer drain、保存 final checkpoint,所有 ranks 关闭本地 client并退出 +10. 给 TQ owner 发送 SIGTERM;仅 owner 执行 tq.close() +11. 等 owner 退出后执行 ray stop +12. 停止 Mooncake 服务 +``` + +外部脚本要用 `trap` 保证异常退出也按“Producer/Consumer → owner → Ray → Mooncake”的顺序清理。不能在 Consumer rank 的 `finally` 中调用全局 `tq.close()`。 + +## 4. Producer 要实现什么 + +### 4.1 Producer 完整顺序 + +```text +读取共享配置 +→ 连接 TQ并校验 owner_ready +→ 初始化 tokenizer +→ 初始化多个 vLLM endpoint clients +→ 流式读取输入文件 +→ 为每条输入分配 sequence_no/sample_id +→ 缺少 response 时由 target vLLM 生成;构造 input_ids/loss_mask +→ 并发请求 vLLM prefill +→ 读取 vLLM hidden-state 临时结果 +→ 转换成 DSpark DraftFeatureSample +→ 构造 SampleMetadata +→ encode_sample 得到 fields/tag/key +→ TQ kv_put 一条 sample +→ 删除该请求临时文件 +→ 所有输入完成后写 EOS +→ close_transfer_queue_client()并退出 +``` + +### 4.2 并发模型 + +Producer 是一个进程,内部并发请求多个 endpoint: + +```text +InputReader +→ bounded asyncio input_queue +→ N 个 RequestWorker +→ bounded publish_queue +→ TQ Publisher +``` + +- `vllm_endpoints` 是列表; +- 每个 endpoint 有独立 semaphore; +- 总并发由 `max_inflight_requests` 限制; +- input/publish queue 必须有上限; +- TQ `kv_put` 如果是同步 API,用 `asyncio.to_thread()` 调用; +- TQ ready 数达到 `max_pending_samples` 时暂停继续请求,形成简单背压。 + +`sequence_no` 在 InputReader 中分配,不按 vLLM 完成顺序分配。因此并发乱序不会改变 sample key。 + +### 4.3 vLLM 结果转换 + +复用现有 `TargetFeatureReplayer` 的: + +- OpenAI-compatible vLLM 请求; +- `prompt_token_ids` 校验; +- `kv_transfer_params.hidden_states_path`; +- safetensors 加载; +- `[seq,layers,hidden]` 校验; +- feature positions 选择; +- aux layers flatten; +- DSpark L1 时拼 final hidden。 + +不要复制 `_feature_from_vllm_payload()`;把纯转换逻辑提取成公共函数,让旧 file backend 和新 Producer 共同调用。 + +临时文件顺序: + +```text +加载 +→ 校验/转换 +→ TQ put 成功 +→ 删除 +``` + +第一版仍允许每个并发请求产生短期临时文件,但不生成全量 hidden-state 数据集。 + +### 4.4 Producer 文件分工 + +| 文件 | 实现内容 | +|---|---| +| `verl_speco/standalone_tq_producer.py` | CLI/Hydra 入口,连接 TQ,启动 asyncio pipeline,写 EOS和汇总指标 | +| `verl_speco/producer/input_reader.py` | 流式读输入、分配 `sequence_no/sample_id`、tokenize、构造 loss mask | +| `verl_speco/producer/vllm_feature_client.py` | endpoint pool、并发控制、HTTP 请求、临时结果读取和删除 | +| `verl_speco/trainer/target_feature_replay.py` | 提取可复用的 vLLM payload → `DraftFeatureSample` 纯转换函数 | +| `tests/unit/test_drafter_sample_protocol.py` | 协议编码、字段和 shape 测试(两人共同) | +| `tests/integration/test_tq_producer_smoke.py` | fake/短 vLLM → TQ sample + EOS | + +Producer 开发者同时负责移植/扩展 `integration/transferqueue_bridge.py` 和实现 `tq_owner.py`,因为这部分与 TQ 写入和连接直接相关。 + +### 4.5 Producer 各文件的函数级实现规格 + +#### `verl_speco/standalone_tq_producer.py` + +需要实现: + +```python +@dataclass +class ProducerStats: + input_count: int + published_count: int + failed_count: int + pending_bytes: int + +async def publish_one(result: PreparedFeature, transport) -> str: ... +async def run_producer(config: DictConfig) -> ProducerStats: ... +def validate_producer_config(config: DictConfig) -> None: ... +def main() -> None: ... +``` + +`main()` 只负责 Hydra/日志/退出码。`run_producer()` 是可测试的业务入口,执行:连接 Ray → 连接 TQ client → 创建 InputReader 和 vLLM client pool → 启动有界 asyncio pipeline → 等待所有请求及 `kv_put` 完成 → 发布 EOS → 关闭本地 client。它不能启动 Ray head、不能创建 TQ owner、不能调用全局 `tq.close()`。 + +`publish_one()` 接收已经完成转换的一条 `PreparedFeature`,调用共享协议的 `encode_sample()` 得到 `(key, fields, tag)`,再通过 bridge 的 `put_sample()` 发布。只有 `put_sample()` 成功返回,`published_count` 才增加,vLLM 临时文件才允许删除。 + +#### `verl_speco/producer/input_reader.py` + +需要实现: + +```python +@dataclass(frozen=True) +class InputRecord: + sequence_no: int + sample_id: str + prompt: str + response: str | None + source_metadata: dict[str, Any] + +def iter_input_records(path: str) -> Iterator[InputRecord]: ... +def tokenize_record(record: InputRecord, tokenizer, config) -> TokenizedRequest: ... +def build_loss_mask(input_ids: Tensor, prompt_length: int) -> Tensor: ... +``` + +`iter_input_records()` 流式读取 JSONL/Parquet,不把全文件载入内存,并按文件顺序分配稳定的 `sequence_no`。已有 response 时 `tokenize_record()` 直接拼接;prompt-only verl 数据通过 chat template 编码后由 target vLLM 生成 response,并设置 `include_output_tokens=true` 同步提取输出 hidden states。 + +#### `verl_speco/producer/vllm_feature_client.py` + +需要实现: + +```python +@dataclass(frozen=True) +class VllmEndpoint: + base_url: str + max_concurrency: int + +class VllmFeatureClientPool: + async def start(self) -> None: ... + async def prefill(self, request: TokenizedRequest) -> RawVllmFeature: ... + async def close(self) -> None: ... + +async def request_prefill(endpoint, request) -> VllmResponse: ... +def choose_endpoint(endpoints, state) -> VllmEndpoint: ... +def load_hidden_state_result(response) -> RawVllmFeature: ... +def delete_temporary_result(raw: RawVllmFeature) -> None: ... +``` + +`prefill()` 必须允许多个 coroutine 同时运行;全局 semaphore 限制总并发,每个 endpoint 另有独立 semaphore。`request_prefill()` 只负责 HTTP 请求和响应校验;`load_hidden_state_result()` 负责读取 vLLM 返回的临时 safetensors/path。临时结果的删除不放在 `load_hidden_state_result()`,而由 `publish_one()` 成功后触发。 + +#### `verl_speco/trainer/target_feature_replay.py` + +把当前类内部的纯转换部分抽成: + +```python +def feature_from_vllm_payload( + payload: RawVllmFeature, + request: TokenizedRequest, + feature_config: FeatureContract, +) -> DraftFeatureSample: ... +``` + +它不进行 HTTP、不访问 TQ、不删除文件,只完成 shape/layout/layer 校验、feature-position 选择、aux layer flatten 和可选 L1 final hidden 拼接。现有 file replay backend 和新 Producer 都调用这一函数,避免两套转换规则。 + +#### Producer 启动配置 + +统一 launcher 向 `verl_speco.standalone_tq_producer` 提供同一套: + +```text +RAY_ADDRESS / Ray namespace +run_id / schema_version / 固定 partition +Mooncake/TQ backend 配置 +输入文件和 tokenizer/model 配置 +vLLM endpoint 列表 +max_inflight_requests / per_endpoint_concurrency +``` + +正式运行不再保留单独的角色 shell wrapper。 + +## 5. Consumer 要实现什么 + +### 5.1 不新写另一套训练器 + +继续使用现有入口: + +```text +draft_train_launcher.py +→ draft_train.py +→ trainer/draft_training_loop.py +→ DrafterBaseTrainer +→ DSparkTrainerBackend +``` + +训练模式继续使用现有 standalone `training.mode=offline`,用 `feature_store.type=tq` 选择 TQ 数据源。不复制 FSDP、loss、optimizer 或 checkpoint 代码。 + +当前 `feature_store.type` 的作用是告诉 `build_feature_store_from_config()` 应创建哪一种训练数据来源: + +| 当前 type | 对象 | 数据来源 | +|---|---|---| +| `torch_shard` | `TorchShardFeatureStore` | 本地 `.pt` shard,当前 separate-training 示例使用它 | +| `token_replay` | `TokenReplayFeatureStore` | 保存 token replay 的 shard | +| `vllm_safetensors` / `safetensors` | `VllmSafetensorsFeatureStore` | 已生成的 safetensors feature | +| `jsonl_token_replay` / `jsonl` | `JsonlTokenReplayFeatureStore` | JSONL token/text replay | + +第一版新增: + +```yaml +actor_rollout_ref: + rollout: + drafter: + training: + mode: offline + feature_store: + type: tq + path: null + shuffle: false + repeat: false +``` + +这里 `offline` 的含义是“独立于 RL 的 standalone drafter training”,不等于数据必须来自磁盘。 + +不能只给现有工厂增加一个 `type=tq` 然后继续使用通用 `DraftFeatureDataLoader`。当前 loader 会: + +```python +keys = list(store.iter_keys(...)) +``` + +它假设数据集 keys 是一个静态快照;空 store 会直接结束,并且各 rank 独立枚举时可能在 Producer 持续写入的过程中看到不同快照。TQ 是流式、会新增并删除 keys 的数据源,所以 `type=tq` 必须选择专用的 `TQFeatureDataLoader`,由 rank 0 统一发现和分配 keys。 + +#### 5.1.1 当前磁盘 feature store 为什么每个 rank 都会取 keys + +`draft_train_launcher.py` 通过 `torchrun` 启动多个训练进程。每个进程都是一个 rank,并且每个 rank 都会独立进入 `run_standalone_draft_training()`、创建 `DraftFeatureDataLoader`、执行它的 `__iter__()`。当前 loader 的核心逻辑是: + +```python +keys = list( + self.store.iter_keys( + shuffle=self.shuffle, + seed=self.seed + epoch, + ) +) +rank_keys = keys[rank::world_size] + +for key in rank_keys: + batch.append(self.store.read(key)) +``` + +因此当前并不是 rank 0 枚举 keys 后再发送给其他 rank,而是所有 rank 都访问相同的静态 feature-store 路径: + +```text +rank 0:iter_keys() 得到完整静态列表 → keys[0::world_size] → read 自己的 samples +rank 1:iter_keys() 得到完整静态列表 → keys[1::world_size] → read 自己的 samples +... +``` + +例如 store 中固定存在: + +```python +keys = ["k0", "k1", "k2", "k3", "k4", "k5", "k6", "k7"] +``` + +当 `world_size=2` 时,两个 rank 都先得到上述完整列表,然后分别计算: + +```python +# rank 0 +rank_keys = keys[0::2] # ["k0", "k2", "k4", "k6"] + +# rank 1 +rank_keys = keys[1::2] # ["k1", "k3", "k5", "k7"] +``` + +这种实现能够成立,是因为 `.pt`、safetensors 或 replay 文件在训练期间是静态数据集。只要所有 rank 使用同一路径和相同 shuffle seed,`iter_keys()` 就会产生相同顺序的 key 快照,各 rank 可以无通信地算出互不重叠的子集。 + +#### 5.1.2 TQ 为什么必须改成 rank 0 发现 keys + +TQ 中的 keys 会在训练期间动态变化:Producer 持续 `kv_put`,训练成功后 Consumer 又执行 `kv_clear`。如果所有 rank 仍然各自调用 `kv_list`,不同调用时刻可能看到不同快照: + +```python +# rank 0 较早调用 +rank0_keys = ["k0", "k1", "k2", "k3"] + +# Producer 随后写入 k4、k5,rank 1 较晚调用 +rank1_keys = ["k0", "k1", "k2", "k3", "k4", "k5"] +``` + +各 rank 再独立切片后,可能得到不同数量的 local samples,进而无法保证它们以相同顺序进入 forward、backward 和梯度 collective。为此,TQ 专用 loader 必须把控制面和数据面分开: + +```text +控制面:rank 0 执行 kv_list,固定本 step 的 global_keys,并向各 rank 分发 local_keys +数据面:每个 rank 使用自己的 local_keys 直接执行 kv_batch_get,从 TQ/Mooncake 读取 hidden-state payload +``` + +rank 0 只发送较小的 key/tag 元数据,不读取并转发其他 rank 的 hidden-state tensor。所有 rank 仍然都会连接同一个 TQ,也都会调用批量读取接口;只有动态 key 的发现和本 step 的 global batch 决策集中在 rank 0。 + +因此 `feature_store.type=tq` 不是只替换底层 `read()`:它同时改变了 key 的发现、分配、等待和删除语义,需要专用 `TQFeatureDataLoader`,并要求训练循环在所有 rank 成功完成 optimizer step 后,由 rank 0 对本 step 的 `global_keys` 执行 `kv_clear`。 + +### 5.2 Consumer 完整顺序 + +```text +torchrun 启动多个 ranks +→ 每个 rank 初始化 torch.distributed +→ 每个 rank 连接同一个 TQ +→ rank 0 校验 owner_ready,并 broadcast 结果 +→ 初始化现有 DSpark trainer +→ rank 0 kv_list 查找 ready sample keys +→ rank 0 选一个 global batch并分给各 rank +→ 每个 rank kv_batch_get 自己的 local keys +→ decode_sample 得到 list[DraftFeatureSample] +→ prepare_training_batch_from_samples() +→ training_step_from_batch() +→ 所有 rank 汇总 success +→ 成功后 rank 0 kv_clear 这个 global batch 的 keys +→ 继续下一批 +→ 看到 EOS 且 ready 为空 +→ 保存 final checkpoint +→ 所有 ranks close_transfer_queue_client()并退出 +``` + +### 5.3 多 rank 如何分 key + +例如: + +```text +world_size=2 +batch_size_per_gpu=2 +global batch size=4 +``` + +rank 0 选出: + +```python +global_keys = ["k10", "k11", "k12", "k13"] +assignments = [ + ["k10", "k11"], # rank 0 + ["k12", "k13"], # rank 1 +] +``` + +通过 `dist.scatter_object_list` 或 broadcast 发送短字符串列表。每个 rank 直接从 TQ 读取自己的 payload,不通过 rank 0 转发 hidden states。 + +### 5.4 从 TQ 到训练 batch + +每个 rank: + +```python +records = tq_transport.get_samples(local_keys) + +samples = [ + decode_sample( + key=key, + tag=tags_by_key[key], + fields=fields, + expected_config=expected_contract, + ) + for key, fields in records +] + +batch = trainer.prepare_training_batch_from_samples( + samples, + step=optimizer_step, +) + +ok = await trainer.training_step_from_batch( + batch, + optimizer_step, +) +``` + +`records` 是多个独立单样本 payload;`samples` 是变长 `DraftFeatureSample` 列表。Consumer transport 层不要直接 `torch.stack()`,现有 trainer/backend 负责对齐和组 batch。 + +### 5.5 删除与结束 + +第一版采用简单逻辑: + +```text +所有 rank get/decode/train 都成功 +→ all_reduce(global_success)=True +→ rank 0 kv_clear(global_batch_keys) +``` + +任何 rank 失败都不 clear。程序报错退出,第一版不自动恢复。 + +EOS 后不足一个 global batch 的尾部,第一版使用 `drop_last=True`:记录数量,rank 0 clear 这些尾部 keys,然后正常结束。 + +checkpoint 仍使用现有 `save_interval_steps` 和 final checkpoint;第一版不保证崩溃后已 clear 数据能够严格重放。 + +### 5.6 Consumer 文件分工 + +| 文件 | 实现内容 | +|---|---| +| `verl_speco/trainer/feature_store.py` | 工厂增加 `type=tq`,返回 `TQFeatureStore`;TQ 类型不要求 `path` | +| `verl_speco/trainer/tq_feature_store.py` | 实现固定 partition 上的 list/get-many/clear/EOS,不伪装静态 `iter_keys()` | +| `verl_speco/trainer/tq_sample_source.py` | 实现专用 `TQFeatureDataLoader`:rank 0 动态 list/filter/sort、key 分配、各 rank get/decode、EOS/tail | +| `verl_speco/trainer/draft_training_loop.py` | 保持 `mode=offline`;`type=tq` 时跳过 `feature_store.path` 必填检查并选择专用 loader;调用现有 prepare/train,global success 后 clear,final checkpoint | +| `verl_speco/draft_train_launcher.py` | `feature_store.type=tq` 时不要求 feature-store path,透传 torchrun/TQ 配置 | +| `verl_speco/config/speco_base.yaml` | 新增 TQ connection、run、poll、batch 等默认配置 | +| `tests/unit/test_tq_sample_source.py` | key 过滤/排序/分 rank、EOS、decode 调用测试 | +| `tests/integration/test_dspark_tq_consumer_smoke.py` | 两 rank 取不同样本,完成训练并 clear | + +### 5.7 Consumer 各文件的函数级实现规格 + +#### `verl_speco/trainer/feature_store.py` + +修改现有工厂: + +```python +def build_feature_store_from_config(feature_store_cfg, read_only=False): + store_type = str(feature_store_cfg.get("type", "torch_shard")).lower() + if store_type == "tq": + return TQFeatureStore.from_config(feature_store_cfg) + ... +``` + +要求: + +- 保留 `torch_shard/token_replay/vllm_safetensors/jsonl_token_replay` 的当前行为; +- `type=tq` 时不读取 `path`; +- 工厂只创建对象,不在 import 阶段连接 Ray/TQ; +- `TQFeatureStore` 是流式数据源适配器,不强行实现有误导性的静态 `iter_keys()` 和逐条 `read()`。 + +#### `verl_speco/trainer/tq_feature_store.py` + +需要实现: + +```python +@dataclass(frozen=True) +class ReadyEntry: + key: str + tag: dict[str, Any] + +class TQFeatureStore: + @classmethod + def from_config(cls, cfg) -> "TQFeatureStore": ... + def connect(self) -> None: ... + def list_ready(self, run_id: str) -> list[ReadyEntry]: ... + def get_many(self, entries: list[ReadyEntry]) -> list[DraftFeatureSample]: ... + def clear_many(self, keys: list[str]) -> None: ... + def read_eos(self, run_id: str) -> EosMetadata | None: ... + def close_local(self) -> None: ... +``` + +`connect()` 调用共享 bridge 的 `connect_ray_cluster()` 和 `connect_transfer_queue_client()`。`list_ready()` 调用 `list_samples()` 后只保留 `tag.status=ready`、匹配 `run_id/schema_version` 的数据 key,并按 `(sequence_no,key)` 排序。`get_many()` 一次批量读取 fields,然后逐条调用共享 `decode_sample()`;返回顺序必须和 entries 相同。`clear_many()` 只允许 rank 0 在全局 step 成功后调用。`close_local()` 只关闭本 rank client,不能关闭 Controller。 + +#### `verl_speco/trainer/tq_sample_source.py` + +需要实现: + +```python +@dataclass +class TQLocalBatch: + local_keys: list[str] + local_samples: list[DraftFeatureSample] + global_keys: list[str] | None + +class TQFeatureDataLoader: + def __iter__(self) -> Iterator[TQLocalBatch]: ... + def _select_global_batch(self) -> tuple[list[str], dict[str, ReadyEntry]]: ... + def _broadcast_assignments(self, assignments) -> list[ReadyEntry]: ... + def _handle_eos_and_tail(self) -> bool: ... + def clear_completed_batch(self, global_keys: list[str] | None) -> None: ... +``` + +执行责任必须明确: + +- 所有 rank 创建 loader 并调用 `store.connect()`; +- 只有 rank 0 执行 `_select_global_batch()` 和 `kv_list`; +- rank 0 生成 `list[list[ReadyEntry]]` assignments,使用 `torch.distributed.scatter_object_list` 或 broadcast 分发小型 key/tag; +- 每个 rank 对自己的 local entries 调用 `store.get_many()`,hidden states 从 TQ/Mooncake 直接进入该 rank; +- loader yield `TQLocalBatch`,不能只 yield samples,因为训练后清理还需要 `global_keys`; +- 暂时为空时轮询等待,不能像静态 loader 一样结束;只有 EOS 已出现并且 ready 数据 drain 完才停止; +- `drop_last=True` 时由 rank 0 记录并清理不足 global batch 的尾部 key。 + +#### `verl_speco/trainer/draft_training_loop.py` + +需要新增或调整: + +```python +def build_training_source(config, rank, world_size): ... +def all_ranks_succeeded(local_ok: bool, device) -> bool: ... +async def run_tq_training_loop(trainer, loader, config) -> dict[str, Any]: ... +``` + +`build_training_source()` 根据 `feature_store.type` 分支:静态类型继续创建 `DraftFeatureDataLoader`;`tq` 创建 `TQFeatureDataLoader`,跳过 `feature_store.path` 必填检查,并禁止 `shuffle/repeat`。`run_tq_training_loop()` 对每个 `TQLocalBatch` 调用现有 `prepare_training_batch_from_samples()` 和 `training_step_from_batch()`;所有 rank 通过 collective 得到 global success 后,才让 rank 0 调用 `clear_completed_batch(global_keys)`。异常路径不 clear,finally 只执行 `store.close_local()`。 + +#### `verl_speco/draft_train_launcher.py` + +保留当前 `torch.distributed.run` 启动方式,只增加配置校验和环境透传: + +```python +def validate_tq_launch_config(overrides, launch_config) -> None: ... +def build_child_env(config) -> dict[str, str]: ... +``` + +它不启动 Ray head、不调用 `tq.init()`。职责是确认 `type=tq` 时提供了 Ray address/namespace,且不要求 feature-store path;随后把同一 Ray connection 配置传给每个 torchrun 子进程。每个子 rank 自己建立 Ray/TQ client,不能在 launcher 父进程建一个 client后期待 fork 继承。 + +#### `verl_speco/config/speco_base.yaml` + +增加默认字段: + +```yaml +feature_store: + type: torch_shard + path: null + shuffle: true + repeat: true + tq: + ray_address: null + ray_namespace: speco-drafter + partition_id: speco_drafter_features + run_id: null + schema_version: 1 + poll_interval_seconds: 0.5 + connect_timeout_seconds: 120 + drop_last: true +``` + +当 `type=tq` 时运行期覆盖为 `shuffle=false/repeat=false`。backend 的完整 owner 配置只需要 TQ owner 使用;Producer/Consumer client 从 named Controller 获取它,不各自重新创建 backend。 + +#### Consumer 测试必须覆盖的函数边界 + +- `test_feature_store_factory_builds_tq_without_path()`; +- `test_rank0_filters_and_sorts_ready_entries()`; +- `test_nonzero_rank_never_calls_kv_list()`; +- `test_assignments_are_disjoint_and_global_batch_complete()`; +- `test_each_rank_gets_only_local_keys()`; +- `test_decode_preserves_hidden_states_layout()`; +- `test_clear_only_after_all_ranks_success()`; +- `test_failure_does_not_clear()`; +- `test_eos_drains_ready_then_stops()`; +- `test_client_close_does_not_kill_owner()`。 + +## 6. 两个人怎么分工 + +### 共同先完成 + +1. `drafter_sample_protocol.py`; +2. Ray/TQ connection 配置字段; +3. 一个小型 golden sample; +4. 启动一个测试 Ray head 和 TQ owner,独立进程 A put、独立进程 B list/get/clear 的 smoke test; +5. 验证 client 退出不会 kill owner,只有 owner shutdown 才销毁 Controller。 + +### 开发者 A:Producer/TQ + +负责: + +```text +integration/transferqueue_bridge.py +tq_owner.py +standalone_tq_producer.py +producer/input_reader.py +producer/vllm_feature_client.py +target_feature_replay.py 的公共转换函数 +owner/producer 启动脚本 +Producer/TQ 测试 +``` + +开发者 A 的可交付接口不是“提供一个 TQ 类”,而是: + +```text +bridge:connect_ray_cluster/start_owner/connect_client/put/list/get_many/clear/client_close/owner_close +owner:run_owner/main/signal handler/owner_ready +producer:run_producer/publish_one/统计与 EOS +input reader:iter_input_records/tokenize_record/build_loss_mask +vLLM client:endpoint pool/request_prefill/load/delete +feature conversion:feature_from_vllm_payload +``` + +开发者 B 可以先针对这些接口写 fake transport,不需要等待真实 vLLM 和 Mooncake 联通。 + +### 开发者 B:Consumer/训练 + +负责: + +```text +feature_store.py 的 type=tq 工厂分支 +tq_feature_store.py +tq_sample_source.py / TQFeatureDataLoader +draft_training_loop.py 的 offline + type=tq 分支 +draft_train_launcher.py 配置适配 +speco_base.yaml Consumer 配置 +Consumer 启动脚本 +Consumer/DSpark 测试 +``` + +开发者 B 的可交付接口是: + +```text +feature-store factory:type=tq 分支 +TQFeatureStore:connect/list_ready/get_many/clear_many/read_eos/close_local +TQFeatureDataLoader:select global batch/distribute local keys/get/yield/EOS-tail +training loop:build source/train/global success/clear/final checkpoint +launcher:TQ 配置校验和 torchrun 子进程环境透传 +``` + +### 联调入口 + +建议再提供: + +```text +examples/run_dspark_tq_pipeline_local.sh +``` + +只用于单机联调,顺序启动: + +```text +ray start --head +→ Mooncake metadata/master +→ TQ owner(ray.init + tq.init(full config)) +→ owner_ready +→ vLLM health check +→ Consumer +→ Producer +→ 等 Producer/Consumer 退出 +→ SIGTERM TQ owner(owner 执行 tq.close) +→ ray stop +→ 停止 Mooncake +``` + +最小联调:Producer 发布 8 条样本,2 个 Consumer ranks、每 rank batch size 2,完成 2 个 optimizer steps,8 个 sample keys 被清理,EOS 后保存 final checkpoint。 + +## 7. 第一版验收标准 + +1. TQ owner、Producer 和所有 Consumer ranks 日志显示相同 Ray address/namespace、TQ Controller、partition 和 run ID。 +2. 在同一 Ray 集群中,独立 owner 创建 TQ 后,独立进程 A put,独立进程 B 能 list/get/clear。 +3. Producer 对多个 vLLM endpoints 并发请求,不串行访问。 +4. 一个输入 record 只生成一个 sample key 和一个 payload。 +5. `kv_list` 只拿 tag;`kv_batch_get` 才拿 hidden states。 +6. Consumer 各 rank 读取不同 local keys,不通过 rank 0 搬运 Tensor。 +7. `decode_sample` 能拒绝模型、layer、layout、dtype 或 shape 不匹配的数据。 +8. 所有 rank 训练成功后才 clear 当前 global batch。 +9. Producer 先完成时,Consumer 能 drain 后再退出。 +10. 不产生长期 hidden-state feature store。 +11. Producer 或任一 Consumer rank 退出不会销毁 TQ Controller;只有 owner 调用全局 `tq.close()`。 +12. Ray object store 中不承载 hidden-state payload,训练 tensor 通过 TQ/Mooncake 路径读取。 + +## 8. 后续建议:第一版跑通后再做 + +以下内容不进入第一版开发: + +- Producer HTTP/TQ 复杂重试和 endpoint 熔断; +- Producer 发布 journal,避免重启后重复生成已 clear 样本; +- 自动重启;第一版失败后人工停止整条 pipeline并使用新 `run_id` 重跑; +- Consumer 从最新 checkpoint 自动恢复; +- checkpoint 成功后再 clear 的严格提交窗口; +- TQ owner/storage 整体丢失后的数据重建; +- 多个独立 Consumer 竞争同一 partition; +- lease、ack、超时回收和 exactly-once; +- 动态扩缩容; +- vLLM server 直接写 TQ。 + +第一版先保证:同一个 TQ 能连通、Producer 能并发生产、Consumer 能正确取数训练、每步成功后能及时清理。 diff --git a/docs/transferqueue_integration_plan.md b/docs/transferqueue_integration_plan.md new file mode 100644 index 00000000..cac16ac1 --- /dev/null +++ b/docs/transferqueue_integration_plan.md @@ -0,0 +1,145 @@ +# verl-SpeCo TransferQueue 落地方案 + +Last updated: 08/21/2026 + +> 目标:在**不修改上游 verl**的前提下,把 SpeCo online 训练里的逐样本特征流 +> 从「`SpecoRayPPOTrainer` driver 中转 + Ray object store」改为「TransferQueue +> 直传」,干掉 driver 这个数据瓶颈,并解锁流式消费与跨副本负载均衡。 +> +> 约束:仅 hook,与 SpeCo 现有 hook 模式一致;TQ 作为独立库使用,**不复用** verl +> 的 `main_ppo_sync` TQ 集成。 + +--- + +## 0. 现状:SpeCo online 特征流的 controller 瓶颈 + +SpeCo online 路径以 `SpecoRayPPOTrainer` 为 hub,所有跨进程大张量都被 driver 串行 +中转,介质是 Ray object store(`ray.put`/`ray.get`/`parallel_put`)。这与 verl 引入 +TQ 想解决的痛点 1:1 对应,只是 verl 干掉的是 `RayPPOTrainer`,我们要干掉的是 +SpeCo 在它之上加的 drafter 管线中转。 + +| # | 流向 | 当前机制 | 是否逐样本 | hook 位置(SpeCo 侧) | +|---|---|---|---|---| +| **a1** | target hidden states(SGLang 采集)-> drafter | `drafter_sample` 塞进 `DataProto.non_tensor_batch` -> driver pop/bucket -> `parallel_put` -> drafter `ray.get` | ✅ | `speco_ray_trainer.py` `generate_sequences_with_speco`;`sglang_adapter.py` `pop_drafter_samples`/`bucket_drafter_samples_by_replica`;`sglang_runtime.py` 组装 `drafter_sample` | +| **a2** | target hidden states(old-logprob hook)-> drafter | actor 前向 hook 截行 -> `ray.put` chunk -> driver 重打包 -> 分发 | ✅ | `oldlogprob_runtime.py` `_install_oldlogprob_hidden_hooks`/`_put_oldlogprob_hidden_refs`;`speco_ray_trainer.py` `_speco_collect_oldlogprob_features` | +| **b2** | target top-logprobs -> drafter(`use_logits=true`) | 随 a1 同一 side-channel | ✅ | `sglang_runtime.py` `target_logprobs`/`hidden_raw_target_logprobs` | +| **d** | rollout tokens -> drafter 训练集 | 随 a1 同一 side-channel(online)/`torch.save` 分片(offline) | ✅ | `sglang_runtime.py`;`speco_worker.py` `_store_rollout_sample` | +| b1 | target **lm_head 权重**(行)-> drafter `TargetHead` | ONE_TO_ALL Ray 分发 | ❌ 参数广播 | `rollout_publish.py` `export_actor_lm_head_weight`/`get_actor_lm_head_weight`;`speco_ray_trainer.py` `_speco_sync_target_lm_head_weight` | +| c | drafter 权重 -> rollout 引擎 | `ray.put` -> driver -> actor;vLLM 末段 ZMQ+SHM,SGLang 进程内 | ❌ 参数广播 | `speco_worker.py` `maybe_publish`;`rollout_publish.py` `update_draft_weights`;`vllm_runtime.py` `BucketedWeightSender` | + +**关键事实**:hidden states 跨进程前一律 CPU 物化(`oldlogprob_runtime.py`、 +`sglang_runtime.py`、`feature_store.py` 均 `.cpu()`),a1 路径下 driver 进程的 +host memory 会真正承载整批 hidden states 并做一次 Ray store 往返。这正是 TQ 要 +消除的往返。 + +--- + +## 1. 为什么不把"替换 feature_store"作为第一刀 + +`TorchShardFeatureStore`(`feature_store.py`)是 `torch.save` 分片 + JSONL manifest, +**只服务于 `collect_only`/`offline`**,不参与 online 热路径。替换它能统一离线存储 +抽象、换更快的分布式后端,但**不解决 controller 瓶颈**,性能收益有限。降级为 +可选尾项(见 §5 P3)。 + +--- + +## 2. 目标方案:TQ 直传逐样本特征流(a1 / a2 / b2 / d) + +### 2.1 角色映射 + +| TQ 角色 | SpeCo 对应 | +|---|---| +| Producer(写) | rollout worker(SGLang 路径,a1/b2/d)/ actor worker(old-logprob 路径,a2)——均在 SpeCo 既有 hook 内 | +| Consumer(读) | drafter worker `collect_rollout_features`(SpeCo 侧) | +| TransferQueueController(control plane) | SpeCo launcher 启动一个 Ray actor;drafter 经 `Sampler`/`StreamingDataLoader` 拉取 | +| Storage backend | `SimpleStorage`(ZMQ,跨节点 CPU 内存);进阶可切 `MooncakeStore`(RDMA,GPU-DRAM) | + +### 2.2 partition / key / 字段设计 + +- `partition_id`:`speco_train`(验证集用 `speco_val`)。 +- `key`:`{uid}_{session_id}_{index}`,与 verl TQ 一致;`uid` SpeCo 已有。 +- `tags`:`global_steps`、`source`∈{`rollout`,`oldlogprob`}、`replica_rank`/`owner_rank`、`status`、`prompt_len`/`response_len`/`seq_len`。ReplayBuffer/负载均衡按 tag 匹配。 +- `fields`(列):`input_ids`、`loss_mask`、`position_ids`、`hidden_states`、`last_hidden_states`/`target`、`target_logprobs`、`hidden_positions`、`prompts`、`responses`。与 `DraftFeatureSample`(`feature_store.py`)字段对齐,便于 online/offline 复用。 + +### 2.3 数据流(目标) + +``` +rollout/actor worker (SpeCo hook) + │ 生成/截取 hidden states 后,就地 tq.kv_batch_put(samples) + ▼ +TransferQueue (SimpleStorage, 跨节点 CPU 内存;可选 MooncakeStore RDMA) + │ control plane 按 sample 粒度追踪 ready 状态,Sampler 跨 drafter 副本均衡 + ▼ +drafter worker + │ tq.kv_batch_get / StreamingDataLoader 消费 → 喂入既有 DataBuffer / collect_online_data + ▼ +drafter 训练 (不变) +``` + +driver 只下发触发与轻量 key/meta,**不再承载 hidden states**。 + +--- + +## 3. 落地改动点(全部在 SpeCo 侧,hook-only) + +### 3.1 启动与配置 +- `draft_train_launcher.py` / `main.py`:`tq.init(config.transfer_queue)`;起 `TransferQueueController.remote(Sampler)`。 +- `config/speco_base.yaml`:新增 `drafter.transfer_queue` 块(backend、partition、enable 开关)。参考 verl `ppo_trainer.yaml` 的 `transfer_queue:` 结构,但**独立配置**,不复用 verl 的。 + +### 3.2 Producer 侧 +- **a1/b2/d(SGLang)**:`sglang_runtime.py` 组装 `drafter_sample` 处(~1594-1648),增加 `tq.kv_batch_put`;返回给 driver 的 `drafter_sample` 只保留 key/meta(或整段不再走 DataProto side-channel,driver 仅触发)。 +- **a2(old-logprob)**:`oldlogprob_runtime.py` `_put_oldlogprob_hidden_refs`(~216),把 `ray.put(hidden_chunk)` 换成 `tq.kv_batch_put`;`OLD_LOGPROB_HIDDEN_CHUNK_REFS_KEY` 改为 TQ key 列表。 + +### 3.3 Consumer 侧 +- `speco_worker.py` `collect_rollout_features`(~665):把 `_resolve_ray_object_ref`/`_resolve_hidden_state_chunks`(`ray.get`)换成 `tq.kv_batch_get`;`_dispatch_nd_compute`(~159)的 `parallel_put` 退化为只传 key(或 drafter 直接从 TQ Sampler 拉,driver 不参与分发)。 +- drafter 内部 `DataBuffer`/`collect_online_data`(`base_trainer.py`)保持不变,只是数据来源由 `ray.get` 改为 TQ get。 + +### 3.4 Driver 侧 +- `speco_ray_trainer.py`:`_speco_collect_rollout_features_rpc`/`speco_collect_rollout_features`(~351)、`_speco_collect_oldlogprob_features`(~1114)不再搬数据,只做触发/传 key;`bucket_drafter_samples_by_replica` 可由 TQ `Sampler` 替代(逐步迁移,先保留作回退)。 + +### 3.5 不改动 +- **b1(lm_head 权重)、c(drafter 权重)**:保持现状。与 verl 上游一致(权重不走 TQ),且 c 的 vLLM 末段已有专用 ZMQ+SHM 通道。 +- verl 本体:零改动。 + +--- + +## 4. 收益与边界(诚实评估) + +### 收益 +1. **去掉 driver 对 hidden states 的 host-memory 中转 + Ray store 往返**:producer 直存 TQ,consumer 直取,driver 不再承载整批特征。 +2. **流式消费**:drafter 在样本 ready 时即可消费,不必等整批 `generate_sequences` 返回,采集与训练可重叠。 +3. **跨 drafter 副本负载均衡**:TQ `Sampler`/`RankAwareSampler` 替代手写 `bucket_drafter_samples_by_replica`/`owner_rank` 分配。 +4. **(若采纳 P3)统一 online/collect_only/offline 存储**:同一 TQ partition,`collect_only` 写、`offline` 读,消掉 on-disk 分片层。 + +### 边界 / 不解决的事 +- 只优化**特征采集**这一子阶段,**不加速** rollout 本身、actor update、reward;e2e 增益取决于该子阶段在 step 中的占比。 SpeCo README 的 20% rollout / 11% e2e 提升来自 acceptance length,与本方案是不同机制,不要混为一谈。 +- **权重同步(b1/c)不放进 TQ**,与 verl 上游保持一致。 +- hidden states 跨进程前**仍需 CPU 物化**(现状如此);要避免物化需切 `MooncakeStore` RDMA,属进阶项。 +- 引入 TQ 依赖与一个 control-plane Ray actor,增加少量运维面。 + +### 风险 +- TQ 与 SpeCo 现有 `owner_rank`/`replica_rank` 路由语义需对齐(Sampler 要复刻「按 owner 分桶」语义,否则样本会错配 drafter 副本)。 +- old-logprob 的 chunk 拆分(`hidden_states_ref_chunks`)映射到 TQ 列式存储时,需保证 chunk meta 与 key 的一致性。 +- 回退路径:保留 `enable_transfer_queue=False` 时走原 Ray 路径,渐进切换。 + +--- + +## 5. 分阶段实施 + +| 阶段 | 范围 | 产出 | +|---|---|---| +| **P0** | a1(SGLang hidden states)走 TQ 直传;drafter `kv_batch_get` 消费;driver 仅触发 | 验证 controller-bypass 闭环 + 正确性 | +| **P1** | a2(old-logprob hidden states)走 TQ;chunk 拆分映射 TQ 列 | 覆盖第二条采集路径 | +| **P2** | b2(top-logprobs)+ d(tokens)随 a1 同 partition 传输;Sampler 替代手写 bucket | 完整特征流 + 跨副本均衡 | +| **P3(可选)** | `TorchShardFeatureStore` → TQ partition,统一 online/collect_only/offline | 离线工作流统一 | + +每个阶段保留 `enable_transfer_queue` 开关与原 Ray 路径回退。 + +--- + +## 6. 待确认决策 + +1. **TQ backend**:`SimpleStorage`(CPU 内存,默认)起步,还是直接上 `MooncakeStore`(RDMA,省 CPU 物化)?后者依赖 RDMA 网络,建议 P0 用 SimpleStorage。 +2. **drafter 消费模式**:`kv_batch_get`(主动拉,改动小)还是 `StreamingDataLoader`(全自动流式,改动大、收益高)?建议 P0 用前者,P2 再考虑后者。 +3. **driver 角色**:P0 先保留 driver 传 key(最小改动),还是直接让 drafter 从 TQ Sampler 自取(driver 彻底退出数据路径)?前者风险低,建议 P0 用前者。 +4. **是否做 P3**:离线统一是否在本次范围内,还是单独立项。 diff --git a/docs/vllm_direct_hidden_state_cotrain_plan.md b/docs/vllm_direct_hidden_state_cotrain_plan.md new file mode 100644 index 00000000..6c8731f4 --- /dev/null +++ b/docs/vllm_direct_hidden_state_cotrain_plan.md @@ -0,0 +1,666 @@ +# Co-train 复用 Producer 直接从 vLLM 获取 Hidden State 的实施方案 + +## 1. 目标与边界 + +本文方案面向 **RL 与草稿模型共同训练(co-train)**,目标是把 target hidden state 的来源从 actor 的 old-logprob 前向切换为 vLLM: + +```text +原方案:vLLM 生成 response + -> actor 为 PPO 计算 old_log_probs + -> 在 actor forward 内通过 hook/output_hidden_states 抓 hidden state + -> 草稿模型训练 + +新方案:vLLM 生成 response + -> 使用 prompt + response 再向 hidden-state vLLM 发起 prefill 请求 + -> vLLM 返回 hidden-state 文件位置,客户端读取并归一化 + -> 现有 scheduler -> SpecoWorker -> drafter buffer -> 草稿模型训练 +``` + +需要特别区分两个动作: + +- actor 的 `compute_log_prob()` **仍然保留**,因为 PPO 更新 actor 需要 `old_log_probs`。 +- 删除的是 old-logprob 前向中专为草稿训练增加的 hidden-state 捕获、拼接、CPU copy 和 Ray put 逻辑。 + +第一版不修改独立训练的 TQ producer/consumer 行为,不让 co-train 使用独立训练的文件读取、EOS、TQ owner 和离线 dataloader。复用的是 producer 中已经验证过的 **vLLM 请求、并发、重试、文件读取、token 对齐和样本归一化能力**。 + +## 2. 当前 co-train 全流程 + +### 2.1 运行时角色 + +| 角色 | 所在位置 | 当前职责 | +|---|---|---| +| `RayPPOTrainer` driver | `verl_speco/trainer/speco_ray_trainer.py` | 驱动 rollout、old-logprob、actor update,并通过 drafter scheduler 决定采样和训练时机 | +| vLLM rollout workers | verl rollout worker group | 根据 prompt 生成 response;当前 vLLM 路径不直接给 drafter hidden state | +| actor workers | `actor_rollout_wg` | 计算 PPO 所需 old log-prob;当前还承担 hidden-state 捕获 | +| drafter workers | `drafter_wg` 中的 `SpecoWorker` | 接收按 owner 分桶的样本,写入在线 buffer,执行 drafter train/publish/checkpoint | +| `DrafterScheduler` | driver 进程内普通 Python 对象 | 决定本 step 是否采集、是否训练,并组织 collection 的 stage/commit/finalize | + +这里的 scheduler 不是另一个服务,也不读取数据;它只负责控制顺序和事务状态。 + +### 2.2 rollout 生成的数据 + +rollout 后,driver 持有 `DataProto batch`。与本方案直接相关的 tensor 通常为: + +```python +batch.batch["prompts"] # [B, Pmax],左 padding +batch.batch["responses"] # [B, Rmax],右 padding +batch.batch["attention_mask"] # [B, Pmax + Rmax] +batch.batch["response_mask"] # [B, Rmax],若存在则表示有效 response token +``` + +单个有效样本会还原为: + +```python +prompt_ids: Tensor[P] +response_ids: Tensor[R] +input_ids = torch.cat([prompt_ids, response_ids]) # Tensor[P + R] +``` + +padding token 不能发给 hidden-state vLLM;必须通过 mask 去掉。 + +### 2.3 当前 old-logprob hidden-state 路径 + +入口位于 `SpecoRayPPOTrainer._speco_online_fit_hooks()` 安装的 `compute_old_log_prob_with_speco()` 包装函数: + +1. `_speco_plan_drafter_collection(OLD_LOGPROB)` 调 scheduler,决定当前 `global_step` 是否采集。 +2. `_speco_build_oldlogprob_collect_plan(batch)` 选择样本、hidden positions 和 owner rank。 +3. driver 把以下控制数据写进 `batch_td`: + + ```python + OLD_LOGPROB_COLLECT_MASK_KEY + OLD_LOGPROB_HIDDEN_POSITIONS_KEY + OLD_LOGPROB_HIDDEN_POSITION_MASK_KEY + OLD_LOGPROB_OWNER_RANK_KEY + OLD_LOGPROB_AUX_LAYER_IDS_KEY + OLD_LOGPROB_HIDDEN_CAPTURE_IMPL_KEY + OLD_LOGPROB_HIDDEN_LAYOUT_KEY + ``` + +4. `actor_rollout_wg.compute_log_prob(batch_td)` 远程执行 actor 前向。 +5. `verl_speco/integration/oldlogprob_runtime.py` 根据 `forward_hook` 或 `output_hidden_states` 捕获指定层,并把结果作为 tensor、Ray object ref 或分块 ref 返回。 +6. driver 的 `_speco_collect_oldlogprob_features()` 把 actor 输出还原为逐样本字典: + + ```python + sample = { + "input_ids": Tensor[1, P + R], + "prompts": Tensor[1, P], + "responses": Tensor[1, R], + "hidden_positions": Tensor[1, Hrows], + "hidden_states": Tensor[1, Hrows, Hdim], # 或 *_ref / *_ref_chunks + "hidden_states_layout": "dflash_aux" | "eagle3_aux_plus_last", + "hidden_position_start": int, + "hidden_position_end": int, + "global_step": int, + "replica_rank": int, + } + ``` + +7. `OldLogProbCollectionAdapter.prepare_payload()` 根据显式 `owners` 把样本分到 drafter owner buckets。 +8. scheduler 执行 collection transaction,Ray RPC 参数本质上是每个 owner 对应的 `list[dict]`。 +9. `SpecoWorker._commit_rollout_features(collection_id, samples)` 解析 hidden tensor/ref,调用 `_store_rollout_sample()`。 +10. `_store_rollout_sample()` 调 `DrafterBaseTrainer.collect_online_data()`,把 CPU 数据写入当前 step 或跨 step buffer。 +11. `update_actor_with_speco()` 调 `_speco_on_before_actor_update()`;scheduler 根据已收集数据产生 training plan,然后执行 actor update 和 drafter training。 + +因此,现有 worker、buffer 和训练后半段并不关心 hidden state 是由 actor 还是 vLLM 产生。需要替换的主要是第 2~6 步的数据生产方式。 + +## 3. 当前 standalone producer 的详细流程 + +### 3.1 哪些部分可以复用 + +`verl_speco/standalone_tq_producer.py` 当前是一个三段式异步流水线: + +```text +read_inputs + -> request_queue + -> N 个 request_worker + -> publish_queue + -> publish_results + -> TQ +``` + +其中只有最后的 TQ publish 和最前面的文件 reader 是 standalone 专属。中间部分已经包含 co-train 需要的核心能力。 + +#### `VllmFeatureClientPool` + +文件:`verl_speco/producer/vllm_feature_client.py` + +职责: + +- 解析多个 `VllmEndpoint`; +- 维护全局并发 semaphore 和每 endpoint semaphore; +- 优先选择当前 inflight 较少的 endpoint; +- 通过 OpenAI completions 接口发送 token IDs; +- 对连接错误、read error、超时等执行指数退避重试; +- 从响应的 `kv_transfer_params.hidden_states_path` 取得 safetensors 路径; +- 等待文件完成,读取 hidden state、token IDs 和相关字段; +- 读取完成后删除临时 safetensors 与 lock 文件。 + +请求不是让 vLLM 再生成一段文本,而是一次 prefill 请求: + +```python +await client_pool.request_prefill( + prompt_token_ids=request.vllm_prompt_token_ids, + sample_id=request.sample_id, +) +``` + +返回的 `RawVllmFeature` 仍是 vLLM 原始坐标系下的数据,例如: + +```python +RawVllmFeature( + token_ids=Tensor[Tprefill], + hidden_states=Tensor[Tprefill, L, D], + hidden_position_start=..., + hidden_position_end=..., + ..., +) +``` + +其中 `L` 是导出的 target layer 数量,`D` 是 target hidden size。 + +#### `prepare_generated_prefill_request` + +文件:`verl_speco/producer/input_reader.py` + +standalone 遇到仅有 prompt 的数据时,先生成 response,再调用该函数把 prompt 和生成结果组成训练请求。核心规则是: + +```python +full_ids = prompt_ids + response_ids +vllm_prompt_token_ids = full_ids[:-1] +``` + +去掉最后一个 token 的原因是:位置 `i` 的 target hidden state 用于预测后续 token,最后一个 token 后面没有本样本内的监督 token。co-train 已经拥有 rollout response,所以只需要直接执行这一步,不需要再次生成 response。 + +当前 `TokenizedRequest` 包含: + +```python +TokenizedRequest( + sequence_no: int, + sample_id: str, + input_ids: list[int], # prompt + response + loss_mask: list[float], # prompt 为 0,有效 response 为 1 + position_ids: list[int], + feature_positions: list[int], # 选中的 target hidden 绝对位置 + draft_position_ids: list[int], + source_metadata: dict, + vllm_prompt_token_ids: list[int], # 发给 vLLM 的 full_ids[:-1] +) +``` + +#### `feature_from_vllm_payload` + +文件:`verl_speco/trainer/target_feature_replay.py` + +该函数把 `RawVllmFeature + TokenizedRequest + FeatureContract` 转成算法训练侧统一使用的 `DraftFeatureSample`。它负责: + +- 检查 vLLM 返回的 token IDs 是否与请求一致; +- 检查 hidden rows、层数、hidden size 和 layout; +- 按 `feature_positions` 选择训练窗口; +- 对齐 `input_ids`、`loss_mask`、positions 与 hidden state; +- hidden state 不完整或位置对不上时拒绝该样本,不把错误样本交给训练。 + +典型结果: + +```python +DraftFeatureSample( + input_ids=Tensor[T], + loss_mask=Tensor[T], + hidden_states=Tensor[Hrows, L * D], + target_logprobs=None, + position_ids=Tensor[T], + feature_positions=Tensor[Hrows], + draft_position_ids=Tensor[Hrows], + metadata={...}, +) +``` + +这里的协议和算法处理应继续由已有 `DraftFeatureSample`、backend 和 contract 决定,不能在新 co-train 组件里再次硬编码 DSpark。 + +### 3.2 哪些部分不能直接搬入 co-train + +以下 standalone 逻辑不能原样调用: + +- 从 JSONL 循环读 epoch;co-train 的输入来自当前 rollout `DataProto`。 +- prompt-only 时调用生成接口;co-train 的 response 已经生成。 +- `sequence_no/run_id/tag/EOS/max_pending_samples`;这些用于 TQ 流式生产消费,不属于单个 RL step。 +- `publish_results()` 和 TQ clear;co-train 已有 scheduler collection transaction 和 worker buffer。 +- standalone owner/consumer 生命周期;co-train 由 Ray trainer 和 worker group 管理。 + +正确的复用方式是抽取“给定 tokenized rollout sample,异步取得并归一化 hidden state”的核心,而不是在 co-train 内启动一个 `standalone_tq_producer` 进程。 + +## 4. 建议的新数据流 + +### 4.1 完整顺序 + +```text +1. rollout vLLM 生成 response +2. driver 得到 DataProto(prompts, responses, masks) +3. scheduler 判断本 step 是否需要采集 VLLM_PREFILL +4. driver 按采样计划选择样本、去 padding、构造 TokenizedRequest +5. CotrainVllmFeatureProducer.submit_batch() 提交并发 prefill +6. hidden-state vLLM endpoint 执行 prompt+response[:-1] prefill +7. client 读取 safetensors,执行 token/shape/position 对齐 +8. 得到 DraftFeatureSample;失败或不完整样本在此处过滤 +9. 将 DraftFeatureSample 转为现有 SpecoWorker collection sample +10. VllmPrefillCollectionAdapter 按 replica owner 分 buckets +11. scheduler stage -> Ray commit RPC -> finalize +12. SpecoWorker._store_rollout_sample() -> collect_online_data() -> buffer +13. scheduler 产生 training plan +14. actor update 与 drafter training 按现有顺序执行 +``` + +步骤 5 提交后不应立刻阻塞等待。driver 可以继续 reward、reference log-prob、advantage、actor old-logprob 等工作;在 drafter collection 必须完成的边界再 `await/result()`。这样 vLLM prefill 与 RL 侧计算重叠。 + +### 4.2 vLLM 请求的具体 token 对齐 + +给定一个 rollout 样本: + +```python +prompt_ids = prompts[i][prompt_mask] # [P] +response_ids = responses[i][response_mask] # [R] +full_ids = cat(prompt_ids, response_ids) # [P + R] +prefill_ids = full_ids[:-1] # [P + R - 1] +``` + +构造: + +```python +loss_mask = zeros(P + R) +loss_mask[P:P + R] = 1 +``` + +然后再应用现有 collection plan 的窗口限制。必须保证: + +```text +返回 token_ids == prefill_ids +hidden rows 能覆盖 feature_positions +feature_positions 非空 +选中区域对应的 loss_mask 中存在有效训练 token +``` + +只含 prompt、有效 response 长度为 0、hidden rows 为 0、token 不一致或窗口为空的样本,都在 producer 转换阶段丢弃,不进入 scheduler payload。这样 producer 的“发布成功数”和 consumer 的“可接收数”天然一致,不会把无效条目带入 collection transaction。 + +### 4.3 scheduler 到 worker 的样本格式 + +建议保留 worker 当前已经支持的 collection sample 外形,不大改训练后半段: + +```python +worker_sample = { + "input_ids": Tensor[1, T], + "prompts": Tensor[1, P], + "responses": Tensor[1, R], + "hidden_positions": Tensor[1, Hrows], + "hidden_states": Tensor[1, Hrows, HiddenWidth], + "hidden_states_layout": str, + "hidden_position_start": int, + "hidden_position_end": int, + "global_step": int, + "replica_rank": int, +} +``` + +`HiddenWidth` 取决于现有 backend/layout。例如多个层已经按最后一维拼接时为 `L * D`。该转换必须调用 `DraftFeatureSample` 已有字段和 metadata,不在 adapter 中按算法猜测。 + +`replica_rank` 不是 vLLM 返回的数据,而是 driver 根据现有 drafter owner 路由计划为样本分配的控制字段。scheduler 只用它决定该样本发给哪个 `SpecoWorker` owner。 + +## 5. 代码修改方案 + +### 5.1 新增 co-train producer 核心 + +新增:`verl_speco/producer/cotrain_vllm_feature_producer.py` + +建议接口: + +```python +@dataclass +class CotrainFeatureRequest: + batch_index: int + owner_rank: int + request: TokenizedRequest + prompt_ids: torch.Tensor + response_ids: torch.Tensor + + +@dataclass +class CotrainFeatureResult: + batch_index: int + owner_rank: int + sample: DraftFeatureSample + + +class CotrainVllmFeatureProducer: + def __init__(self, config, *, contract: FeatureContract): ... + + def submit_batch( + self, + requests: list[CotrainFeatureRequest], + ) -> Future[list[CotrainFeatureResult]]: ... + + async def _produce_one( + self, + request: CotrainFeatureRequest, + ) -> CotrainFeatureResult | None: ... + + def close(self) -> None: ... +``` + +内部直接复用: + +```python +raw = await self.client_pool.request_prefill(...) +sample = feature_from_vllm_payload(raw, request.request, self.contract) +``` + +组件应持有一个长期存在的 `VllmFeatureClientPool`,不能每 step 重建 HTTP client、semaphore 和线程池。由于 PPO driver 主流程通常是同步代码,第一版可让组件内部持有一个后台 asyncio event loop thread,`submit_batch()` 返回 `concurrent.futures.Future`。训练结束时统一 `close()`,取消未完成任务并关闭 HTTP client。 + +### 5.2 增加从 rollout tensor 构造请求的函数 + +修改:`verl_speco/producer/input_reader.py` + +新增纯函数,复用现有长度截断、feature window、position 和 loss-mask 规则: + +```python +def build_rollout_prefill_request( + *, + sample_id: str, + sequence_no: int, + prompt_ids: Sequence[int], + response_ids: Sequence[int], + producer_cfg, + source_metadata: dict, +) -> TokenizedRequest: + ... +``` + +它不接收文本、不调用 tokenizer、不生成 response,只做: + +1. 拼接有效 prompt/response token; +2. 按已有 `max_feature_length` 等规则截取 response; +3. 建 loss mask、positions; +4. 设置 `vllm_prompt_token_ids=full_ids[:-1]`。 + +必须把 standalone 与 co-train 的公共构造逻辑下沉到同一个私有 helper,避免两个路径以后出现 off-by-one 或截断规则差异。 + +### 5.3 扩展 scheduler 的 collection source + +修改: + +- `verl_speco/trainer/scheduler/schedule_types.py` +- `verl_speco/trainer/scheduler/collection_adapter.py` +- `verl_speco/trainer/scheduler/drafter_scheduler.py` + +新增: + +```python +class DrafterCollectionSource(str, Enum): + SGLANG = "sglang" + OLD_LOGPROB = "oldlogprob" + VLLM_PREFILL = "vllm_prefill" +``` + +新增 `VllmPrefillCollectionAdapter`。它只负责: + +- 校验每个 sample 有 `replica_rank`; +- 使用 `_build_payload()` 按 owner 分桶; +- 设置 `CollectionPayload.source=VLLM_PREFILL`。 + +它不负责请求 vLLM、不解码 hidden state、不实现算法逻辑。 + +同时更新 collection source 的稳定排序值、adapter registry 和 metrics source label。 + +### 5.4 在 `SpecoRayPPOTrainer` 接入异步生产 + +修改:`verl_speco/trainer/speco_ray_trainer.py` + +新增或调整以下职责: + +```python +def _speco_vllm_prefill_collection_requested(self) -> bool: ... +def _speco_vllm_prefill_collection_enabled(self) -> bool: ... +def _speco_get_cotrain_vllm_producer(self) -> CotrainVllmFeatureProducer: ... +def _speco_build_vllm_prefill_requests(self, batch, collection_plan): ... +def _speco_submit_vllm_prefill_collection(self, batch): ... +def _speco_finish_vllm_prefill_collection(self) -> int: ... +def _speco_close_vllm_prefill_producer(self) -> None: ... +``` + +接入点建议如下: + +1. rollout 返回 `gen_batch_output` 并合并成训练 batch 后,调用 `_speco_submit_vllm_prefill_collection(batch)`。 +2. 提交函数先调用 scheduler 的 `plan_collection(VLLM_PREFILL)`;未命中 interval 时不发 HTTP 请求。 +3. 继续执行 reward、ref、old-logprob 和 advantage。 +4. 在 `_speco_on_before_actor_update()` 生成 training plan 之前调用 `_speco_finish_vllm_prefill_collection()`: + - 等待 Future; + - 过滤失败样本; + - 转成 worker sample; + - adapter 分桶; + - `_speco_execute_collection()`。 +5. 原 `compute_old_log_prob_with_speco()` 在该模式下走普通 `original_compute_old_log_prob()`,不再注入 hidden capture keys。 +6. fit 的 `finally` 中关闭 producer。 + +这里“等待点必须在 training plan 之前”是必要条件。否则 scheduler 看到的 buffer version 仍是旧值,本 step 可能错误判断没有可训练样本。 + +### 5.5 worker 和训练侧尽量不改 + +`verl_speco/workers/speco_worker.py` 的以下路径可以直接复用: + +```text +collect_rollout_features / collection transaction RPC + -> _commit_rollout_features + -> _store_rollout_sample + -> DrafterBaseTrainer.collect_online_data +``` + +第一版只在必要时增加一个从 `DraftFeatureSample` 转现有 sample dict 的小 helper;不要新增另一套 buffer,也不要让 worker 连接 standalone TQ。 + +如果 `DraftFeatureSample.to_training_item()` 与 `collect_online_data()` 的 metadata 表达存在差异,应在一个公共转换函数中补齐,而不是在 driver、adapter、worker 分别写一套字段映射。 + +### 5.6 禁用 old-logprob hidden capture,但保留 PPO old-logprob + +修改配置判定和 hook 分支: + +```yaml +collect_hidden_states_from_sgl: false +collect_hidden_states_from_old_logprob: false +collect_hidden_states_from_vllm: true +``` + +当 `collect_hidden_states_from_vllm=true` 时: + +- 不设置 `OLD_LOGPROB_HIDDEN_*` keys; +- 不安装/启用 `oldlogprob_runtime` hidden hooks; +- 不调用 `_speco_collect_oldlogprob_features()`; +- 仍执行标准 `actor_rollout_wg.compute_log_prob()`,得到 PPO 的 log-prob 和 entropy。 + +三种来源第一版必须互斥: + +```python +sum([ + collect_hidden_states_from_sgl, + collect_hidden_states_from_old_logprob, + collect_hidden_states_from_vllm, +]) <= 1 +``` + +## 6. 配置建议 + +修改:`verl_speco/config/actor/actor.yaml` 或本仓实际承载 drafter training 默认值的配置文件,并在示例脚本暴露关键参数。 + +建议结构: + +```yaml +actor_rollout_ref: + rollout: + drafter: + training: + mode: online + collect_hidden_states_from_sgl: false + collect_hidden_states_from_old_logprob: false + collect_hidden_states_from_vllm: true + + vllm_feature_source: + endpoints: + - http://127.0.0.1:8000/v1 + - http://127.0.0.1:8001/v1 + model: /path/to/target-model + max_inflight_requests: 128 + per_endpoint_concurrency: 64 + request_timeout_seconds: 600 + max_retries: 3 + retry_base_delay_seconds: 1.0 + max_sequence_length: 8192 +``` + +`target_layer_ids`、`max_feature_length`、hidden layout、算法类型等应继续读取现有 drafter 配置,不在 `vllm_feature_source` 重复定义。`FeatureContract` 也从同一份运行配置构建,从而保证 vLLM 导出层和 trainer 预期一致。 + +### vLLM 服务要求 + +当前 `VllmFeatureClientPool` 使用 OpenAI HTTP endpoint 和 `kv_transfer_params.hidden_states_path`。因此第一版要求: + +- co-train 可访问一个或多个已启动的 hidden-state vLLM 服务; +- 服务加载的 target model/tokenizer 与 rollout/actor 使用的版本一致; +- `extract_hidden_states` 的 layer IDs 与 drafter contract 一致; +- vLLM hidden 文件目录对 driver 可见; +- prefix caching 对 hidden-state 导出必须关闭或已验证能返回所有所需 rows。 + +verl 内部 rollout vLLM worker 不一定天然暴露当前 client 所需的 OpenAI 地址和共享 hidden 文件路径。第一版建议使用独立启动的 hidden-state vLLM endpoints。后续若 rollout 服务能够暴露同等接口,再把 endpoint discovery 接入 worker group,producer 核心无需改变。 + +## 7. 是否在 co-train 中使用 TQ + +### 7.1 第一版建议:不使用 TQ + +第一版直接把 CPU hidden tensor 放入现有 scheduler payload,必要时沿用 Ray object ref/chunk ref 机制。理由: + +- co-train 已经有 scheduler collection transaction、owner 路由和 worker buffer; +- 当前 standalone TQ 的 run ID、pending、EOS、clear 语义针对跨进程无限流,不适合直接套在单个 RL step 上; +- 少改 worker 和生命周期,能先验证 vLLM hidden 与 actor hidden 的数值/训练等价性。 + +此时的控制面和数据面是: + +```text +控制面:driver -> scheduler -> Ray RPC(collection_id, owner bucket) +数据面:CPU tensor 随 Ray 参数,或先 ray.put 后传 ObjectRef +``` + +### 7.2 第二阶段可选:TQ 承载大 tensor + +若 Ray object store 压力明显,再让 producer 将 `DraftFeatureSample` 编码后写 TQ,而 scheduler sample 只传: + +```python +{ + "feature_key": str, + "replica_rank": int, + "global_step": int, +} +``` + +worker commit 时按 key `get + decode`,commit 成功后 clear,rollback 时保留或清理。这需要定义 co-train 专属的 step-scoped key 和事务清理规则,不能复用 standalone EOS。该阶段会增加失败恢复复杂度,不建议和第一版一起提交。 + +## 8. 错误处理与一致性 + +### 8.1 单样本错误 + +以下错误在 `_produce_one()` 内记录 sample ID、batch index、endpoint 和原因,然后丢弃该样本: + +- response 为空; +- token IDs 不一致; +- hidden rows 为 0 或覆盖不了选择位置; +- layer/hidden size/layout 不符合 contract; +- 截断后没有有效训练 token。 + +只有成功转换成 `DraftFeatureSample` 的样本才计入 `CollectionPayload.collected_samples`。 + +### 8.2 请求级错误 + +连接/读取错误先使用现有 client pool 重试。超过最大重试后,第一版建议默认让当前 collection 失败并终止本 step,而不是静默用不完整 batch 训练;可以后续增加 `failure_policy=fail_step|skip_sample|skip_collection`。 + +### 8.3 多 rank 一致性 + +vLLM 请求和样本过滤都发生在 driver;driver 形成最终成功样本列表后才按 owner 分桶并发 RPC。因此每个 drafter owner 收到的数量是 scheduler 已知的,不让各训练 rank 自行请求 vLLM、各自过滤。这避免某 rank 接受、另一个 rank 拒绝后进入不同 collective 顺序。 + +## 9. 指标与日志 + +建议增加: + +```text +drafter/vllm_prefill/candidate_samples +drafter/vllm_prefill/submitted_samples +drafter/vllm_prefill/succeeded_samples +drafter/vllm_prefill/dropped_samples +drafter/vllm_prefill/request_elapsed_sec +drafter/vllm_prefill/wait_elapsed_sec +drafter/vllm_prefill/overlap_elapsed_sec +drafter/vllm_prefill/payload_mib +drafter/vllm_prefill/retry_count +drafter/vllm_prefill/per_endpoint_inflight +``` + +每次 collection 至少记录:`global_step`、`collection_id`、候选数、提交数、成功数、按 owner 分桶数量、hidden rows、payload bytes 和等待时间。单样本拒绝日志记录 sample ID 和失败检查项,但不要打印完整 token 或 hidden tensor。 + +## 10. 测试计划 + +### 10.1 单元测试 + +1. padded `prompts/responses` 能还原正确有效 token。 +2. `prefill_ids == prompt_ids + response_ids[:-1]` 的边界测试,包括 response 长度 0/1。 +3. co-train request builder 与 standalone builder 对同一 token 输入产生一致的 loss mask、feature positions 和截断结果。 +4. 多 endpoint 并发、重试和 endpoint 选择沿用现有 client pool 测试。 +5. token 不一致、hidden rows=0、缺层、空窗口均被 producer 拒绝,且不进入 payload。 +6. `VllmPrefillCollectionAdapter` 能按 replica owner 正确分桶。 +7. vLLM source 开启时 old-logprob batch 不包含任何 hidden capture key。 +8. 标准 old-logprob 结果仍正确返回给 PPO。 +9. producer Future 在训练计划生成前完成 collection;关闭时无残留线程和 HTTP client。 + +### 10.2 集成测试 + +用小模型和两个 hidden-state endpoints 运行数个 co-train steps,对比: + +- old-logprob capture 与 vLLM prefill 的 token IDs、positions、hidden shape; +- 相同权重和样本下 drafter loss/metrics 是否接近; +- actor PPO metrics 是否不变; +- collect interval 未命中时没有 vLLM hidden 请求; +- endpoint 临时失败时重试后能继续; +- 多 drafter owner 下每个 owner 收到预期样本数。 + +## 11. 建议实施顺序与文件清单 + +### 阶段 A:抽取公共 producer 能力 + +- 修改 `verl_speco/producer/input_reader.py`:增加 rollout-token request builder,共享截断/对齐 helper。 +- 新增 `verl_speco/producer/cotrain_vllm_feature_producer.py`:长期 client pool、异步 batch submit、转换与过滤、关闭逻辑。 +- 不改 standalone TQ publish 行为。 + +### 阶段 B:接入 scheduler 和 driver + +- 修改 `verl_speco/trainer/scheduler/schedule_types.py`:增加 `VLLM_PREFILL`。 +- 修改 `verl_speco/trainer/scheduler/collection_adapter.py`:增加 owner 分桶 adapter。 +- 修改 `verl_speco/trainer/scheduler/drafter_scheduler.py`:注册 adapter。 +- 修改 `verl_speco/trainer/speco_ray_trainer.py`:提交 future、等待、构造 payload、执行 collection、关闭 producer。 + +### 阶段 C:配置、示例和验证 + +- 修改默认 drafter training 配置:增加 source 开关与 vLLM client 参数。 +- 修改 co-train example:关闭 old-logprob hidden capture,填写 hidden-state endpoints。 +- 增加 request builder、adapter、driver hook 和端到端测试。 +- 用同一批固定 token 对照 actor-captured hidden 与 vLLM hidden,再进行正式性能测试。 + +### 阶段 D:可选 TQ 数据面 + +- 仅在 Ray object store 成为瓶颈后实施。 +- 新增 co-train TQ key/ref adapter、worker decode 和 collection finalize/rollback 清理。 +- 不改变独立训练现有 TQ 协议。 + +## 12. 最终推荐 + +推荐第一版采用: + +```text +独立 hidden-state vLLM endpoints + + 复用 VllmFeatureClientPool + + 复用 TokenizedRequest / FeatureContract / feature_from_vllm_payload + + 新增 VLLM_PREFILL scheduler source + + 复用现有 SpecoWorker collection/buffer/train + + 暂不在 co-train 中引入 TQ +``` + +这样修改范围集中在“hidden-state 来源”和“异步接入点”,不会重写已经稳定的 drafter worker 与训练逻辑,也不会影响 PPO 必需的 actor old-logprob 计算。等这一版验证 hidden 数值、loss 和吞吐后,再决定是否把 Ray 中的大 tensor 数据面替换为 TQ。 diff --git a/tests/unit/test_standalone_resume.py b/tests/unit/test_standalone_resume.py new file mode 100644 index 00000000..9708c0f1 --- /dev/null +++ b/tests/unit/test_standalone_resume.py @@ -0,0 +1,55 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from pathlib import Path + +import pytest + +from verl_speco.trainer.standalone_resume import ( + load_standalone_resume, + save_standalone_resume, +) + + +def test_standalone_resume_round_trip_and_input_validation(tmp_path: Path) -> None: + input_path = tmp_path / "train.jsonl" + input_path.write_text('{"prompt":"q","response":"a"}\n', encoding="utf-8") + checkpoint_path = tmp_path / "draft_step_7" + + save_standalone_resume( + checkpoint_path, + [5, 1, 5, 3], + optimizer_step=7, + input_path=input_path, + ) + + consumed, metadata = load_standalone_resume( + checkpoint_path, input_path=input_path + ) + assert consumed == {1, 3, 5} + assert metadata is not None + assert metadata["optimizer_step"] == 7 + assert metadata["consumed_count"] == 3 + + input_path.write_text('{"prompt":"changed","response":"a"}\n', encoding="utf-8") + with pytest.raises(ValueError, match="input file changed"): + load_standalone_resume(checkpoint_path, input_path=input_path) + + +def test_missing_standalone_resume_is_not_a_resume_checkpoint( + tmp_path: Path, +) -> None: + consumed, metadata = load_standalone_resume(tmp_path / "pretrained") + assert consumed == set() + assert metadata is None diff --git a/tests/unit/test_standalone_tq_training_launcher.py b/tests/unit/test_standalone_tq_training_launcher.py index f0921c0e..a8992740 100644 --- a/tests/unit/test_standalone_tq_training_launcher.py +++ b/tests/unit/test_standalone_tq_training_launcher.py @@ -22,6 +22,7 @@ from verl_speco.standalone_tq_training_launcher import ( _preflight_input_file, + _producer_max_samples, _target_final_layer_id, build_pipeline_commands, resolve_pipeline_config, @@ -53,6 +54,17 @@ def test_pipeline_config_derives_transport_identity_from_training_args() -> None assert config.run_id.startswith("dspark-") +def test_producer_max_samples_uses_remaining_total_steps() -> None: + args = [ + *_training_args(), + "actor_rollout_ref.rollout.drafter.training.batch_size_per_gpu=2", + "speco.draft_training.nproc_per_node=4", + "speco.draft_training.nnodes=1", + ] + + assert _producer_max_samples(args, resumed_optimizer_step=6) == 32 + + def test_pipeline_config_reads_non_dspark_algorithm_from_training_args() -> None: args = [ item.replace("speculative_algorithm=DSPARK", "speculative_algorithm=DFLASH") diff --git a/tests/unit/test_tq_consumer.py b/tests/unit/test_tq_consumer.py index 59b5736b..92f33770 100644 --- a/tests/unit/test_tq_consumer.py +++ b/tests/unit/test_tq_consumer.py @@ -196,6 +196,7 @@ def test_world_size_one_streams_trains_then_clears_and_stops() -> None: batch = next(iterator) assert batch.local_keys == [entry.key for entry in entries] assert batch.global_keys == batch.local_keys + assert batch.global_sequence_nos == [0, 1] assert store.get_calls == [batch.local_keys] loader.clear_completed_batch(batch.global_keys) diff --git a/tests/unit/test_tq_producer.py b/tests/unit/test_tq_producer.py index 1204b8c1..d6873397 100644 --- a/tests/unit/test_tq_producer.py +++ b/tests/unit/test_tq_producer.py @@ -24,6 +24,7 @@ from verl_speco.producer.vllm_feature_client import RawVllmFeature from verl_speco.standalone_tq_producer import run_producer, validate_producer_config +from verl_speco.trainer.standalone_resume import save_standalone_resume from verl_speco.transport.drafter_sample_protocol import PROTOCOL_SCHEMA_VERSION @@ -281,6 +282,42 @@ def test_run_producer_restarts_input_until_max_samples(tmp_path: Path) -> None: assert eos_tags[0]["total_samples"] == 5 +def test_run_producer_skips_consumed_sequences_before_vllm(tmp_path: Path) -> None: + input_path = tmp_path / "input.jsonl" + _write_input(input_path) + checkpoint_path = tmp_path / "draft_step_1" + save_standalone_resume( + checkpoint_path, + [0], + optimizer_step=1, + input_path=input_path, + ) + config = _config(input_path) + producer_cfg = config["speco"]["standalone_tq_producer"] + producer_cfg["resume_checkpoint_path"] = str(checkpoint_path) + producer_cfg["max_samples"] = 2 + transport = _Transport() + pool = _Pool(tmp_path) + + stats = asyncio.run( + run_producer( + config, + transport=transport, + tokenizer=_Tokenizer(), + client_pool=pool, + ) + ) + + sequence_nos = sorted( + int(tag["sequence_no"]) + for tag in transport.records.values() + if tag.get("record_type") == "sample" + ) + assert stats.input_count == 2 + assert pool.prefill_calls == 2 + assert sequence_nos == [1, 2] + + def test_run_producer_generates_response_for_verl_chat_prompt(tmp_path: Path) -> None: input_path = tmp_path / "dapo.jsonl" input_path.write_text( diff --git a/verl_speco/config/speco_base.yaml b/verl_speco/config/speco_base.yaml index 02f57a5e..0484fc06 100644 --- a/verl_speco/config/speco_base.yaml +++ b/verl_speco/config/speco_base.yaml @@ -25,6 +25,9 @@ speco: # target vLLM; rows with a response are replayed directly for hidden states. standalone_tq_producer: input_path: null + # Set internally by the unified launcher when model_path is a resumable + # standalone checkpoint. Direct Producer users may leave it null. + resume_checkpoint_path: null tokenizer_path: null tokenizer_fingerprint: null target_model_id: null diff --git a/verl_speco/standalone_tq_producer.py b/verl_speco/standalone_tq_producer.py index 3097de64..e8df7170 100644 --- a/verl_speco/standalone_tq_producer.py +++ b/verl_speco/standalone_tq_producer.py @@ -41,6 +41,7 @@ delete_temporary_result, ) from verl_speco.trainer.feature_store import DraftFeatureSample +from verl_speco.trainer.standalone_resume import load_standalone_resume from verl_speco.trainer.target_feature_replay import ( FeatureContract, HiddenStateAlignmentError, @@ -199,6 +200,17 @@ async def run_producer( ) logger.info("Standalone TQ Producer observed owner_ready run_id=%s", run_id) + consumed_sequence_nos, resume_metadata = load_standalone_resume( + producer_cfg.get("resume_checkpoint_path"), + input_path=str(producer_cfg["input_path"]), + ) + logger.info( + "Standalone TQ Producer resume progress checkpoint=%s consumed=%s step=%s", + producer_cfg.get("resume_checkpoint_path"), + len(consumed_sequence_nos), + None if resume_metadata is None else resume_metadata.get("optimizer_step"), + ) + if tokenizer is None: logger.info( "Standalone TQ Producer loading tokenizer path=%s", @@ -245,19 +257,26 @@ async def run_producer( async def read_inputs() -> None: max_samples = int(producer_cfg.get("max_samples", 0) or 0) epoch = 0 + source_sequence_no = 0 while True: epoch_count = 0 + scanned_count = 0 for source_record in iter_input_records( str(producer_cfg["input_path"]) ): if max_samples > 0 and stats.input_count >= max_samples: break + sequence_no = source_sequence_no + source_sequence_no += 1 + scanned_count += 1 + if sequence_no in consumed_sequence_nos: + continue # iter_input_records restarts sequence_no at zero on every # pass. TQ keys require a run-global sequence number so a # repeated sample never overwrites an earlier pending copy. record = replace( source_record, - sequence_no=stats.input_count, + sequence_no=sequence_no, ) request = ( prepare_generation_request(record, tokenizer, producer_cfg) @@ -276,7 +295,7 @@ async def read_inputs() -> None: record.sample_id, record.response is not None, ) - if epoch_count == 0 and stats.input_count == 0: + if scanned_count == 0 and stats.input_count == 0: raise ValueError("Standalone TQ Producer input contains no samples") if max_samples <= 0 or stats.input_count >= max_samples: break diff --git a/verl_speco/standalone_tq_training_launcher.py b/verl_speco/standalone_tq_training_launcher.py index 4c9b6fa1..3e23da83 100644 --- a/verl_speco/standalone_tq_training_launcher.py +++ b/verl_speco/standalone_tq_training_launcher.py @@ -39,6 +39,8 @@ from urllib.request import urlopen import uuid +from verl_speco.trainer.standalone_resume import load_standalone_resume + logger = logging.getLogger(__name__) @@ -264,7 +266,9 @@ def _positive_int_override( return value -def _producer_max_samples(training_args: Sequence[str]) -> int: +def _producer_max_samples( + training_args: Sequence[str], *, resumed_optimizer_step: int = 0 +) -> int: """Return samples needed for exactly max_steps complete global batches.""" raw_max_steps = _find_override(training_args, _MAX_STEPS_KEY) @@ -278,7 +282,8 @@ def _producer_max_samples(training_args: Sequence[str]) -> int: ) nproc = _positive_int_override(training_args, _NPROC_KEYS, default=1) nnodes = _positive_int_override(training_args, _NNODES_KEYS, default=1) - return max_steps * batch_size * nproc * nnodes + remaining_steps = max(max_steps - int(resumed_optimizer_step), 0) + return remaining_steps * batch_size * nproc * nnodes def _stable_path_identity(kind: str, path: str) -> str: @@ -399,6 +404,16 @@ def build_pipeline_commands( ) -> PipelineCommands: """Build the internal commands without exposing transport options.""" + drafter_path = _strip_quotes(_find_override(training_args, _DRAFTER_PATH_KEY) or "") + _, resume_metadata = load_standalone_resume( + drafter_path or None, + input_path=config.input_path, + ) + resumed_optimizer_step = ( + int(resume_metadata.get("optimizer_step", 0)) + if resume_metadata is not None + else 0 + ) tq_overrides = [ f"{_TQ_PREFIX}.enable=true", f"{_TQ_PREFIX}.ray.address={ray_address}", @@ -480,6 +495,8 @@ def build_pipeline_commands( *tq_overrides, *producer_tuning_overrides, f"speco.standalone_tq_producer.input_path={config.input_path}", + "speco.standalone_tq_producer.resume_checkpoint_path=" + + (drafter_path if resume_metadata is not None else "null"), f"speco.standalone_tq_producer.tokenizer_path={config.tokenizer_path}", "speco.standalone_tq_producer.tokenizer_fingerprint=" + _stable_path_identity("tokenizer", config.tokenizer_path), @@ -492,7 +509,12 @@ def build_pipeline_commands( + _hydra_list(config.vllm_endpoints), f"speco.standalone_tq_producer.vllm_model={config.model_path}", "speco.standalone_tq_producer.max_samples=" - + str(_producer_max_samples(training_args)), + + str( + _producer_max_samples( + training_args, + resumed_optimizer_step=resumed_optimizer_step, + ) + ), ] consumer_internal = [ f"{_FEATURE_STORE_PREFIX}.type=tq", diff --git a/verl_speco/trainer/draft_training_loop.py b/verl_speco/trainer/draft_training_loop.py index e61e7c2a..61dd58f4 100644 --- a/verl_speco/trainer/draft_training_loop.py +++ b/verl_speco/trainer/draft_training_loop.py @@ -40,6 +40,10 @@ build_feature_store_from_config, ) from verl_speco.trainer.standalone_checkpoint import rewrite_standalone_runtime_config +from verl_speco.trainer.standalone_resume import ( + load_standalone_resume, + save_standalone_resume, +) from verl_speco.trainer.tq_sample_source import TQFeatureDataLoader, TQLocalBatch logger = logging.getLogger(__name__) @@ -131,6 +135,8 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: attempted_batches = 0 last_save_result: dict[str, Any] | None = None last_saved_step = 0 + consumed_sequence_nos: set[int] = set() + standalone_input_path = _standalone_input_path(config) store = None feature_replayer = None feature_producer = None @@ -155,6 +161,22 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: initial_optimizer_step = int(trainer.optimizer_steps_total) optimizer_step = initial_optimizer_step last_saved_step = optimizer_step + if feature_store_type == "tq": + consumed_sequence_nos, resume_metadata = load_standalone_resume( + drafter_cfg.get("model_path"), + input_path=standalone_input_path, + ) + if initial_optimizer_step > 0 and resume_metadata is None: + raise ValueError( + "Standalone TQ checkpoint restored optimizer state but has no " + "standalone_resume.json; exact data resume is unavailable" + ) + logger.info( + "[standalone rank=%s] resume progress optimizer_step=%s consumed=%s", + rank, + initial_optimizer_step, + len(consumed_sequence_nos), + ) current_stage = "open_feature_store" stage_started = time.perf_counter() @@ -294,7 +316,7 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: ) sample_source = feature_producer sample_iterator = iter(sample_source) - while max_steps <= 0 or successful_steps < max_steps: + while max_steps <= 0 or optimizer_step < max_steps: current_stage = "load_next_batch" loaded_batch = _next_batch_across_ranks( sample_iterator, @@ -419,6 +441,10 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: rank=rank, device=trainer.runtime_device, ) + if rank == 0: + consumed_sequence_nos.update( + tq_local_batch.global_sequence_nos or [] + ) successful_steps += 1 optimizer_step = int(trainer.optimizer_steps_total) if optimizer_step <= initial_optimizer_step: @@ -436,7 +462,14 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: _log_standalone_step_metrics(step_metrics, rank=rank) if save_interval > 0 and optimizer_step % save_interval == 0: current_stage = "save_checkpoint" - last_save_result = _save_standalone_checkpoint(trainer, optimizer_step) + last_save_result = _save_standalone_checkpoint( + trainer, + optimizer_step, + consumed_sequence_nos=consumed_sequence_nos, + input_path=( + standalone_input_path if feature_store_type == "tq" else None + ), + ) if _sync_any_rank_saved_checkpoint(last_save_result.get("saved")): last_saved_step = optimizer_step _barrier() @@ -445,7 +478,13 @@ async def _run_standalone_draft_training_async(config) -> dict[str, Any]: if final_save and successful_steps > 0 and optimizer_step != last_saved_step: current_stage = "save_final_checkpoint" last_save_result = _save_standalone_checkpoint( - trainer, optimizer_step, wait=True + trainer, + optimizer_step, + wait=True, + consumed_sequence_nos=consumed_sequence_nos, + input_path=( + standalone_input_path if feature_store_type == "tq" else None + ), ) _barrier() except Exception: @@ -503,8 +542,16 @@ def _build_backend(draft_config): def _save_standalone_checkpoint( - trainer: DrafterBaseTrainer, step: int, *, wait: bool = False + trainer: DrafterBaseTrainer, + step: int, + *, + wait: bool = False, + consumed_sequence_nos: set[int] | None = None, + input_path: str | None = None, ) -> dict[str, Any]: + consumed_snapshot = torch.tensor( + sorted(consumed_sequence_nos or ()), dtype=torch.int64 + ) save_checkpoint = getattr(trainer, "save_checkpoint", None) if callable(save_checkpoint): result = save_checkpoint(int(step), wait=wait) @@ -513,6 +560,12 @@ def _save_standalone_checkpoint( if result.get("saved") and checkpoint_path and is_export_leader: if wait: _rewrite_standalone_block_runtime_config(trainer, checkpoint_path) + _save_resume_sidecar( + checkpoint_path, + consumed_snapshot, + step=step, + input_path=input_path, + ) else: future = getattr(trainer, "_pending_full_checkpoint_future", None) if future is not None: @@ -521,6 +574,9 @@ def _save_standalone_checkpoint( trainer, checkpoint_path, completed, + consumed_snapshot=consumed_snapshot, + step=step, + input_path=input_path, ) ) return result @@ -551,12 +607,21 @@ def _save_standalone_checkpoint( future.result() trainer._pending_full_checkpoint_future = None _rewrite_standalone_block_runtime_config(trainer, checkpoint_path) + _save_resume_sidecar( + checkpoint_path, + consumed_snapshot, + step=step, + input_path=input_path, + ) elif future is not None: future.add_done_callback( - lambda completed: _rewrite_standalone_block_runtime_config( + lambda completed: _finalize_standalone_checkpoint( trainer, checkpoint_path, completed, + consumed_snapshot=consumed_snapshot, + step=step, + input_path=input_path, ) ) return { @@ -592,6 +657,10 @@ def _finalize_standalone_checkpoint( trainer: DrafterBaseTrainer, checkpoint_path: str, completed_future, + *, + consumed_snapshot: torch.Tensor | None = None, + step: int | None = None, + input_path: str | None = None, ) -> None: try: completed_future.result() @@ -602,6 +671,41 @@ def _finalize_standalone_checkpoint( return _rewrite_standalone_block_runtime_config(trainer, checkpoint_path) + _save_resume_sidecar( + checkpoint_path, + consumed_snapshot, + step=step, + input_path=input_path, + ) + + +def _save_resume_sidecar( + checkpoint_path: str, + consumed_snapshot: torch.Tensor | None, + *, + step: int | None, + input_path: str | None, +) -> None: + if consumed_snapshot is None or step is None or input_path is None: + return + save_standalone_resume( + checkpoint_path, + consumed_snapshot, + optimizer_step=step, + input_path=input_path, + ) + + +def _standalone_input_path(config: Any) -> str | None: + data_cfg = getattr(config, "data", None) + train_files = getattr(data_cfg, "train_files", None) + if train_files is None and isinstance(data_cfg, dict): + train_files = data_cfg.get("train_files") + if isinstance(train_files, str): + return train_files + if train_files is not None and len(train_files) == 1: + return str(train_files[0]) + return None def _torch_load_cpu(path: str) -> Any: diff --git a/verl_speco/trainer/standalone_resume.py b/verl_speco/trainer/standalone_resume.py new file mode 100644 index 00000000..a8c274aa --- /dev/null +++ b/verl_speco/trainer/standalone_resume.py @@ -0,0 +1,131 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Standalone-only data progress stored beside a drafter checkpoint.""" + +from __future__ import annotations + +import json +import os +from pathlib import Path +from typing import Any, Iterable + +import torch + + +RESUME_METADATA_NAME = "standalone_resume.json" +CONSUMED_SEQUENCE_NAME = "consumed_sequence_nos.pt" +RESUME_SCHEMA_VERSION = 1 + + +def build_input_fingerprint(path: str | os.PathLike[str]) -> dict[str, Any]: + input_path = Path(path).resolve() + stat = input_path.stat() + return { + "path": str(input_path), + "size_bytes": int(stat.st_size), + "mtime_ns": int(stat.st_mtime_ns), + } + + +def save_standalone_resume( + checkpoint_path: str | os.PathLike[str], + consumed_sequence_nos: Iterable[int] | torch.Tensor, + *, + optimizer_step: int, + input_path: str | os.PathLike[str], +) -> None: + checkpoint_dir = Path(checkpoint_path) + checkpoint_dir.mkdir(parents=True, exist_ok=True) + if isinstance(consumed_sequence_nos, torch.Tensor): + values = consumed_sequence_nos.detach().to(device="cpu", dtype=torch.int64) + else: + values = torch.tensor( + sorted({int(value) for value in consumed_sequence_nos}), + dtype=torch.int64, + ) + values = torch.unique(values.flatten(), sorted=True) + if values.numel() and int(values[0]) < 0: + raise ValueError("consumed sequence numbers must be non-negative") + + tensor_path = checkpoint_dir / CONSUMED_SEQUENCE_NAME + tensor_temporary = tensor_path.with_suffix(tensor_path.suffix + ".incomplete") + torch.save(values, tensor_temporary) + os.replace(tensor_temporary, tensor_path) + + metadata = { + "schema_version": RESUME_SCHEMA_VERSION, + "optimizer_step": int(optimizer_step), + "consumed_count": int(values.numel()), + "consumed_sequence_file": CONSUMED_SEQUENCE_NAME, + "input_fingerprint": build_input_fingerprint(input_path), + } + metadata_path = checkpoint_dir / RESUME_METADATA_NAME + metadata_temporary = metadata_path.with_suffix( + metadata_path.suffix + ".incomplete" + ) + metadata_temporary.write_text( + json.dumps(metadata, ensure_ascii=False, indent=2) + "\n", + encoding="utf-8", + ) + os.replace(metadata_temporary, metadata_path) + + +def load_standalone_resume( + checkpoint_path: str | os.PathLike[str] | None, + *, + input_path: str | os.PathLike[str] | None = None, +) -> tuple[set[int], dict[str, Any] | None]: + if not checkpoint_path: + return set(), None + checkpoint_dir = Path(checkpoint_path) + metadata_path = checkpoint_dir / RESUME_METADATA_NAME + if not metadata_path.is_file(): + return set(), None + metadata = json.loads(metadata_path.read_text(encoding="utf-8")) + if int(metadata.get("schema_version", 0)) != RESUME_SCHEMA_VERSION: + raise ValueError(f"Unsupported standalone resume metadata: {metadata_path}") + tensor_name = metadata.get("consumed_sequence_file") + if tensor_name != CONSUMED_SEQUENCE_NAME: + raise ValueError(f"Invalid consumed sequence file in {metadata_path}") + tensor_path = checkpoint_dir / tensor_name + try: + values = torch.load(tensor_path, map_location="cpu", weights_only=True) + except TypeError: + values = torch.load(tensor_path, map_location="cpu") + if not isinstance(values, torch.Tensor) or values.dtype != torch.int64: + raise ValueError(f"Invalid consumed sequence tensor: {tensor_path}") + values = values.flatten() + if values.numel() and int(values[0]) < 0: + raise ValueError(f"Negative consumed sequence number in {tensor_path}") + consumed = {int(value) for value in values.tolist()} + if len(consumed) != int(metadata.get("consumed_count", -1)): + raise ValueError(f"Consumed sequence count mismatch in {checkpoint_dir}") + if input_path is not None: + saved_fingerprint = metadata.get("input_fingerprint") + current_fingerprint = build_input_fingerprint(input_path) + if saved_fingerprint != current_fingerprint: + raise ValueError( + "Standalone resume input file changed since the checkpoint was saved: " + f"saved={saved_fingerprint!r}, current={current_fingerprint!r}" + ) + return consumed, metadata + + +__all__ = [ + "CONSUMED_SEQUENCE_NAME", + "RESUME_METADATA_NAME", + "build_input_fingerprint", + "load_standalone_resume", + "save_standalone_resume", +] diff --git a/verl_speco/trainer/tq_sample_source.py b/verl_speco/trainer/tq_sample_source.py index 3ec0ce50..07ab9d06 100644 --- a/verl_speco/trainer/tq_sample_source.py +++ b/verl_speco/trainer/tq_sample_source.py @@ -36,6 +36,7 @@ class TQLocalBatch: local_keys: list[str] local_samples: list[DraftFeatureSample] global_keys: list[str] | None + global_sequence_nos: list[int] | None = None def build_assignments( @@ -117,6 +118,9 @@ def __iter__(self) -> Iterator[TQLocalBatch]: command = { "kind": "batch", "global_keys": [entry.key for entry in selected], + "global_sequence_nos": [ + int(entry.tag["sequence_no"]) for entry in selected + ], "assignments": [ [_entry_to_wire(entry) for entry in rank_entries] for rank_entries in assignments @@ -166,10 +170,16 @@ def __iter__(self) -> Iterator[TQLocalBatch]: if self.rank == 0 else None ) + global_sequence_nos = ( + [int(value) for value in command.get("global_sequence_nos", [])] + if self.rank == 0 + else None + ) yield TQLocalBatch( local_keys=[entry.key for entry in local_entries], local_samples=samples, global_keys=global_keys, + global_sequence_nos=global_sequence_nos, ) def clear_completed_batch(self, global_keys: Sequence[str] | None) -> None: From 92d9432bd4c87b231cdc252e157b99dc6f04a00b Mon Sep 17 00:00:00 2001 From: xyr <1615681145@qq.com> Date: Tue, 1 Sep 2026 10:03:38 +0800 Subject: [PATCH 46/50] remove local personal planning docs from tracking --- ...sync_vllm_mooncake_dspark_training_plan.md | 2902 ----------------- ...eature_sample_tq_protocol_refactor_plan.md | 549 ---- docs/standalone_tq_consumer_implementation.md | 713 ---- docs/standalone_tq_drafter_resume_plan.md | 448 --- ...standalone_tq_foundation_implementation.md | 1030 ------ docs/standalone_tq_producer.md | 212 -- docs/standalone_tq_training_parameters.md | 169 - ...standalone_vllm_tq_dspark_training_plan.md | 1130 ------- docs/transferqueue_integration_plan.md | 145 - docs/vllm_direct_hidden_state_cotrain_plan.md | 666 ---- 10 files changed, 7964 deletions(-) delete mode 100644 docs/async_vllm_mooncake_dspark_training_plan.md delete mode 100644 docs/draft_feature_sample_tq_protocol_refactor_plan.md delete mode 100644 docs/standalone_tq_consumer_implementation.md delete mode 100644 docs/standalone_tq_drafter_resume_plan.md delete mode 100644 docs/standalone_tq_foundation_implementation.md delete mode 100644 docs/standalone_tq_producer.md delete mode 100644 docs/standalone_tq_training_parameters.md delete mode 100644 docs/standalone_vllm_tq_dspark_training_plan.md delete mode 100644 docs/transferqueue_integration_plan.md delete mode 100644 docs/vllm_direct_hidden_state_cotrain_plan.md diff --git a/docs/async_vllm_mooncake_dspark_training_plan.md b/docs/async_vllm_mooncake_dspark_training_plan.md deleted file mode 100644 index 525da0fb..00000000 --- a/docs/async_vllm_mooncake_dspark_training_plan.md +++ /dev/null @@ -1,2902 +0,0 @@ -# verl-SpeCo 异步 vLLM → TransferQueue(Mooncake 后端)→ DSpark 流式训练方案 - -Last updated: 08/21/2026 - -## 1. 文档范围 - -本文只讨论当前 `verl-SpeCo-ls` 项目的独立草稿模型训练入口: - -```text -examples/run_qwen3-8b_drafter_separate_training.sh - → python -m verl_speco.standalone_tq_training_launcher - ├─→ TransferQueue owner - ├─→ vLLM hidden-state producer - └─→ TransferQueue consumer - → python -m verl_speco.draft_train_launcher - → torch.distributed.run - → python -m verl_speco.draft_train - → run_standalone_draft_training() -``` - -目标是训练 `mode=train/offline` 的草稿模型,不是 RL 与草稿模型一起训练,也不是先生成全量 hidden states 再长期保存到磁盘。 - -输入既可以是已有 response 的 replay 文件,也可以是 verl prompt-only Parquet。新流水线需要: - -1. Producer 读取 prompt;缺少 response 时由 target vLLM 生成; -2. Producer 并行请求一个或多个 vLLM endpoint 做 prefill; -3. hidden states 写入 TransferQueue(简称 TQ),TQ 的数据后端使用 Mooncake; -4. DSpark 的各训练 rank 从 TQ/Mooncake 并行读取自己的 local batch; -5. 所有 rank 完成同一个 optimizer step 后,清理这一批 TQ 数据; -6. 不使用自研 Stream Coordinator;第一版复用 PR #48 已验证的 TQ KV 传输模式,由 TQ 保存 payload 和 tags,rank 0 通过 `kv_list` 发现 ready key 并组成 global batch。 - -本文不照搬 AngelSpec 的进程组织。实现依据是旁边只读参考仓库 `verl-SpeCo` 的 PR #48 分支,尤其是: - -```text -verl_speco/integration/transferqueue_bridge.py -verl_speco/integration/sglang_runtime.py -verl_speco/integration/oldlogprob_runtime.py -verl_speco/workers/speco_worker.py -verl_speco/integration/task_runner.py -``` - -参考仓库只用于理解和复用设计,不修改其代码。真正实现仍放在本项目 `verl-SpeCo-ls`,服务于 standalone drafter training。 - -建议按下面顺序阅读: - -1. 第一部分先介绍 verl/SpeCo 中的进程、`TokenOutput`、`DataProto`、WorkerGroup、replica 和 owner rank; -2. 再完整解释 PR #48 的 SGLang TQ 路径; -3. 再解释 PR #48 的 old-logprob TQ 路径; -4. 再解释 hidden 恢复后如何进入 online buffer 和 optimizer step; -5. 第二部分才讨论如何把相同 TQ 传输能力改造成独立 Producer/DSpark Consumer。 - -## 第一部分:PR #48 原始 TQ 流程 - -这一部分只解释只读参考仓库 `../verl-SpeCo` 的 PR #48。这里的 Producer、Consumer、Ray driver、数据结构和生命周期全部是 PR #48 当前代码的行为,不是本文为 standalone 设计的行为。 - -### 1A. 阅读 PR #48 前必须知道的项目对象 - -#### SGLang server - -SGLang server 是 rollout 推理进程。它接收 prompt,执行 target model 推理并生成 response。SpeCo 的 SGLang patch 还会在推理过程中收集指定层 hidden states。 - -它不是 drafter trainer,也不执行 optimizer step。在 PR #48 中它是 hidden-state Producer。 - -#### TokenOutput - -`TokenOutput` 是一次 rollout request 的返回对象。核心字段可以概括为: - -```python -TokenOutput( - token_ids=list[int], - log_probs=..., - routed_experts=..., - extra_fields={ - "global_steps": int, - "drafter_sample": dict | None, - }, -) -``` - -`token_ids` 是生成结果;`extra_fields` 是 SpeCo 添加的旁路字段。`drafter_sample` 不影响正常 response 返回,它用于把草稿训练所需信息从 rollout 侧带回训练控制层。 - -#### DataProto 和 non_tensor_batch - -verl 将一批 rollout 结果整理成 `DataProto`。它通常分为: - -```python -DataProto( - batch=TensorDict(...), - non_tensor_batch={...}, - meta_info={...}, -) -``` - -- `batch`:规则的批量 tensor,例如 prompts、responses、attention mask; -- `non_tensor_batch`:不能直接组成规则 dense batch 的 Python 对象或 object array; -- `meta_info`:批次级配置和指标。 - -每个 request 的 `TokenOutput.extra_fields["drafter_sample"]` 最终会被 rollout/agent-loop 聚合到: - -```python -gen_batch_output.non_tensor_batch["drafter_sample"] -``` - -所以 SGLang 侧写的是单 request `TokenOutput`,driver 侧拿到的是批量 `DataProto`。 - -#### RayPPOTrainer driver - -`SpecoRayPPOTrainer` 所在进程是控制进程,简称 driver。它负责按训练 step 依次调用 rollout、old-logprob、actor update,以及向各 WorkerGroup 发 RPC。 - -driver 不执行 drafter 模型 forward。PR #48 之前,它会接触包含 hidden tensor 的 `drafter_sample`;PR #48 之后,它只处理中小 tensor、标量 metadata 和 TQ key。 - -#### WorkerGroup - -WorkerGroup 是 verl 对一组 Ray worker 的调用封装。代码: - -```python -self.drafter_wg.collect_rollout_features(buckets) -``` - -不是普通本地函数调用,而是根据注册的 dispatch rule,将参数分发到多张卡上的 `SpecoWorker.collect_rollout_features()`。 - -#### Rollout replica - -rollout replica 是一组共同承载一个 rollout 模型副本的进程/GPU。一个任务可能存在多个 data-parallel rollout replica。`replica_rank` 标识当前 sample 是哪个 rollout replica 生成的: - -```text -replica_rank = 0, 1, 2, ... -``` - -#### Drafter training replica、DP rank 和 SP rank - -drafter 训练也可能按 data parallel 和 sequence parallel 组织: - -```text -drafter replica / DP rank 0 - ├─ SP rank 0 - └─ SP rank 1 - -drafter replica / DP rank 1 - ├─ SP rank 0 - └─ SP rank 1 -``` - -同一个 drafter DP replica 内的 SP ranks 共同执行一个模型副本的训练。`replica_rank` 用来把 rollout replica 产生的数据路由到对应 drafter DP replica。 - -#### Owner rank - -`collect_rollout_features` 注册了: - -```python -@register( - dispatch_mode=make_nd_compute_dispatch_fn( - mesh_name="drafter_owner_route" - ) -) -``` - -每个 drafter DP replica 的 `SP rank 0` 被标记为 collect leader/owner。driver 传入的是按 replica 分好的 bucket,dispatch 层负责把 bucket 发到对应训练组。一个 replica 内可能有多个 rank 需要共同训练,但只有指定 leader 负责汇总 RPC 返回。 - -这里的 owner route 是 PR #48 为什么不能简单让任意 worker 随机拿 sample 的原因:sample 必须进入与当前 drafter device mesh 一致的训练组。 - -## 2. PR #48 改造前的 online 特征流程 - -PR #48 改造的是 SpeCo 的 online drafter feature transport。hidden states 有两条主要来源。 - -### 2.1 SGLang rollout hidden 路径 - -改造前: - -```text -SGLang server - → 生成 drafter_sample,其中直接包含 hidden_states CPU tensor - → TokenOutput.extra_fields - → RayPPOTrainer driver 收集 drafter_sample - → driver 按 drafter replica/owner 分桶 - → Ray dispatch / object store - → SpecoWorker.collect_rollout_features(samples) - → _store_rollout_sample() - → online drafter buffer/train -``` - -此时 `drafter_sample` 类似: - -```python -{ - "input_ids": int64[1, total_len], - "prompts": int64[1, prompt_len], - "responses": int64[1, response_len], - "hidden_states": bf16[1, hidden_rows, hidden_dim], - "hidden_positions": int64[1, hidden_rows], - "target_logprobs": tensor | None, - "global_step": 42, - "replica_rank": 1, -} -``` - -问题是整个字典经过 driver,而 `hidden_states` 是其中最大的字段。driver 的 host memory 和 Ray object store 都会承载这些 tensor。 - -### 2.2 old-logprob hook hidden 路径 - -另一条路径在 actor old-logprob forward 中捕获 hidden states。改造前,大 chunk 通过: - -```python -chunk_ref = ray.put(hidden_chunk) -``` - -sample 不直接带 tensor,而是带: - -```python -{ - "hidden_states_ref_chunks": [ - { - "ref": ray_object_ref, - "start": 0, - "length": 512, - }, - ], -} -``` - -drafter worker 再 `ray.get(ref)`,根据 `start/length` 切出每条 sample 所需行。 - -### 2.3 PR #48 要改变的边界 - -PR #48 没有改变: - -- rollout 什么时候产生 sample; -- driver 如何触发 drafter worker; -- drafter worker 如何调用 `_store_rollout_sample()`; -- drafter model 的训练逻辑; -- drafter 权重发布。 - -它只改变大 tensor 的跨进程介质: - -```text -改造前:Producer → Ray driver/object store → Consumer -改造后:Producer → TQ storage → Consumer - key 仍走原 Ray 控制路径 -``` - -## 3. PR #48 改造后的完整 TQ 流程 - -### 3.0 总览 - -PR #48 的目标不是让 TQ 自己产生训练 batch,而是把原来经过 Ray driver/Ray object store 的大 hidden tensor 攁到 TQ。原来的控制路径继续存在,只是控制路径上从“大 tensor”变成“小 key”。 - -整体结构是: - -```text - 原 Ray 控制路径 - drafter_sample / chunk ref -Producer ───────────────────── key ───────────────────▶ Consumer - │ │ - │ kv_put(large tensor) │ kv_batch_get(key) - ▼ ▼ -TransferQueue storage ─────────────────────────────────────┘ -``` - -因此 PR #48 同时保留两条通道: - -```text -控制通道:Producer → Ray driver → drafter worker -数据通道:Producer → TQ storage → drafter worker -``` - -控制通道负责告诉 Consumer “本次应该处理哪个样本、对应哪个 TQ key”;数据通道负责传输 hidden states 等大 tensor。 - -#### 3.0.1 配置放在哪里 - -PR #48 在 drafter training 配置下增加: - -```yaml -actor_rollout_ref: - rollout: - drafter: - training: - transfer_queue: - enable: false - backend: - storage_backend: SimpleStorage - SimpleStorage: - total_storage_size: 100000 - num_data_storage_units: 8 -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/config/speco_base.yaml -``` - -这里的 `transfer_queue` 是 SpeCo drafter 自己的配置,不是 standalone `feature_store.type`,也不是简单把上游 verl 的 `transfer_queue.enable` 打开。 - -#### 3.0.2 TaskRunner 创建整套 TQ - -RL 任务启动时,`SpecoTaskRunner.run()` 在创建 workers 之前调用: - -```python -from verl_speco.integration.transferqueue_bridge import ( - close_transfer_queue, - init_transfer_queue, -) - -transfer_queue_started = init_transfer_queue(config) -try: - trainer.init_workers() - trainer.fit() -finally: - if transfer_queue_started: - close_transfer_queue() -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/integration/task_runner.py:319 -``` - -`init_transfer_queue(config)` 内部读取: - -```python -config.actor_rollout_ref.rollout.drafter.training.transfer_queue -``` - -然后执行: - -```python -tq.init(_to_plain_dict(tq_cfg)) -``` - -并记录: - -```python -_state["initialized"] = True -_state["owner"] = True -``` - -这里的 owner 是“创建/拥有 TQ 生命周期的进程”。只有 owner 在任务结束时执行 `tq.close()`。 - -关键顺序是: - -```text -SpecoTaskRunner -→ tq.init(完整配置) -→ trainer.init_workers() -→ Ray workers 启动 -``` - -也就是说,PR #48 假设 TQ Controller/storage 已经由 TaskRunner 在 Ray 集群环境中建立,后启动的 worker 只需要连接。 - -#### 3.0.3 每个 Producer/Consumer 进程怎么连接 TQ - -bridge 中的 `_ensure_initialized()` 是进程级懒初始化: - -```python -def _ensure_initialized(): - if _state["initialized"]: - return - - with _state_lock: - if _state["initialized"]: - return - - tq.init() - _state["initialized"] = True -``` - -注意这里是: - -```python -tq.init() -``` - -不是: - -```python -tq.init(config) -``` - -无参初始化的含义是连接 TaskRunner 已创建的同一套 TQ。它依赖 PR #48 所处的 Ray 运行环境完成服务发现。 - -因此 PR #48 不是“每个 worker 各创建一套 TQ”,而是: - -```text -TaskRunner:tq.init(config),创建一次 -SGLang producer:tq.init(),连接 -actor producer:tq.init(),连接 -drafter consumer:tq.init(),连接 -``` - -#### 3.0.4 SGLang Producer 怎么写 hidden states - -SGLang 已经完成 rollout,并组装出 `drafter_sample` 后,PR #48 执行: - -```python -configure_transfer_queue(training_cfg) - -if is_transfer_queue_enabled(): - tq_key = make_sample_key( - collection_global_steps, - self.replica_rank, - request_id, - ) - - tq_payload = { - "hidden_states": hidden_states.unsqueeze(0).cpu(), - } - - if target_logprobs is not None: - tq_payload["target_logprobs"] = ( - target_logprobs.unsqueeze(0).cpu() - ) - - put_sample( - tq_key, - tq_payload, - tag={ - "global_step": collection_global_steps, - "replica_rank": self.replica_rank, - }, - ) - - drafter_sample["hidden_states_tq_key"] = tq_key - drafter_sample["hidden_states"] = None -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/integration/sglang_runtime.py:2020 -``` - -这里发生了两条不同的数据流: - -```text -大 tensor:SGLang → TQ -小字典/key:SGLang → 原 Ray/driver 路径 → drafter worker -``` - -写进 TQ 后将: - -```python -drafter_sample["hidden_states"] = None -``` - -是为了避免 hidden states 继续经过 driver/Ray object store。driver 仍收到 sample,但大 tensor 已替换成: - -```python -drafter_sample["hidden_states_tq_key"] -``` - -#### 3.0.5 `put_sample()` 实际怎么写 - -bridge 中: - -```python -def put_sample(key, tensor_dict, *, tag=None): - payload = { - k: v - for k, v in tensor_dict.items() - if torch.is_tensor(v) - } - - _ensure_initialized() - - tq.kv_put( - key=key, - partition_id="speco_drafter_features", - fields=payload, - tag=tag or {}, - ) -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/integration/transferqueue_bridge.py:206 -``` - -这里可以明确看到: - -- PR #48 使用 TQ 高层 KV API; -- 一个 key 对应一个 sample; -- `fields` 是 tensor 字典; -- `tag` 是小 metadata; -- partition 当前写死为 `speco_drafter_features`; -- 写入前 tensor 已 `.cpu()`; -- 写入失败直接抛异常,不静默回退。 - -key 的生成代码是: - -```python -def make_sample_key(global_step, replica_rank, request_id): - return f"speco:{global_step}:{replica_rank}:{request_id}" -``` - -这个 key 对 RL rollout 是合理的,因为它用 step、rollout replica 和 request ID 标识一次在线采样。 - -#### 3.0.5.1 PR #48 中“数据”和“元数据”分别长什么样 - -PR #48 一条 SGLang 样本在写 TQ 之前,`drafter_sample` 同时包含大 tensor 和控制信息。简化后类似: - -```python -drafter_sample = { - # 普通训练输入,仍走原 sample/Ray 控制路径 - "input_ids": int64[1, prompt_len + response_len], - "prompts": int64[1, prompt_len], - "responses": int64[1, response_len], - - # 大 tensor,开启 TQ 后从这个字典移除 - "hidden_states": bf16[1, hidden_rows, hidden_dim], - "target_logprobs": fp32[1, rows, topk_or_vocab] | None, - - # hidden 与 token 对齐所需的小字段 - "hidden_positions": int64[1, hidden_rows] | None, - "hidden_position_start": int, - "hidden_position_end": int, - "hidden_window_start": int, - "hidden_window_end": int, - - # 控制信息 - "global_step": int, - "replica_rank": int, -} -``` - -执行 `put_sample()` 时,并不是把整个 `drafter_sample` 放进 TQ。PR #48 只抽取占用大的 tensor: - -```python -tq_payload = { - "hidden_states": bf16[1, hidden_rows, hidden_dim], - "target_logprobs": fp32[1, rows, ...], # 可选 - "hidden_raw_target_logprobs": ..., # 可选 - "hidden_raw_target_logprobs_positions": ..., # 可选 -} -``` - -这就是 TQ 的 data payload。它被传给: - -```python -tq.kv_put(fields=tq_payload) -``` - -另外还有 TQ tag: - -```python -tag = { - "global_step": 42, - "replica_rank": 1, -} -``` - -tag 是 TQ 侧轻量 metadata,用于描述/检索对象,不承载 hidden tensor。 - -写入完成后,仍经 Ray 传递的轻量 `drafter_sample` 变成: - -```python -drafter_sample = { - "input_ids": int64[1, total_len], - "prompts": int64[1, prompt_len], - "responses": int64[1, response_len], - - "hidden_states": None, - "target_logprobs": None, - "hidden_states_tq_key": "speco:42:1:req-007", - - "hidden_positions": int64[1, hidden_rows] | None, - "hidden_position_start": 128, - "hidden_position_end": 640, - "global_step": 42, - "replica_rank": 1, -} -``` - -因此 PR #48 实际存在三类对象: - -| 对象 | 内容 | 传输路径 | 作用 | -|---|---|---|---| -| TQ fields/payload | `hidden_states` 等大 tensor | Producer → TQ storage → Consumer | 避免 driver 搬运大 tensor | -| TQ tag | `global_step`、`replica_rank` | TQ control/index metadata | 描述该 key | -| 轻量 `drafter_sample` | tokens、位置、标量 metadata、`hidden_states_tq_key` | 原 Ray driver 路径 | 告诉 Consumer 取哪个 key,以及如何解释取回的 tensor | - -代码实现解耦的关键不是“所有内容都进 TQ”,而是: - -```python -drafter_sample["hidden_states_tq_key"] = tq_key -drafter_sample["hidden_states"] = None -``` - -第一行给 Consumer 留下寻址信息;第二行阻止大 tensor 继续沿旧路径传输。 - -#### 3.0.5.2 Consumer 如何把两部分重新合成一个训练样本 - -Consumer 最初拿到的是轻量 sample: - -```python -sample["hidden_states"] is None -sample["hidden_states_tq_key"] == "speco:42:1:req-007" -``` - -它执行: - -```python -payload = get_sample(sample["hidden_states_tq_key"]) -sample["hidden_states"] = payload["hidden_states"] -``` - -合并后: - -```python -sample = { - "input_ids": ..., - "prompts": ..., - "responses": ..., - "hidden_positions": ..., - "hidden_states": bf16[1, hidden_rows, hidden_dim], - "hidden_states_tq_key": "speco:42:1:req-007", - ... -} -``` - -后面的 `_store_rollout_sample()` 看到的结构与关闭 TQ 时基本一致,所以训练主体不需要增加 TQ 分支。TQ bridge 只改变“大 tensor 从哪里恢复”,不改变 drafter trainer 的输入语义。 - -#### 3.0.6 old-logprob Producer 怎么写 chunk - -PR #48 的另一条 Producer 路径来自 actor old-logprob hidden hook。原来是: - -```python -chunk_ref = ray.put(hidden_chunk) -``` - -开启 TQ 后改成: - -```python -tq_key = make_sample_key( - global_step, - owner, - f"chunk{len(chunk_refs)}", -) - -put_sample( - tq_key, - {"hidden": hidden_chunk}, - tag={ - "global_step": global_step, - "owner": owner, - }, -) - -chunk_ref = tq_key -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/integration/oldlogprob_runtime.py:533 -``` - -后面的 driver 不需要知道 ref 是 Ray ObjectRef 还是 TQ string key,它只把 ref 当不透明 token 继续传递。 - -#### 3.0.7 Consumer 怎么根据 key 读取 - -drafter worker 收到原来的 sample 小字典后: - -```python -tq_key = sample.get("hidden_states_tq_key") - -if tq_key is not None and self._speco_tq_enabled: - payload = get_sample(tq_key) - - for field in ( - "hidden_states", - "target_logprobs", - "hidden_raw_target_logprobs", - "hidden_raw_target_logprobs_positions", - ): - if payload.get(field) is not None: - sample[field] = payload[field] - - if sample.get("hidden_states") is None: - raise RuntimeError( - "TQ key exists but hidden_states payload is missing" - ) -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/workers/speco_worker.py:848 -``` - -恢复 tensor 后,后面的逻辑仍使用原 `sample`/`batch`,drafter trainer 不需要知道 tensor 来自 Ray 还是 TQ。 - -#### 3.0.8 `get_sample()` 实际怎么读 - -```python -def get_sample(key): - _ensure_initialized() - - result = tq.kv_batch_get( - keys=[key], - partition_id="speco_drafter_features", - ) - - value = _extract_value(result, key) - return _tensordict_to_dict(value) -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/integration/transferqueue_bridge.py:237 -``` - -`_extract_value()` 兼容三种返回形态: - -```python -if isinstance(result, dict): - return result.get(key) -if isinstance(result, (list, tuple)): - return result[0] -return result -``` - -这是因为不同 TQ 版本/后端返回包装可能不同。 - -#### 3.0.9 为什么需要 `_densify_tq_tensor()` - -PR #48 后续修复发现,TQ 把单样本 tensor 放进 TensorDict 后,`kv_batch_get` 可能返回 NestedTensor,并额外带 batch 维。旧代码要执行: - -```python -tensor[start:start + length] -``` - -但 NestedTensor 不支持在 jagged dim 直接 slice。因此加入: - -```python -def _densify_tq_tensor(tensor): - if tensor.is_nested: - parts = [ - part - for part in tensor.unbind() - if part.numel() > 0 - ] - tensor = torch.cat(parts, dim=0) - - if tensor.dim() == 3: - tensor = tensor.squeeze(0) - elif tensor.dim() == 1: - tensor = tensor.unsqueeze(0) - - return tensor.contiguous() -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/workers/speco_worker.py:72 -``` - -standalone Consumer 同样必须做这个转换,不能假定 `kv_batch_get` 返回普通 dense `[seq, hidden]`。 - -#### 3.0.10 为什么需要 per-step cache - -old-logprob 路径里,一个 owner hidden chunk 可能被约 16 个 sample 共同引用。如果每个 sample 都: - -```python -get_sample(same_tq_key) -``` - -就会重复传输同一个数百 MB chunk。PR #48 在每次 `collect_rollout_features()` 开始时创建: - -```python -self._tq_chunk_cache = {} -``` - -解析 ref 时: - -```python -cache_key = ref if isinstance(ref, str) else id(ref) - -if cache_key not in cache: - cache[cache_key] = _resolve_tq_or_ray_ref(ref) - -tensor = cache[cache_key] -``` - -对应参考代码: - -```text -../verl-SpeCo/verl_speco/workers/speco_worker.py:98 -../verl-SpeCo/verl_speco/workers/speco_worker.py:854 -``` - -独立训练如果每个 sample 都是独立 TQ key,主要使用 `kv_batch_get(keys=[...])` 批量读取,不一定需要跨 sample chunk cache;但 prefetch 重试或共享 packed object 时仍应保留 key cache。 - -#### 3.0.11 PR #48 什么时候删除数据 - -PR #48 没有在 `get_sample()` 后删除。其注释明确说明:同一 drafter replica 的多个 TP/SP rank 可能读取同一 key,第一次读取后立刻删除会让剩余 rank 失败。 - -当前策略是任务结束时由 owner: - -```python -tq.close() -``` - -统一结束 TQ 生命周期。也就是说,PR #48 当前没有实现精细的逐 step `kv_clear`。 - -#### 3.0.12 PR #48 的完整时序 - -```text -SpecoTaskRunner - → tq.init(config) - → 启动 Ray workers - -SGLang/actor Producer process - → configure_transfer_queue() - → 第一次 put 时 tq.init() - → kv_put(key, tensor fields, tag) - → 把 key 塞回原 sample/ref - -Ray driver - → 只中转小 sample/key - -drafter worker Consumer process - → 第一次 get 时 tq.init() - → kv_batch_get([key]) - → 解包 TensorDict/NestedTensor - → 恢复 sample["hidden_states"] - → 原 drafter collect/train 逻辑 - -任务结束 - → TaskRunner owner tq.close() -``` - -### 3.1 已经实现的可复用能力 - -PR #48 新增 `verl_speco/integration/transferqueue_bridge.py`,锁定思路是把 TQ 当成独立传输库,不修改上游 verl。它提供: - -```python -configure_transfer_queue(training_cfg) -init_transfer_queue(config) -make_sample_key(global_step, replica_rank, request_id) -put_sample(key, tensor_dict, tag=...) -get_sample(key) -close_transfer_queue() -``` - -实际写入调用是: - -```python -tq.kv_put( - key=key, - partition_id="speco_drafter_features", - fields=payload, - tag=tag, -) -``` - -实际读取调用是: - -```python -result = tq.kv_batch_get( - keys=[key], - partition_id="speco_drafter_features", -) -``` - -另外,PR #48 已经处理了多项 standalone 方案也需要的问题: - -1. TQ 返回值可能是 direct value、`{key: value}` 或 list,需要统一解包; -2. TQ 返回的 tensor 可能是 NestedTensor,必须通过 `unbind + cat` 恢复为生产端写入的 dense tensor; -3. 多个样本引用同一 hidden chunk 时,需要 per-step cache,避免重复 `kv_batch_get` 同一个大对象; -4. TQ key 存在但 payload 缺失时 fail loud,不能静默丢样本; -5. `enable=false` 时保留原传输路径。 - -这些逻辑应直接作为本项目 TQ adapter 的参考。 - -### 3.2 PR #48 的数据流 - -PR #48 优化的是 RL online 路径: - -```text -SGLang/actor worker - → kv_put(hidden states) - → 把 hidden_states_tq_key 塞进原 drafter_sample - → 原 Ray driver 继续传递小 sample/key - → drafter worker collect_rollout_features() - → kv_batch_get(key) -``` - -它没有让 consumer 自己从 TQ 发现下一批 key;key 仍沿原来的 Ray 控制路径到达 drafter worker。 - -### 3.3 PR #48 没有提供的 standalone 能力 - -PR #48 当前没有实现: - -- 从预生成 response 文件读取数据的独立 Producer; -- Producer 并行请求外部 vLLM endpoint; -- standalone DSpark trainer 主动发现 ready key; -- global batch 到各 torchrun rank 的分片; -- 每个 optimizer step 后精确 `kv_clear`; -- EOS; -- standalone 无 Ray 的 TQ bootstrap; -- MooncakeStore 的实际运行验证。 - -PR #48 当前配置是: - -```yaml -transfer_queue: - enable: false - backend: - storage_backend: SimpleStorage - SimpleStorage: - total_storage_size: 100000 - num_data_storage_units: 8 -``` - -并且 bridge 注释针对 `TransferQueue==0.1.7`。旁边当前 verl 主线已使用 `0.1.8` 文案并包含 `MooncakeStore` 配置。当前机器没有安装 `transfer_queue` 包,因此正式实现前必须锁定版本并实机验证 API 签名,不能把 0.1.7 和 0.1.8 混用。 - -### 3.4 standalone 方案对 PR #48 的扩展 - -不再自行实现 `publish_ready/claim/lease/ack` HTTP 服务。第一版在 PR #48 KV 模式上补四个操作: - -```python -tq.kv_batch_put(...) # Producer 批量写 -tq.kv_list(...) # rank 0 列出 key + tag -tq.kv_batch_get(...) # 各 rank 并行读 -tq.kv_clear(...) # optimizer step 成功后删 -``` - -第一版不依赖 TQ `Sampler/StreamingDataLoader/get_meta`,因为 PR #48 并未使用或验证这些接口。等 KV 流程稳定后再升级为 TQ StreamingDataLoader。 - -### 3.5 PR #48 与 standalone 独立训练逐项映射 - -| PR #48 online RL | standalone drafter training | -|---|---| -| `SpecoTaskRunner` 调用 `tq.init(config)` | 新增独立 TQ owner/bootstrap 进程,或在确认安全后由 Producer owner 调用 `tq.init(config)` | -| SGLang/actor worker 是 Producer | `feature_producer.py` 是独立 Producer | -| rollout 过程中已经得到 hidden states | Producer 读取预生成 response,再请求外部 vLLM prefill | -| `make_sample_key(global_step, replica_rank, request_id)` | `make_replay_sample_key(dataset,row,tokens,target fingerprint)` | -| `put_sample()` / `kv_put()` | 继续复用同一写入模式,可扩展为 `kv_batch_put()` | -| key 通过 Ray driver/sample 传给 consumer | 没有 driver;rank 0 用 `kv_list()` 主动发现 ready keys | -| drafter worker `get_sample(key)` | 每个 torchrun rank `kv_batch_get(local_keys)` | -| `collect_rollout_features()` 恢复 sample | 转换为 `DraftFeatureSample` 后调用现有 `prepare_training_batch_from_samples()` | -| 多 TP/SP rank 可能读同一 key,所以不立即删除 | data-parallel rank 读取互不重叠的 key;全 rank step 成功后统一 `kv_clear(global_keys)` | -| TaskRunner 结束时 `tq.close()` | 每 step clear;输入 drain 完成后 owner 最后 `tq.close()` | - -standalone 需要新增的控制流是: - -```text -Producer DSpark rank 0 其他 ranks - │ │ │ - │ kv_put(sample key, fields, tag) │ │ - ├────────────────────────────────────▶│ │ - │ │ kv_list READY keys │ - │ │ │ - │ │ broadcast selected_keys ──▶│ - │ │ │ - │ │ kv_batch_get(local keys) │ kv_batch_get(local keys) - │ │ │ - │ ├──── DSpark synchronized step ────┤ - │ │ │ - │ │ kv_clear(global keys) │ -``` - -这个映射中,TQ 同时承担: - -- 大 tensor 存储/传输; -- key、tag 和 partition 的轻量索引。 - -但第一版 global batch 的选择仍由单个 DSpark job 的 rank 0 完成。这样最接近 PR #48 的 KV API,避免在同一次改造中再引入未经该 PR 验证的 Sampler/StreamingDataLoader。 - -### 3.6 PR #48 的 SGLang 路径:逐函数、逐对象完整流程 - -下面从一次生成请求开始,不省略中间层。 - -#### 阶段 1:SGLang完成生成并收集 hidden states - -执行进程:SGLang rollout server。 - -输入是一次 request 对应的 prompt 和生成配置。生成结束时,代码已经持有: - -```python -prompt_tensor: int64[prompt_len] -response_tensor: int64[response_len] -hidden_states: bf16[hidden_rows, hidden_dim] -hidden_positions: int64[hidden_rows] | None -target_logprobs: tensor | None -request_id: str -collection_global_steps: int -self.replica_rank: int -``` - -这些变量的语义: - -- `prompt_tensor`:输入 prompt token IDs; -- `response_tensor`:SGLang生成的 response token IDs; -- `hidden_states`:target model 指定层在部分 token positions 上的输出; -- `hidden_positions`:每个 hidden row 对应完整 `prompt+response` 序列中的哪个 token position; -- `target_logprobs`:可选的目标概率监督; -- `request_id`:当前 rollout request 标识; -- `replica_rank`:执行该 request 的 rollout replica。 - -SGLang 先构造完整 sample: - -```python -drafter_sample = { - "input_ids": torch.cat( - [prompt_tensor, response_tensor], dim=0 - ).unsqueeze(0), - "prompts": prompt_tensor.unsqueeze(0), - "responses": response_tensor.unsqueeze(0), - "hidden_states": hidden_states.unsqueeze(0).cpu(), - "hidden_positions": hidden_positions.unsqueeze(0).cpu(), - "target_logprobs": ( - target_logprobs.unsqueeze(0).cpu() - if target_logprobs is not None - else None - ), - "global_step": collection_global_steps, - "replica_rank": self.replica_rank, - # 还有 hidden window/alignment metadata -} -``` - -前导维 `1` 表示这是一个单样本 batch。`.cpu()` 表示跨进程传输前把大 tensor 放到 CPU 内存。 - -#### 阶段 2:PR #48 将大 fields 写入 TQ - -同一个 SGLang进程执行: - -```python -tq_key = make_sample_key( - collection_global_steps, - self.replica_rank, - request_id, -) - -tq_payload = { - "hidden_states": hidden_states.unsqueeze(0).cpu(), -} - -put_sample( - tq_key, - tq_payload, - tag={ - "global_step": collection_global_steps, - "replica_rank": self.replica_rank, - }, -) -``` - -调用展开后是: - -```python -tq.init() # 当前进程第一次使用时 -tq.kv_put( - key=tq_key, - partition_id="speco_drafter_features", - fields=tq_payload, - tag=tag, -) -``` - -效果是 TQ 中增加一行: - -```text -partition = speco_drafter_features -key = speco:42:1:req-007 -fields = {hidden_states: bf16[1, H, D], ...} -tag = {global_step: 42, replica_rank: 1} -``` - -`kv_put` 返回后,SGLang侧将旧 sample 改成: - -```python -drafter_sample["hidden_states_tq_key"] = tq_key -drafter_sample["hidden_states"] = None -``` - -此后该 Python 字典不再携带 hidden tensor,只携带定位它的 key。 - -#### 阶段 3:把 drafter_sample 放进 TokenOutput.extra_fields - -SGLang返回: - -```python -TokenOutput( - token_ids=token_ids, - log_probs=log_probs, - routed_experts=routed_experts, - extra_fields={ - "global_steps": collection_global_steps, - "drafter_sample": drafter_sample, - }, -) -``` - -此时 `TokenOutput` 中有两类输出: - -- 正常 rollout 输出:`token_ids/log_probs`; -- SpeCo 训练旁路输出:`extra_fields.drafter_sample`。 - -TQ payload 不在 `TokenOutput` 中,只有 `hidden_states_tq_key` 在其中。 - -#### 阶段 4:多个 TokenOutput 聚合成 gen_batch_output - -rollout/agent-loop 层将多个 request 的输出合并为批量 `DataProto`: - -```python -gen_batch_output.non_tensor_batch["drafter_sample"] -``` - -可能是 object array: - -```python -array([ - {"hidden_states_tq_key": "speco:42:0:req-A", ...}, - {"hidden_states_tq_key": "speco:42:1:req-B", ...}, -], dtype=object) -``` - -之所以进入 `non_tensor_batch`,是因为每条 sample 的 hidden window、Python metadata 和可选字段不一定具有统一 dense shape。 - -#### 阶段 5:driver 从 DataProto 取出 drafter samples - -`generate_sequences_with_speco()` 包装原 rollout 调用: - -```python -gen_batch_output = original_generate_sequences(...) -collected = self._speco_collect_generation_samples(gen_batch_output) -``` - -`_speco_collect_generation_samples()` 调用: - -```python -samples = pop_drafter_samples(gen_batch_output) -``` - -`pop_drafter_samples()` 实际执行: - -```python -non_tensor_batch = gen_batch_output.non_tensor_batch -samples_array = non_tensor_batch.pop("drafter_sample", None) -samples = normalize_drafter_samples(samples_array) -``` - -这里 `pop` 有两个作用: - -1. 取得 SpeCo drafter side-channel samples; -2. 从正常 PPO 的 `gen_batch_output` 中移除该旁路字段,避免后续 PPO batch 继续携带它。 - -`normalize_drafter_samples()` 将 dict、object array 或 list 统一成: - -```python -samples: list[dict] -``` - -#### 阶段 6:driver 按 replica_rank 分桶 - -假设有两个 rollout/drafter replicas,收到: - -```python -samples = [ - {"replica_rank": 1, "hidden_states_tq_key": "k1", ...}, - {"replica_rank": 0, "hidden_states_tq_key": "k2", ...}, - {"replica_rank": 1, "hidden_states_tq_key": "k3", ...}, -] -``` - -执行: - -```python -buckets = bucket_drafter_samples_by_replica( - samples, - num_replicas=2, -) -``` - -结果: - -```python -buckets = [ - [sample_k2], # bucket 0 - [sample_k1, sample_k3], # bucket 1 -] -``` - -分桶依据只有: - -```python -owner_rank = int(sample["replica_rank"]) -buckets[owner_rank].append(sample) -``` - -这一步没有读取 TQ,也没有处理 hidden tensor;只对小字典做路由。 - -#### 阶段 7:driver 通过 WorkerGroup RPC 分发 buckets - -driver 调用: - -```python -self._speco_set_drafter_global_step() -self._speco_collect_rollout_features_rpc( - "rollout", - buckets, -) -``` - -RPC 内部调用: - -```python -self.drafter_wg.collect_rollout_features(buckets) -``` - -因为 worker 方法注册了 `drafter_owner_route` dispatch,WorkerGroup 将 `buckets[0]` 发给 drafter DP replica 0,将 `buckets[1]` 发给 drafter DP replica 1。一个训练 replica 内的 SP ranks 根据 mesh dispatch 规则参与对应调用。 - -这里传输的对象仍是: - -```python -list[dict] -``` - -其中包含 tokens、position metadata 和 TQ key,不包含被置空的 hidden states。 - -#### 阶段 8:SpecoWorker 根据 key 从 TQ 恢复 tensor - -目标 worker 执行: - -```python -def collect_rollout_features(self, samples): - for sample in samples: - tq_key = sample.get("hidden_states_tq_key") - payload = get_sample(tq_key) - sample["hidden_states"] = payload["hidden_states"] -``` - -`get_sample()` 展开为: - -```python -tq.init() # 此 Consumer 进程第一次使用时 -result = tq.kv_batch_get( - keys=[tq_key], - partition_id="speco_drafter_features", -) -payload = _extract_value(result, tq_key) -payload = _tensordict_to_dict(payload) -``` - -现在 `sample` 再次包含: - -```python -{ - "input_ids": ..., - "hidden_positions": ..., - "hidden_states": bf16[1, hidden_rows, hidden_dim], - "hidden_states_tq_key": "...", -} -``` - -这与关闭 TQ 时 worker 收到的逻辑内容一致。 - -#### 阶段 9:worker 构造 DrafterBaseTrainer 所需 batch dict - -worker 先保留 token fields: - -```python -batch = { - "input_ids": sample["input_ids"], - "prompts": sample["prompts"], - "responses": sample["responses"], -} -``` - -再复制 hidden alignment metadata,例如: - -```python -batch["hidden_positions"] -batch["hidden_position_start"] -batch["hidden_position_end"] -batch["hidden_states_layout"] -batch["global_step"] -``` - -hidden tensor 单独作为参数: - -```python -self._store_rollout_sample( - batch=batch, - hidden_states=hidden, - target_logprobs=target_logprobs, -) -``` - -#### 阶段 10:样本进入在线 buffer 或落盘 - -`_store_rollout_sample()` 根据 training mode 分支: - -```python -if mode == "collect_only": - self._write_rollout_feature_sample( - batch, - hidden_states, - target_logprobs, - ) -else: - self.trainer.collect_online_data( - batch, - hidden_states, - target_logprobs, - ) -``` - -`collect_only` 会转换成 `DraftFeatureSample` 并写 `TorchShardFeatureStore`。online 模式则进入 `DrafterBaseTrainer.collect_online_data()`。 - -`collect_online_data()` 做: - -1. 将 `input_ids/hidden_states/positions/logprobs` 规范化到 CPU; -2. 按 batch 维拆成逐样本; -3. 根据 `hidden_positions` 校验 hidden row 与 token position; -4. 截取可训练窗口; -5. 构造内部 training item; -6. 保存到当前 step 的 `collected_data`,或在启用 data buffer 时保存到跨 step buffer。 - -因此 TQ get 完成不代表立即 optimizer step。它先恢复 online training sample,再进入现有数据准备逻辑。 - -#### 阶段 11:driver 在 actor update 周期触发 drafter 训练 - -driver 包装了 `update_actor()`: - -```python -should_train_drafter = ( - self._speco_should_attempt_drafter_train_this_step() -) - -actor_output = original_update_actor(...) - -if should_train_drafter: - drafter_trained, metrics = self._speco_train_drafter() -``` - -`_speco_train_drafter()` 再向 WorkerGroup 发: - -```python -self.drafter_wg.train_drafter() -``` - -每个 `SpecoWorker.train_drafter()`: - -1. 检查是否属于 drafter training group; -2. 检查 `training_interval_steps`; -3. 激活 drafter training model; -4. 循环 `train_steps_per_trigger` 次; -5. 每次调用 `self.trainer.training_step(global_step)`; -6. 成功时准备需要发布的 drafter state dict; -7. 清理训练期间临时状态。 - -`training_step()` 从刚才的 online `collected_data/DataBuffer` 组成 batch,执行 drafter forward、loss、backward 和 optimizer step。 - -所以 SGLang TQ 路径的最终效果是: - -```text -TQ 只替换 hidden tensor 跨进程传输 -→ sample 收集逻辑不变 -→ online buffer 不变 -→ drafter training trigger 不变 -→ loss/optimizer 不变 -``` - -### 3.7 PR #48 old-logprob 路径的完整差异 - -old-logprob 路径没有 `TokenOutput.extra_fields`。它从 PPO 的 actor old-logprob forward 开始。 - -#### 阶段 1:driver 构造 collect plan - -driver 根据 batch、collect interval 和 drafter owner 数量决定: - -```python -collect_mask: bool[batch] -hidden_positions: list/tensor per sample -owner_rank: int64[batch] -prompt_lens: int64[batch] -response_lens: int64[batch] -``` - -并把 `global_step` 等控制字段放入 old-logprob micro-batch。 - -#### 阶段 2:actor forward hook 选择 hidden rows - -actor worker 在 old-logprob forward 中捕获指定层 hidden states,根据 `collect_mask/position_mask` 只保留需要训练的样本和 token rows。 - -输出可以是 dense selected tensor,也可以是 sparse rows。随后 `_put_oldlogprob_hidden_refs()` 将同一个 owner 的多条 sample rows 拼成一个较大的 `hidden_chunk`。 - -#### 阶段 3:hidden chunk 写入 TQ - -改造前: - -```python -chunk_ref = ray.put(hidden_chunk) -``` - -PR #48: - -```python -tq_key = make_sample_key( - global_step, - owner, - f"chunk{chunk_index}", -) - -put_sample( - tq_key, - {"hidden": hidden_chunk}, - tag={"global_step": global_step, "owner": owner}, -) - -chunk_ref = tq_key -``` - -TQ fields: - -```python -{"hidden": bf16[total_owner_rows, hidden_dim]} -``` - -控制路径中的 chunk metadata: - -```python -chunk_info = { - "sample_indices": [0, 3, 5], - "starts": [0, 128, 384], - "lengths": [128, 256, 96], - "row_indices": [...], - "dtype": "bfloat16", - "shape": [480, hidden_dim], -} -``` - -`starts/lengths` 描述每条 sample 在共享 chunk 中对应的行区间。 - -#### 阶段 4:driver 将 chunk ref 还原为逐样本引用 - -driver 的 `_speco_collect_oldlogprob_features()` 读取: - -```python -chunk_refs = ["speco:42:0:chunk0", ...] -chunk_meta = [chunk_info, ...] -``` - -然后为每个 batch sample 构造: - -```python -sample["hidden_states_ref_chunks"] = [ - { - "ref": "speco:42:0:chunk0", - "chunk_start": 128, - "chunk_length": 256, - "chunk_row_indices": ..., - "dtype": "bfloat16", - "shape": [480, hidden_dim], - } -] -``` - -同时构造该 sample 的: - -```python -input_ids -prompts -responses -hidden_positions -hidden_states_layout -replica_rank=owner -``` - -再按 owner 放入 `buckets[owner]`,通过同一个 `collect_rollout_features()` RPC 发给 drafter worker。 - -#### 阶段 5:Consumer 获取共享 chunk 并切片 - -drafter worker 发现: - -```python -sample.get("hidden_states") is None -sample.get("hidden_states_ref_chunks") is not None -``` - -于是调用 `_resolve_hidden_state_chunks()`。对字符串 ref: - -```python -if ref.startswith("speco:"): - full_chunk = get_sample(ref)["hidden"] - full_chunk = _densify_tq_tensor(full_chunk) -``` - -然后按 sample metadata 取行: - -```python -sample_hidden = full_chunk[ - chunk_start : chunk_start + chunk_length -] -``` - -同一个 chunk 被多个 sample 复用,所以使用: - -```python -self._tq_chunk_cache[ref] = full_chunk -``` - -保证一次 `collect_rollout_features()` 中同一个 TQ key 只 get 一次。 - -得到逐样本 hidden 后,后续 `_store_rollout_sample → collect_online_data → train_drafter` 与 SGLang 路径相同。 - -### 3.8 PR #48 数据生命周期和清理 - -PR #48 的 TQ row 生命周期是: - -```text -TaskRunner tq.init(config) -→ Producer kv_put -→ key 经 Ray 控制路径传递 -→ 一个或多个 drafter TP/SP rank kv_batch_get -→ online drafter 收集/训练继续执行 -→ 整个 trainer.fit() 结束 -→ TaskRunner finally 调用 tq.close() -``` - -当前没有: - -```python -tq.kv_clear(key) -``` - -原因是同一 sample/key 可能被一个 drafter replica 的多个 rank 读取。若第一个 rank get 后删除,其余 rank 可能 get 失败。 - -因此 PR #48 采用任务级生命周期,而不是样本级确认和回收。这简化了并发正确性,但意味着长任务中的 TQ storage 占用可能持续增长;代码注释也把“leader 在 barrier 后精细 clear”留作后续工作。 - -### 3.9 PR #48 开启与关闭时的行为差异 - -`configure_transfer_queue()` 返回: - -```python -enabled_in_config and transfer_queue_importable -``` - -关闭时: - -```text -SGLang drafter_sample 继续内联 hidden_states -old-logprob 继续 ray.put(hidden_chunk) -Consumer 继续 ray.get/ref resolve -``` - -开启时: - -```text -SGLang hidden fields → TQ,sample 只带 key -old-logprob hidden chunk → TQ,ref 变成字符串 key -Consumer 根据 key 类型走 TQ get -``` - -如果配置开启但 `transfer_queue` 包未安装,bridge 会记录 warning,并让 `is_transfer_queue_enabled()` 返回 false,保留旧 Ray 路径。若已经进入 `put_sample/get_sample` 却发生 TQ 错误,则抛异常,不静默丢 hidden data。 - -## 第二部分:基于 PR #48 的 standalone drafter training 适配 - -从这一部分开始才讨论 `verl-SpeCo-ls` 的独立训练。下面的 `kv_list`、rank 0 选 global keys、逐 step `kv_clear`、独立 vLLM Producer 都是需要在本项目新增的逻辑,不是 PR #48 已有逻辑。 - -### 当前 standalone 基线 - -当前独立训练是: - -```text -draft_train_launcher -→ torch.distributed.run -→ 每个 rank 创建 DraftFeatureDataLoader -→ 每个 rank 在训练循环内同步 TargetFeatureReplayer.materialize() -→ vLLM/file hidden payload -→ prepare_training_batch_from_samples() -→ training_step_from_batch() -``` - -新方案要把 `TargetFeatureReplayer.materialize()` 从训练 rank 的同步取数路径移到独立 Producer,同时保持后两步训练接口不变。 - -## 4. TQ metadata 到底记录什么 - -### 4.1 Partition - -一次训练运行使用一个独立 partition: - -```python -partition_id = f"speco:{run_id}:dspark_train" -``` - -partition 用来隔离: - -- 不同训练 run; -- train 和 validation; -- 不同 target checkpoint 生成的 hidden states。 - -不能让两个 target 模型共用同一 partition,否则训练侧可能消费错误的 hidden states。 - -### 4.2 Sample key - -每条输入样本使用稳定 key: - -```python -sample_key = sha256( - dataset_id - + row_id - + prompt_token_ids - + response_token_ids - + tokenizer_fingerprint - + target_model_fingerprint - + target_layer_ids - + hidden_states_layout -).hexdigest() -``` - -稳定 key 用于: - -- vLLM HTTP 请求重试时不生成不同对象; -- Producer 重启后识别相同样本; -- 检查 hidden states 是否属于正确模型和正确层; -- TQ/Mooncake 清理时准确定位对象。 - -### 4.3 Fields 与 READY 约定 - -每个样本包含固定字段: - -```python -{ - "input_ids": int64[seq], - "loss_mask": float32[seq], - "position_ids": int64[seq], - "hidden_states": bf16[seq, aux_hidden_dim or aux_hidden_dim+hidden_size], -} -``` - -这里必须保持当前 `DraftFeatureSample` 的契约:`TargetFeatureReplayer._feature_from_vllm_payload()` 会把多个 aux layers flatten;当 layout 是 `dflash_aux_plus_last` 时,还会把 final hidden 拼到同一个 `hidden_states` tensor 尾部。`DSparkTrainerBackend.preprocess_individual_items()` 再根据 metadata 中的 `hidden_states_layout` 将 final hidden 切出来,生成训练 batch 的 `target_last_hidden_states`。 - -因此 TQ 不需要新增一个独立 `target_last_hidden_states` field。开启 DSpark L1 时,要求: - -```text -metadata.hidden_states_layout = dflash_aux_plus_last -hidden_states.shape[-1] = num_context_layers * hidden_size + hidden_size -``` - -完整 TQ native metadata API 可以做字段级 ready 判定,但 PR #48 当前走的是高层 KV API:一次 `kv_put` 把一个样本的多个 tensor fields 一起写入。因此第一版采用更直接的约定: - -```python -required_fields = [ - "input_ids", - "loss_mask", - "position_ids", - "hidden_states", -] -``` - -Producer 先在内存中验证所有必需字段,再进行一次 `kv_put`。只有 `kv_put` 成功返回,key 才会出现在 `kv_list` 结果中,并带有: - -```python -tag={ - "status": "ready", - "run_id": run_id, - "sample_id": sample_key, -} -``` - -Consumer 只选择 `status=ready` 且 `run_id` 匹配的 key。不要先 put `input_ids`、再单独 put `hidden_states`,否则 Consumer 可能观察到半成品。 - -### 4.4 Tags - -tags 是轻量 metadata,不放大 tensor: - -```python -tags = { - "sample_id": sample_key, - "source_row": row_id, - "seq_len": seq_len, - "payload_bytes": payload_bytes, - "target_model_fp": target_model_fingerprint, - "target_layers": "8,16,24", - "hidden_layout": "dflash_aux_plus_last", - "producer_status": "success", -} -``` - -tags 可用于过滤、监控、背压统计和错误排查。它不能替代 tensor shape/dtype 校验。 - -### 4.5 Run ID,而不是先依赖 task_name - -PR #48 的 `kv_put/kv_batch_get` 路径没有使用 `task_name` 或 Sampler consumption history。第一版按它的已验证接口,在 partition 和 tags 中放 `run_id`: - -```python -partition_id = "speco_drafter_features" -tag = { - "run_id": run_id, - "status": "ready", -} -``` - -不同 run 最好直接使用不同 partition: - -```python -partition_id = f"speco_drafter_features_{run_id}" -``` - -这样 job 重启和清理更简单。`task_name=dspark_train` 留到后续迁移 TQ native metadata/Sampler 时再使用。 - -### 4.6 standalone 中一条样本的完整对象形态 - -standalone 没有 PR #48 的轻量 `drafter_sample → Ray driver` 路径,因此需要让 TQ tag 承担“如何找到和解释 payload”的 metadata 作用。 - -#### Producer 读到的原始记录 - -```python -source_record = { - "dataset_id": "math-train", - "row_id": 12345, - "prompt": "...", - "response": "已经提前生成的 response", -} -``` - -#### Token replay 样本 - -分词和对齐后: - -```python -replay_sample = DraftReplaySample( - input_ids=int64[full_seq], - loss_mask=float32[full_seq], - position_ids=int64[full_seq], - feature_positions=int64[feature_rows], - draft_position_ids=int64[feature_rows], - metadata={ - "dataset_id": "math-train", - "row_id": 12345, - }, -) -``` - -这里的 `full_seq` 是 prompt 与预生成 response 拼接后的长度;`feature_positions` 指明哪些 token 位置最终进入 drafter 监督样本。 - -#### vLLM 返回的原始 hidden payload - -当前文件协议要求 safetensors 至少包含: - -```python -vllm_payload = { - "token_ids": int64[prefill_rows], - "hidden_states": bf16[prefill_rows, returned_layers, hidden_size], -} -``` - -这还不能直接给 DSpark。Producer 应复用当前: - -```python -TargetFeatureReplayer._feature_from_vllm_payload(...) -``` - -完成 token 校验、position 对齐、选层和 flatten。 - -#### Producer 最终得到的 DraftFeatureSample - -```python -feature = DraftFeatureSample( - algorithm="DSpark", - input_ids=int64[feature_rows], - loss_mask=float32[feature_rows], - position_ids=int64[feature_rows], - hidden_states=bf16[feature_rows, feature_hidden_dim], - metadata={ - "hidden_states_layout": "dflash_aux_plus_last", - "target_layer_ids": [8, 16, 24], - "target_model_path": "...", - "target_config_fingerprint": "...", - "feature_start": 128, - "feature_end": 640, - "sequence_length": 512, - }, -) -``` - -若有 3 个 aux layer、target hidden size 为 4096,并包含 final hidden: - -```text -feature_hidden_dim = 3 * 4096 + 4096 = 16384 -hidden_states.shape = [feature_rows, 16384] -``` - -前 `12288` 维是 aux/context hidden,最后 `4096` 维是 DSpark L1 所需 final hidden。现有 DSpark backend 会根据 `dflash_aux_plus_last` 自动拆分。 - -#### 写入 TQ 的 data fields - -第一版建议一个 key 对应一个已经规范化完成的 `DraftFeatureSample`: - -```python -tq_fields = { - "input_ids": feature.input_ids.cpu(), - "loss_mask": feature.loss_mask.cpu(), - "position_ids": feature.position_ids.cpu(), - "hidden_states": feature.hidden_states.cpu(), -} -``` - -这里只放 tensor,因为 PR #48 的 `put_sample()` 会过滤非 tensor: - -```python -payload = { - key: value - for key, value in tensor_dict.items() - if torch.is_tensor(value) -} -``` - -#### 写入 TQ 的 tag metadata - -```python -tq_tag = { - "run_id": "run-20260818-001", - "status": "ready", - "sample_id": sample_key, - "sequence_no": 12345, - "algorithm": "DSpark", - "hidden_states_layout": "dflash_aux_plus_last", - "target_model_fingerprint": "sha256:...", - "target_layer_ids": "8,16,24", - "feature_rows": 512, - "hidden_dim": 16384, - "payload_bytes": 16777216, -} -``` - -tag 中只放 TQ 版本支持序列化的小标量/字符串。列表等复杂对象可以编码为稳定字符串或 JSON。`payload_bytes` 用于背压统计。 - -#### TQ 中逻辑上保存的 row - -```text -partition: speco_drafter_features_run-20260818-001 -key: 86a4...ef2 - -fields: - input_ids → int64[512] - loss_mask → float32[512] - position_ids → int64[512] - hidden_states → bf16[512, 16384] - -tag: - status → ready - sequence_no → 12345 - hidden_layout → dflash_aux_plus_last - target_model_fp → sha256:... -``` - -#### Consumer 恢复出的对象 - -rank 0 通过 `kv_list` 同时取得 key 和 tag;每个 rank 用 local keys 调用 `kv_batch_get` 取得 fields,然后组合: - -```python -feature = DraftFeatureSample( - algorithm=tag["algorithm"], - input_ids=densify(fields["input_ids"]).reshape(-1), - loss_mask=densify(fields["loss_mask"]).reshape(-1), - position_ids=densify(fields["position_ids"]).reshape(-1), - hidden_states=densify(fields["hidden_states"]), - metadata={ - "hidden_states_layout": tag["hidden_states_layout"], - "target_model_fingerprint": tag["target_model_fingerprint"], - }, -) - -feature.validate(strict=True) -``` - -这样传给: - -```python -trainer.prepare_training_batch_from_samples([feature, ...]) -``` - -的数据结构,与现有文件 feature store 读出的 `DraftFeatureSample` 一致。也就是说,TQ 替换的是存取介质和调度方式,不改变 DSpark backend 的样本契约。 - -## 5. 新的整体架构 - -```text - 小 metadata - ┌──────────────────────────┐ - │ TransferQueueController │ - │ KV metadata / key / tags │ - │ partition / storage map │ - └────────────┬─────────────┘ - │ -JSONL/token replay │ - │ │ - ▼ │ -Feature Producer │ - ├─ tokenizer/window │ - ├─ asyncio bounded concurrency │ - ├─ vLLM endpoint pool │ - ├─ validate/pack │ - └─ TQ put ─────────────────────┤ - ▼ - TQ Mooncake backend - hidden-state tensors - │ - ┌───────────────────┼───────────────────┐ - ▼ ▼ ▼ - DSpark rank 0 DSpark rank 1 DSpark rank N - TQ get TQ get TQ get - └───────────────────┼───────────────────┘ - ▼ - synchronized optimizer step - │ - ▼ - TQ clear after success -``` - -大 tensor 的路径是: - -```text -vLLM/Producer memory → TQ Mooncake backend → each training rank -``` - -不会走: - -```text -Mooncake → rank 0 → rank 1/2/3 -``` - -rank 0 最多只广播 `kv_list` 得到的 key 字符串列表;tags 只在选 batch 时由 rank 0 使用。 - -### 5.1 standalone 每一步为什么能实现推理和训练异步 - -#### 步骤 A:Producer 独立推进输入 cursor - -Producer 自己维护: - -```python -reader_cursor = 12346 -``` - -它不等待 Trainer 请求样本。只要 TQ ready bytes 没超过背压上限,就继续读取文件并创建 vLLM task。 - -效果是 Producer 的执行进度与 `optimizer_step` 解耦: - -```text -Producer sequence_no: 1200,1201,1202,... -Trainer optimizer_step: 87 -``` - -两者通过 TQ 中的 ready rows 衔接,不互相直接调用。 - -#### 步骤 B:并发 vLLM task 完成顺序可以乱序 - -例如 Producer 同时提交: - -```text -sequence_no 100 → endpoint 0 -sequence_no 101 → endpoint 1 -sequence_no 102 → endpoint 0 -``` - -完成顺序可能是: - -```text -101 → 100 → 102 -``` - -每个 task 完成后独立执行 `kv_put`,所以慢请求不会阻塞已经完成的请求写入。tag 中的 `sequence_no` 保留原数据顺序。 - -#### 步骤 C:`kv_put` 成功是 READY 可见性的边界 - -Producer 在调用前已经得到完整 `DraftFeatureSample`。一次 `kv_put` 写入该 sample 的全部 tensor fields,并在 tag 中标记 `status=ready`。 - -因此 Consumer 的判断规则是: - -```text -kv_list 能列出该 key -且 tag.run_id 匹配 -且 tag.status == ready -→ 可以尝试 kv_batch_get -``` - -Consumer 仍需对取回 fields 做完整性校验;tag 是调度 metadata,不是正确性证明。 - -#### 步骤 D:rank 0 只负责选 key - -rank 0 执行: - -```python -entries = list_ready_keys() -selected = sorted(entries, key=sequence_no)[:global_batch_size] -``` - -这一步处理的数据只是: - -```python -[ - {"key": "k100", "sequence_no": 100, ...}, - {"key": "k101", "sequence_no": 101, ...}, -] -``` - -不包含 `[seq, hidden_dim]` hidden tensor,所以 rank 0 不成为大数据中转瓶颈。 - -#### 步骤 E:广播保证所有 rank 对同一个 global step 达成一致 - -所有 rank 调用同一次: - -```python -dist.broadcast_object_list(holder, src=0) -``` - -广播结束后,每个 rank 看到完全相同的 global key list。然后按确定性区间切分: - -```text -rank 0: keys[0:per_rank] -rank 1: keys[per_rank:2*per_rank] -... -``` - -这样不会出现 rank 0 训练 batch A、rank 1 训练 batch B 数量不同,或者某个 rank 没进入 backward 的情况。 - -#### 步骤 F:各 rank 直接读取 Mooncake 后端 - -每个 rank 执行: - -```python -tq.kv_batch_get(keys=local_keys, partition_id=partition_id) -``` - -TQ 根据 key 找到 storage backend 中的数据。使用 MooncakeStore 时,大 tensor 数据路径是 Mooncake → 本 rank;rank 0 不读取其他 rank 的 local payload。 - -因此: - -```text -控制面:rank 0 → broadcast small keys -数据面:Mooncake → each rank directly -``` - -#### 步骤 G:恢复现有 DraftFeatureSample 契约 - -每个 rank 将 TQ fields 与 tag 合并、densify、校验,得到 `list[DraftFeatureSample]`。从这里开始继续执行项目现有代码: - -```python -batch = trainer.prepare_training_batch_from_samples( - materialized_samples, - step=optimizer_step, -) - -ok = await trainer.training_step_from_batch( - batch, - optimizer_step, -) -``` - -所以 TQ 不进入 DSpark model/backend 内部,训练数学逻辑不变。 - -#### 步骤 H:全 rank 成功以后才能清理 - -每个 rank 的 `ok` 通过现有 `_all_ranks_true()` 聚合: - -```text -rank 0 ok = true -rank 1 ok = true -rank 2 ok = true -rank 3 ok = true -→ global_ok = true -``` - -只有此时 rank 0 执行: - -```python -tq.kv_clear(keys=global_keys, partition_id=partition_id) -``` - -这样保证被删除的数据已经参与完成的 optimizer step。若任何 rank get/OOM/backward 失败,不执行 clear,便于作业失败后的诊断或恢复。 - -#### 步骤 I:异步重叠如何形成 - -时间线上: - -```text -时间 ─────────────────────────────────────────▶ - -Producer: vLLM(batch N+1) ─ put ─ vLLM(batch N+2) ─ put -Trainer: get(batch N) ─ train(batch N) ─ get/train(batch N+1) -``` - -Producer 和 Trainer 是不同进程,互相不调用;TQ ready rows 是缓冲区。因此 vLLM prefill、网络传输和 DSpark GPU 训练可以重叠。背压只在缓冲区达到容量上限时暂停 Producer。 - -## 6. Producer:读取预生成 response 并并行请求 vLLM - -### 6.1 输入处理 - -Producer 从现有 JSONL/token replay 数据源读取: - -```python -sample = { - "row_id": "12345", - "prompt": "...", - "response": "提前生成好的文本", -} -``` - -构造: - -```python -prompt_ids = tokenizer.encode(sample["prompt"]) -response_ids = tokenizer.encode(sample["response"]) -input_ids = prompt_ids + response_ids -``` - -同时产生: - -```python -loss_mask -position_ids -feature_positions -sample_key -``` - -### 6.2 有界并发 - -不能按样本串行请求: - -```python -for sample in samples: - result = request_vllm(sample) -``` - -改成: - -```python -async def run_producer(samples): - semaphore = asyncio.Semaphore(max_inflight_requests) - - async def run_one(sample): - async with semaphore: - result = await vllm_pool.prefill(sample) - feature = validate_and_pack(sample, result) - await tq_transport.put(feature) - - async with asyncio.TaskGroup() as group: - for sample in samples: - group.create_task(run_one(sample)) -``` - -`max_inflight_requests` 是 Producer 同时未完成的请求数量,不是 vLLM batch size。vLLM 服务端仍会对同时到达的请求做自己的 continuous batching。 - -### 6.3 多 endpoint - -多个 endpoint 例如: - -```yaml -vllm_endpoints: - - http://node0:8000/v1 - - http://node1:8000/v1 - - http://node2:8000/v1 -``` - -调度器维护每个 endpoint 的 inflight 数: - -```python -endpoint = min( - endpoints, - key=lambda item: item.inflight, -) -``` - -请求前 `inflight += 1`,在 `finally` 中 `inflight -= 1`。失败只重试对应样本,不阻塞全部 Producer。 - -### 6.4 当前 vLLM 文件桥接 - -当前客户端协议期望: - -```python -response.kv_transfer_params["hidden_states_path"] -``` - -所以第一阶段仍然是: - -```text -vLLM 写临时 safetensors -→ Producer load_file -→ 校验 token_ids/hidden_states -→ TQ put 到 Mooncake backend -→ TQ put 成功后删除临时文件 -``` - -删除必须发生在 TQ put 成功之后: - -```python -path = request_vllm_hidden_file(sample) -try: - feature = load_and_validate(path) - await tq_transport.put(feature) -finally: - if put_succeeded: - Path(path).unlink(missing_ok=True) -``` - -### 6.5 目标版本:vLLM 直接写 TQ/Mooncake - -目标响应可改成: - -```json -{ - "kv_transfer_params": { - "backend": "transfer_queue", - "partition_id": "speco:run-1:dspark_train", - "sample_key": "abc123" - } -} -``` - -服务端顺序必须是: - -```text -prefill -→ 捕获指定层 hidden states -→ TQ/Mooncake put 完成 -→ 返回 HTTP success 和 sample key -``` - -这需要定制 vLLM exporter;当前 `verl-SpeCo-ls` 中没有服务端 writer 实现。 - -## 7. 按 PR #48 扩展 TQ bridge - -不要重新发明一套 transport。将 PR #48 的 bridge 设计移植到本项目并增加 standalone 所需方法: - -```python -class StandaloneTQTransport: - def put_sample(self, key, tensor_dict, tag): ... - def list_ready_keys(self, run_id): ... - def get_samples(self, keys, fields=None): ... - def clear_samples(self, keys): ... - def put_control(self, key, tag): ... - def close(self): ... -``` - -写入延续 PR #48 的真实形式: - -```python -tq.kv_put( - key=key, - partition_id=partition_id, - fields={ - "input_ids": feature.input_ids.cpu(), - "loss_mask": feature.loss_mask.cpu(), - "position_ids": feature.position_ids.cpu(), - "hidden_states": feature.hidden_states.cpu(), - }, - tag={ - "run_id": run_id, - "status": "ready", - "sequence_no": sequence_no, - "payload_bytes": payload_bytes, - }, -) -``` - -批量读取延续 PR #48 的 `kv_batch_get`: - -```python -result = tq.kv_batch_get( - keys=keys, - partition_id=partition_id, - fields=required_fields, # 0.1.7 是否支持该参数需实机确认 -) -``` - -新增发现和清理: - -```python -items = tq.kv_list(partition_id=partition_id) -tq.kv_clear(keys=keys, partition_id=partition_id) -``` - -这里的 `kv_list/kv_clear` 参数名必须根据锁定的 TQ 版本验证。当前参考仓库只实际调用了 `kv_put/kv_batch_get`,没有为这两个接口提供运行证据。 - -读取结果继续复用 PR #48 的两个适配函数: - -```python -value = _extract_value(result, key) -row = _tensordict_to_dict(value) -row["hidden_states"] = _densify_tq_tensor(row["hidden_states"]) -``` - -## 8. DSpark 多 rank 如何消费 - -### 8.1 第一版:rank 0 用 kv_list 发现 READY keys - -PR #48 中 key 由 Ray driver 传给 drafter worker;standalone 没有这条控制路径,所以 rank 0 需要主动列出 key: - -```python -rank = dist.get_rank() -world_size = dist.get_world_size() -global_batch_size = batch_size_per_gpu * world_size - -if rank == 0: - entries = tq_transport.list_ready_keys(run_id=run_id) - entries.sort(key=lambda x: (x.tag["sequence_no"], x.key)) - selected_keys = [x.key for x in entries[:global_batch_size]] -else: - selected_keys = None - -holder = [selected_keys] -dist.broadcast_object_list(holder, src=0) -selected_keys = holder[0] -``` - -`sequence_no` 是输入文件顺序。它让多个 Producer 并发完成顺序不同的情况下,Trainer 仍能确定性地组成 batch。 - -rank 0 广播的是字符串 key 列表,不是 hidden-state tensor。 - -### 8.2 各 rank 切自己的 keys - -例如 global batch keys: - -```text -[s0, s1, s2, s3, s4, s5, s6, s7] -``` - -world size 为 4、每卡 batch size 为 2: - -```text -rank 0 → [s0, s1] -rank 1 → [s2, s3] -rank 2 → [s4, s5] -rank 3 → [s6, s7] -``` - -代码: - -```python -def shard_keys(keys, rank, world_size): - assert len(keys) % world_size == 0 - per_rank = len(keys) // world_size - start = rank * per_rank - end = start + per_rank - return keys[start:end] -``` - -### 8.3 每个 rank 并行 get - -所有进程执行: - -```python -local_keys = shard_keys( - selected_keys, - rank=rank, - world_size=world_size, -) - -local_payloads = tq_transport.get_samples(local_keys) -``` - -数据路径: - -```text -rank 0 ← Mooncake(s0,s1) -rank 1 ← Mooncake(s2,s3) -rank 2 ← Mooncake(s4,s5) -rank 3 ← Mooncake(s6,s7) -``` - -不是 rank 0 get 全部后再 scatter。 - -### 8.4 转成当前训练格式 - -TQ 返回的数据需要先按 PR #48 的规则解包、densify,再转换成现有 `DraftFeatureSample`: - -```python -def tq_row_to_feature(row, tag): - return DraftFeatureSample( - algorithm="DSpark", - input_ids=row["input_ids"], - loss_mask=row["loss_mask"], - position_ids=row["position_ids"], - hidden_states=row["hidden_states"], - metadata={ - "hidden_states_layout": tag["hidden_states_layout"], - **row.get("metadata", {}), - }, - ) -``` - -`hidden_states_layout` 不能只放在无法恢复的临时 Python 对象里;应随 TQ tag 保存。rank 0 从 `kv_list` 取得 key/tag 后,需要把每个 local key 对应的 tag 一起交给本 rank 的转换逻辑。Consumer 必须把 layout 放回 `DraftFeatureSample.metadata`。DSpark L1 所需的 `target_last_hidden_states` 随后由现有 backend 从拼接的 `hidden_states` 中切出,transport 层不计算 loss,也不拆 layout。 - -## 9. 修改当前训练循环 - -在 `run_standalone_draft_training()` 中增加数据源分支: - -```python -feature_store_type = str(feature_store_cfg.get("type", "torch_shard")) - -if feature_store_type == "transfer_queue": - tq_stream = build_transfer_queue_stream( - config=config, - rank=rank, - world_size=world_size, - ) - store = None - loader = None - feature_replayer = None -else: - store = build_feature_store_from_config( - feature_store_cfg, - read_only=True, - ) - loader = DraftFeatureDataLoader(...) -``` - -流式训练循环: - -```python -while successful_steps < max_steps: - global_keys, materialized_samples = tq_stream.next_local_batch() - - batch = trainer.prepare_training_batch_from_samples( - materialized_samples, - step=optimizer_step, - ) - - has_batch = batch is not None - if not _all_ranks_true(has_batch, trainer.runtime_device): - raise RuntimeError("at least one rank failed to fetch its TQ batch") - - ok = await trainer.training_step_from_batch( - batch, - optimizer_step, - ) - - if not _all_ranks_true(ok, trainer.runtime_device): - raise RuntimeError("DSpark step failed on at least one rank") - - dist.barrier() - if rank == 0: - tq_stream.clear_global_batch(global_keys) - dist.barrier() -``` - -TQ 接在数据输入层,不放进 `DSparkTrainerBackend`。Backend 只负责模型、forward、loss、backward 和 optimizer。 - -## 10. READY key、inflight key 和训练提交 - -### 10.1 Ready - -在本方案中 ready 表示: - -> `dspark_train` 所需的所有 fields 已经成功写入 TQ storage,训练侧可以读取。 - -第一版不是通过 PR #48 尚未验证的 Sampler 自动判定,而是 Producer 完整校验 payload 后一次 `kv_put`,并设置 `tag.status=ready`。rank 0 的 `kv_list` 只选择这些 key。 - -### 10.2 Inflight key - -rank 0 选出一个 global batch 后,在本地保存: - -```python -inflight_global_keys = selected_keys -``` - -其他 step 不应再次选中这批 key。因为 standalone 只有一个同步 DSpark job,第一版由 rank 0 进程内 `set` 排除 inflight key 即可: - -```python -ready = [x for x in listed if x.key not in inflight_keys] -``` - -若以后允许多个独立 Trainer job 同时消费同一 partition,进程内 set 就不够,届时必须使用 TQ Controller/Sampler 的原子消费分配能力。 - -### 10.3 Optimizer committed - -optimizer committed 表示所有 DSpark rank 已经完成: - -```text -forward → backward → gradient synchronization → optimizer.step -``` - -它比 `kv_list` 发现 key、甚至 `kv_batch_get` 取出 tensor 都更晚。 - -第一版推荐简单语义: - -```text -TQ 负责 key/tag 和 tensor 传输 -rank 0 负责单 Trainer job 的 batch 选择和 inflight set -训练失败 → 整个作业 fail-fast -训练成功 → kv_clear payload,并从 inflight set 移除 -恢复 → 从最近 checkpoint + 输入 cursor 重新启动 -``` - -这样不需要重新实现 lease/ack 状态机,也没有夸大 PR #48 当前尚未使用的 Sampler 能力。 - -## 11. 为什么训练完一个 step 才清理 - -不能在 `kv_batch_get()` 后立即 clear: - -```text -get 成功 -→ clear -→ forward OOM -→ 数据已不存在,无法重试 -``` - -正确顺序: - -```text -rank 0..N get -→ 所有 rank 确认 batch 有效 -→ training_step_from_batch -→ _all_ranks_true(ok) -→ rank 0 kv_clear global keys -``` - -当前 `_all_ranks_true()` 已经是项目现有的跨 rank 同步工具,可以继续复用。 - -## 12. 背压 - -背压表示 Producer 生成速度高于 Trainer 消费速度时,Producer 必须暂停继续提交,以免 Mooncake 内存无限增长。 - -建议限制: - -```yaml -max_vllm_inflight_requests: 32 -max_pending_put_bytes: 8589934592 -max_tq_ready_samples: 256 -max_tq_ready_bytes: 68719476736 -``` - -Producer 在 tags 中写: - -```python -{"payload_bytes": payload_bytes} -``` - -周期性通过 TQ list/metadata 统计当前 partition 尚未消费的数据量: - -```python -while ready_bytes >= max_tq_ready_bytes: - await asyncio.sleep(backpressure_poll_interval) -``` - -如果锁定的 TQ 版本提供容量/ready 统计接口,应直接使用,避免全量扫描 keys。 - -## 13. Stable ID、幂等和孤儿数据 - -### 13.1 幂等 - -幂等表示同一操作重复执行,最终逻辑结果仍只有一份。 - -Producer 对同一样本重试时必须使用相同 `sample_key`: - -```python -await tq.put(key="abc123", ...) -await tq.put(key="abc123", ...) -``` - -不能每次生成随机 key: - -```text -abc123-retry-1 -abc123-retry-2 -``` - -否则一个输入可能训练多次并持续占用 Mooncake。 - -### 13.2 孤儿数据 - -孤儿数据表示 tensor 已经写进 storage,但由于 Producer 崩溃或 metadata 更新失败,没有进入正常消费路径。 - -使用 TQ 后,metadata 和 storage 由同一套系统管理,可以减少“Mooncake有对象、自研 Coordinator 没记录”的双系统窗口,但仍要配置 TTL/partition cleanup: - -```text -训练正常结束 → clear partition -训练异常退出 → 下次启动检查旧 partition -超过 TTL → 清理未消费数据 -``` - -## 14. EOS 和 drop-last - -EOS 表示 Producer 已经读完输入,并且所有 vLLM/TQ put 都已完成。 - -TQ 中需要一种结束条件,具体使用 Controller API、特殊 metadata 或单独的运行状态记录取决于锁定版本。不能把“当前暂时没有 ready sample”当成 EOS,因为 Producer 可能仍在请求 vLLM。 - -最后不足一个 global batch 时: - -```python -global_batch_size = batch_size_per_gpu * world_size -``` - -第一版使用 `drop_last=true`,避免某些 rank 有数据、某些 rank 没数据导致分布式训练不同步。 - -结束条件: - -```text -producer_done == true -and ready_samples < global_batch_size -and inflight_requests == 0 -and pending_puts == 0 -``` - -## 15. 双缓冲预取 - -训练 batch N 时,CPU 后台线程预取 batch N+1: - -```python -next_future = executor.submit(tq_stream.next_local_batch) - -current_batch = first_batch -while current_batch is not None: - next_batch = next_future.result() - next_future = executor.submit(tq_stream.next_local_batch) - - train(current_batch) - current_batch = next_batch -``` - -实际顺序应调整为避免等待 future 后才训练。推荐: - -```python -current = tq_stream.next_local_batch() - -while current is not None: - future = executor.submit(tq_stream.next_local_batch) - train_and_clear(current) - current = future.result() -``` - -第一版只预取一个 global batch,避免 Trainer 崩溃时大量数据已被采样但未训练。 - -如果 Mooncake/TQ Python get 是阻塞函数,使用专用 `ThreadPoolExecutor`,不要阻塞 Producer 的 asyncio event loop。 - -## 16. 建议代码结构 - -```text -verl_speco/ - trainer/ - tq_transport.py # TQ client、put/get/meta/clear 封装 - tq_feature_stream.py # kv_list、global keys、rank shard、decode/prefetch - feature_producer.py # JSONL → 并发 vLLM → TQ - draft_training_loop.py # 增加 transfer_queue 数据源分支 - target_feature_replay.py # 复用 vLLM payload 校验/转换逻辑 -``` - -不要新增: - -```text -coordinator.py -coordinator_client.py -``` - -建议抽象: - -```python -class StreamingFeatureSource(Protocol): - def next_local_batch(self) -> tuple[list[str], list[DraftFeatureSample]]: ... - def clear_global_batch(self, keys: list[str]) -> None: ... - def close(self) -> None: ... -``` - -这样训练循环不依赖 TQ 的具体类型。 - -## 17. 配置草案 - -```yaml -actor_rollout_ref: - rollout: - drafter: - training: - mode: offline - backend: dspark - batch_size_per_gpu: 2 - max_steps: 1000 - - feature_store: - type: transfer_queue - partition_id: speco_drafter_features_${run_id} - drop_last: true - prefetch_steps: 1 - - transfer_queue: - # 与 PR #48 的配置层级和 init 方式保持一致。 - enable: true - package_version: 0.1.8 # 最终以实测版本为准 - backend: - storage_backend: MooncakeStore - MooncakeStore: - auto_init: false - metadata_server: localhost:50123 - master_server_address: localhost:50124 - local_hostname: localhost - protocol: tcp - global_segment_size: 4294967296 - local_buffer_size: 1073741824 - device_name: "" - - required_fields: - - input_ids - - loss_mask - - position_ids - - hidden_states - - producer: - input_path: /path/to/generated_responses.jsonl - vllm_endpoints: - - http://node0:8000/v1 - - http://node1:8000/v1 - max_inflight_requests: 32 - max_pending_put_bytes: 8589934592 - max_ready_samples: 256 - max_ready_bytes: 68719476736 -``` - -当前 examples 中的: - -```bash -transfer_queue.enable=False -``` - -属于 verl RL 主入口配置,当前 standalone `draft_train_launcher` 不读取它。不能只改成 `True`;必须实现上述 `feature_store.type=transfer_queue` 分支。 - -## 18. 启动顺序 - -逻辑顺序: - -```text -1. 启动 Mooncake metadata/master 服务; -2. 启动 standalone TQ owner 进程,调用一次 `tq.init(tq_config)` 创建/连接 TQ Controller 和 MooncakeStore backend; -3. 启动一个或多个定制 vLLM server -4. 启动 Feature Producer -5. Producer 和所有 Trainer rank 调用无参 `tq.init()`,连接 owner 创建的同一套 TQ; -6. 启动 verl_speco.draft_train_launcher -7. torchrun 启动所有 DSpark rank -8. 各 rank 连接 TQ -9. rank 0 用 `kv_list` 选择 global keys,各 rank 并行 `kv_batch_get`/train; -10. 输入耗尽后 Producer 发布 done 状态 -11. Trainer drain 完整 global batches 后退出 -12. 清理 partition,停止 Producer、vLLM、TQ、Mooncake -``` - -PR #48 的 `init_transfer_queue()` 在 Ray `SpecoTaskRunner` 中调用 `tq.init(config)`,worker 的无参 `tq.init()`依靠同一个 Ray 集群发现 named Controller;这部分不能原样复制到 standalone torchrun。 - -本项目要求一开始不使用 Ray,因此 Phase 0 必须先证明锁定的 TQ 版本支持独立 owner/controller 进程,以及 Producer/torchrun rank 如何获得连接信息。若 TQ 0.1.7/0.1.8 实际只能通过 Ray named actor 完成发现,那么有两个选择: - -1. 接受仅用 Ray 承载 TQ 控制面的最小方案; -2. 给 TQ 增加或使用其已有的独立 ZMQ/server-info bootstrap。 - -在这项验证完成前,文档不能声称 PR #48 已经提供“无 Ray TQ standalone 启动”。 - -## 19. 故障处理 - -### vLLM 请求失败 - -- 对单个 sample 按稳定 key 重试; -- 指数退避; -- 超过次数记录失败,并根据配置 fail-fast 或跳过; -- 不写不完整 TQ fields。 - -### vLLM 文件读取成功,但 TQ put 失败 - -- 暂时保留临时文件; -- 重试 TQ put; -- put 成功后再删除; -- 不把样本视为 ready。 - -### 某个训练 rank get 失败 - -- 该 rank 报告 `local_ok=false`; -- `_all_ranks_true()` 使全部 rank 得到一致失败结果; -- 第一版整个训练 fail-fast; -- 不 clear global batch。 - -### OOM/optimizer step 失败 - -- 不 clear; -- 所有 rank 一致退出; -- 从最近训练 checkpoint 恢复; -- 根据 TQ 消费提交语义决定是否重放当前 batch。 - -### clear 失败 - -- optimizer 已成功,不能再次训练这批; -- 将 batch keys 写入本地小型 `gc_pending` 日志; -- 后台重试 clear; -- checkpoint 保存最近 committed sample IDs,避免恢复时重复消费。 - -## 20. 观测指标 - -Producer: - -```text -producer/vllm_inflight -producer/vllm_requests_per_sec -producer/vllm_prefill_tokens_per_sec -producer/vllm_p50_latency -producer/vllm_p95_latency -producer/tq_put_bytes_per_sec -producer/tq_put_failures -producer/pending_put_bytes -``` - -TQ/Mooncake: - -```text -tq/ready_samples -tq/ready_bytes -tq/consumed_samples -tq/storage_bytes -tq/clear_failures -mooncake/put_bandwidth -mooncake/get_bandwidth -``` - -Trainer: - -```text -trainer/tq_wait_seconds -trainer/tq_get_seconds -trainer/tq_get_bytes_per_sec -trainer/decode_seconds -trainer/h2d_seconds -trainer/step_seconds -trainer/data_stall_ratio -trainer/successful_steps -``` - -## 21. 实施阶段 - -### Phase 0:锁定依赖和契约 - -- 从 PR #48 的 `TransferQueue==0.1.7` 起验证,同时对比当前 verl 文档使用的 0.1.8; -- 实测 `tq.init(config)`、worker `tq.init()`、`kv_put`、`kv_batch_get`、`kv_list`、`kv_clear`; -- 验证该版本是否支持无 Ray owner/controller 部署以及连接信息传递; -- 先用 PR #48 已配置的 `SimpleStorage` 做最小闭环; -- 再把 backend 换成 `MooncakeStore`,验证 tcp,最后再验证 rdma; -- 固定 fields、tags、partition 和 `run_id/sequence_no/status`; -- 写 fake TQ 单元测试。 - -### Phase 1:文件桥接 + TQ KV 模式 - -- 新增独立 Producer; -- 32 个有界并发 vLLM 请求; -- 读取 vLLM 临时 safetensors; -- TQ put 成功后删除文件; -- standalone trainer 由 rank 0 `kv_list` 并广播 global keys; -- 各 rank 并行 `kv_batch_get`; -- 复用 PR #48 的返回值解包、NestedTensor densify 和 per-step cache; -- optimizer 成功后 `kv_clear`。 - -验收:连续训练 1000 step,临时文件数量、TQ ready bytes 和 Mooncake占用均保持有界。 - -### Phase 2:双缓冲与多 endpoint - -- 增加多 endpoint 最少 inflight 调度; -- 增加一个 global batch 预取; -- 动态背压; -- 注入单 rank get 失败,确认所有 rank 一致退出而非死锁。 - -### Phase 3:vLLM 直接写 TQ/Mooncake - -- 修改外部定制 vLLM exporter; -- 去掉 `hidden_states_path` 临时文件; -- HTTP 响应返回 partition/sample key; -- 验证 HTTP 重试的幂等性。 - -### Phase 4:可选升级到 TQ StreamingDataLoader - -- 在当前保守方案稳定后再引入 RankAwareSampler; -- 让每个 rank 自动取得 local micro-batch; -- 去掉 rank 0 手工 key-list 广播; -- 验证与 torchrun/DSpark 的 global step 对齐。 - -## 22. 最终推荐 - -针对当前 `verl-SpeCo-ls`,推荐的第一版不是 AngelSpec 的完整架构,也不是单独写 Coordinator,而是: - -```text -当前预生成 response 文件 -→ 独立 asyncio Producer -→ 并行访问多个 vLLM endpoint -→ 读取并校验临时 hidden-state 文件 -→ TransferQueue put -→ Mooncake storage backend -→ rank 0 kv_list 获取 READY global keys -→ broadcast key list -→ 各 DSpark rank 并行 kv_batch_get -→ 现有 prepare_training_batch_from_samples() -→ 现有 training_step_from_batch() -→ 全 rank 成功 -→ TQ clear -``` - -这套方案保留当前 standalone DSpark 训练主体,只替换 `DraftFeatureDataLoader + TargetFeatureReplayer.materialize()` 所在的数据输入路径。它真正建立在 PR #48 已实现的 KV transport 之上,而不是假设 PR #48 已经实现了 standalone Sampler/StreamingDataLoader。 - -## 23. 参考 - -- verl TransferQueue: -- TransferQueue: -- Mooncake Store: -- 本地只读参考:`../verl-SpeCo/verl_speco/integration/transferqueue_bridge.py`(PR #48) -- 本地只读参考:`../verl-SpeCo/docs/transferqueue_integration_plan.md`(PR #48) diff --git a/docs/draft_feature_sample_tq_protocol_refactor_plan.md b/docs/draft_feature_sample_tq_protocol_refactor_plan.md deleted file mode 100644 index 10398c85..00000000 --- a/docs/draft_feature_sample_tq_protocol_refactor_plan.md +++ /dev/null @@ -1,549 +0,0 @@ -# TQ `DraftFeatureSample` 通用传输协议重构方案 - -> Last updated: 08/27/2026 - -## 1. 目标与结论 - -当前 standalone TQ 流程已经在 Consumer 侧恢复为 `DraftFeatureSample`,然后调用既有的 -`DrafterTrainer.prepare_training_batch_from_samples()`。但是传输层仍然通过一套较重的 -`SampleMetadata` 重新描述 hidden-state 布局、shape 和训练字段,导致: - -- `DraftFeatureSample.metadata` 不能完整往返; -- TQ codec 了解过多 DSpark/hidden-state 布局细节; -- 新算法即使已经能构造 `DraftFeatureSample`,仍可能需要修改 TQ 协议; -- Producer 和 Consumer 分别实现 ready tag 过滤,容易产生统计口径不一致。 - -本次重构采用以下边界: - -1. 保留 `SampleMetadata`,但将它缩减为 **TQ 控制信封**; -2. 完整训练数据只由 `DraftFeatureSample` 表达; -3. TQ codec 对 `DraftFeatureSample` 做通用、无损、算法无关的编码和解码; -4. Consumer 解码后直接把 `DraftFeatureSample` 交给现有训练流程; -5. 算法差异只保留在 Producer 的样本构造和 Trainer backend 中; -6. Producer 和 Consumer 复用同一个 ready-tag 解析函数。 -7. Producer 写入资格与 Consumer 读取资格使用同一个共享判定,禁止两端分别实现近似校验; -8. hidden-state token/position/row 对齐失败的样本在 Producer 侧直接丢弃,不得部分截取后写入 TQ。 - -TQ 不能直接存放 Python dataclass 实例。TQ 0.1.7 的数据面接收 tensor fields,因此仍然 -需要 `encode_sample()` / `decode_sample()`;这里要删除的是自定义的训练数据定义,而不是 -必要的传输编码。 - -## 2. 重构后的职责划分 - -### 2.1 `SampleMetadata`:只负责队列控制 - -建议保留以下字段: - -```python -@dataclass(frozen=True) -class SampleMetadata: - protocol_schema_version: int - run_id: str - sample_id: str - sequence_no: int -``` - -字段含义: - -| 字段 | 用途 | -| --- | --- | -| `protocol_schema_version` | TQ key/tag/fields 编码格式的版本,不是模型算法版本 | -| `run_id` | 隔离不同 standalone 训练任务 | -| `sample_id` | 保留输入样本身份,便于定位错误 | -| `sequence_no` | 为并发完成的样本建立确定顺序,并生成唯一 key | - -从 `SampleMetadata` 删除以下字段: - -```text -algorithm -target_model_id -target_model_revision -tokenizer_fingerprint -target_layer_ids -hidden_states_layout -hidden_dtype -hidden_shape -feature_length -full_sequence_length -feature_start -feature_end -use_logits -``` - -这些字段如果训练需要,应保存在 `DraftFeatureSample.algorithm` 或 -`DraftFeatureSample.metadata` 中。TQ 控制层不再验证其算法语义。 - -### 2.2 TQ tag:控制面索引 - -tag 保留可在不加载 tensor payload 的情况下完成发现、排序和 run 隔离所需的信息: - -```python -tag = { - "record_type": "sample", - "status": "ready", - "protocol_schema_version": 2, - "run_id": "dspark-a1b2c3", - "sequence_no": 6500, - "sample_id": "train-006500", -} -``` - -tag 不再携带 `algorithm`、layer IDs、hidden shape 等训练信息。EOS 仍使用独立 control tag: - -```python -tag = { - "record_type": "control", - "status": "eos", - "protocol_schema_version": 2, - "run_id": "dspark-a1b2c3", - "total_samples": 15000, -} -``` - -### 2.3 TQ fields:完整 `DraftFeatureSample` - -fields 使用原生 tensor 字段加一个 JSON manifest: - -```python -fields = { - "sample__input_ids": Tensor, - "sample__loss_mask": Tensor, - "sample__hidden_states": Tensor, - "sample__position_ids": Tensor, # 可选 - "sample__last_hidden_states": Tensor, # 可选 - "sample__target": Tensor, # 可选 - "sample__target_logprobs": Tensor, # 可选 - "sample__metadata_tensor__000000": Tensor, - "sample__manifest_json": UInt8Tensor, -} -``` - -`sample__manifest_json` 的逻辑内容示例: - -```json -{ - "draft_feature_schema_version": 1, - "algorithm": "DSPARK", - "present_fields": [ - "input_ids", - "loss_mask", - "hidden_states", - "position_ids" - ], - "hidden_states_kind": "tensor", - "metadata": { - "hidden_states_layout": "dflash_aux_plus_last", - "target_layer_ids": [1, 12, 23, 34, 45], - "feature_start": 31, - "feature_end": 543, - "hidden_positions": { - "__tq_tensor_ref__": "sample__metadata_tensor__000000" - } - } -} -``` - -manifest 和 tensor fields 合起来必须能够完整恢复: - -```python -DraftFeatureSample.from_dict(payload, strict=True) -``` - -## 3. 通用 metadata codec - -`DraftFeatureSample.metadata` 不能简单 `json.dumps()`,因为当前代码会在其中保存 -`hidden_positions` 等 tensor。新 codec 采用递归 tree 编码。 - -直接写入 JSON 的类型: - -```text -None、bool、int、float、str -dict[str, value] -list[value] -tuple[value](manifest 记录 tuple 类型,解码后恢复 tuple) -``` - -tensor 的处理方式: - -```text -metadata中的Tensor -→ 转为CPU contiguous Tensor -→ 单独写入fields -→ manifest原位置写tensor field引用 -``` - -例如: - -```python -metadata = { - "feature_start": 31, - "hidden_positions": torch.tensor([31, 32, 33]), -} -``` - -编码为: - -```python -fields["sample__metadata_tensor__000000"] = tensor([31, 32, 33]) - -manifest["metadata"] = { - "feature_start": 31, - "hidden_positions": { - "__tq_tensor_ref__": "sample__metadata_tensor__000000" - }, -} -``` - -不支持的对象不能静默执行 `str(value)`,否则协议不是无损的。第一版应 fail closed,错误中打印 -metadata 路径和实际类型。后续如果确实存在 NumPy scalar/array,可显式增加稳定编码规则。 - -## 4. `hidden_states` 两种表示 - -`DraftFeatureSample.hidden_states` 支持: - -```python -torch.Tensor | list[torch.Tensor] -``` - -单 tensor: - -```python -fields["sample__hidden_states"] = hidden -manifest["hidden_states_kind"] = "tensor" -``` - -tensor list: - -```python -fields["sample__hidden_states__000000"] = hidden_0 -fields["sample__hidden_states__000001"] = hidden_1 -manifest["hidden_states_kind"] = "list" -manifest["hidden_states_fields"] = [ - "sample__hidden_states__000000", - "sample__hidden_states__000001", -] -``` - -这样协议不会再因某个算法使用 tensor list 而报错。 - -## 5. 需要修改的文件和函数 - -### 5.1 `verl_speco/transport/drafter_sample_protocol.py` - -这是主要重构文件。 - -修改内容: - -1. 将 `SampleMetadata` 缩减为控制信封; -2. 将 `PROTOCOL_SCHEMA_VERSION` 从 1 升到 2; -3. 修改 `make_sample_key()`,继续使用 protocol version、run、sequence 和 sample ID; -4. 修改 `make_ready_tag()`,只生成控制面字段; -5. 重写 `encode_sample(sample, meta)`: - - 调用 `DraftFeatureSample.to_dict()`; - - 编码所有 dataclass tensor 字段; - - 编码 hidden-state tensor list; - - 递归编码 metadata; - - 生成 manifest; -6. 重写 `decode_sample(key, tag, fields, expected_config)`: - - 解析并校验控制信封; - - 解析 manifest; - - 恢复所有 tensor 和 metadata; - - 调用 `DraftFeatureSample.from_dict(..., strict=True)`; -7. 将 `_validate_primary_tensors()` 中与具体 hidden layout/shape 的约束删除; -8. 新增并导出统一函数: - -```python -parse_ready_tag(tag) -> SampleMetadata | None -is_ready_sample_tag(tag, *, run_id, protocol_schema_version) -> bool -``` - -Producer backpressure 和 Consumer discovery 必须复用这两个函数,禁止再分别复制过滤条件。 - -此外增加统一的发布资格函数: - -```python -validate_publishable_sample(sample) -> None -``` - -`encode_sample()` 和 Consumer 的 `decode_sample()` 都调用同一组 sample 结构校验。Producer 只有 -通过该校验后才能生成 ready tag;这样不存在“Producer 写入成功,但 Consumer 按另一套规则过滤”的 -中间状态。协议错误必须在 `put_sample()` 前暴露。 - -建议新增内部函数: - -```python -_encode_metadata_tree(value, fields, path) -> JSONValue -_decode_metadata_tree(value, fields, path) -> Any -_encode_hidden_states(value, fields, manifest) -> None -_decode_hidden_states(fields, manifest) -> Tensor | list[Tensor] -_json_to_uint8_tensor(value) -> Tensor -_uint8_tensor_to_json(value) -> Any -``` - -### 5.2 `verl_speco/standalone_tq_producer.py` - -修改内容: - -1. 保留 `PreparedFeature.metadata: SampleMetadata`,但它现在只是 TQ 信封; -2. 简化 `_sample_metadata()`,只读取: - -```text -run_id -request.sample_id -request.sequence_no -protocol_schema_version -``` - -3. 删除 `_sample_metadata()` 对 feature shape、layout、target model 和 logits 的复制; -4. `publish_one()` 仍保持: - -```python -fields = encode_sample(result.sample, result.metadata) -tag = make_ready_tag(result.metadata) -transport.put_sample(key, fields, tag=tag) -``` - -5. `_wait_for_pending_capacity()` 使用协议模块的 - `is_ready_sample_tag()`,与 Consumer 使用完全相同的过滤规则; -6. 保持“只有 TQ put 成功后才删除 vLLM 临时文件”的生命周期不变。 - -Producer 还必须对 hidden-state 对齐失败做样本级丢弃: - -```text -token_ids 与请求 token IDs 不一致 -feature positions 超出 hidden-state rows -hidden-state rows 不能覆盖完整训练窗口 -hidden-state layer 数不足 -→ 记录 sample_id/sequence_no/原因 -→ 删除本次 vLLM 临时文件和 lock -→ dropped_count += 1 -→ 不进入 publish_queue -→ 不写 ready tag/fields -→ request worker 继续处理下一条样本 -``` - -不能沿用当前“只丢弃越界 positions、使用剩余 positions 继续训练”的行为。TQ Producer 应启用严格 -对齐模式:只要一个目标位置无法与 hidden-state row 对应,整条样本就无效。普通网络错误、TQ put -错误和协议编程错误仍然 fail fast,不能被误当作脏样本吞掉。 - -Producer 的算法相关职责仍然保留在: - -```python -feature_from_vllm_payload(raw, request, feature_contract) -``` - -也就是说,Producer 必须先构造正确且完整的 `DraftFeatureSample`,TQ codec 不负责推断算法布局。 - -### 5.3 `verl_speco/trainer/tq_feature_store.py` - -修改内容: - -1. `list_ready()` 使用统一 `parse_ready_tag()`; -2. 删除本文件中重复的 tag 字段校验; -3. `get_many()` 继续调用 `decode_sample()`,返回类型仍为 - `list[DraftFeatureSample]`; -4. `ExpectedFeatureConfig` 只检查: - -```text -run_id -protocol_schema_version -``` - -5. 不再在 TQ store 中检查 algorithm、target model、layer IDs、dtype 和 layout; -6. EOS 解析使用同一个 protocol version 字段命名。 - -### 5.4 `verl_speco/trainer/tq_sample_source.py` - -主流程不需要改变: - -```text -rank 0 list_ready -→ 按sequence_no排序 -→ 为各rank分配key -→ 各rank get_many -→ 得到DraftFeatureSample -``` - -只需确保诊断日志中的 ready 统计也调用统一 tag parser,避免日志口径和正式读取口径不同。 - -### 5.5 `verl_speco/trainer/feature_store.py` - -第一版不修改 `DraftFeatureSample` 公共字段,避免影响现有离线 store、PR #48 和非 TQ 路径。 - -可以新增一个小的公共字段列表,供 store 和 TQ codec 复用,例如: - -```python -DRAFT_FEATURE_OPTIONAL_TENSOR_FIELDS = ( - "last_hidden_states", - "target", - "target_logprobs", - "position_ids", -) -``` - -不要让 TQ codec 再维护一份不同的 optional-field 列表。 - -### 5.6 `verl_speco/trainer/target_feature_replay.py` - -不修改训练转换逻辑。`feature_from_vllm_payload()` 继续负责把不同算法的 vLLM 输出转换为 -`DraftFeatureSample`。 - -当前它支持: - -```text -EAGLE3、DFLASH、DSPARK -``` - -以后增加新算法时,只需要在这里或对应算法 converter 中实现: - -```text -RawVllmFeature + TokenizedRequest → DraftFeatureSample -``` - -如果新的 `DraftFeatureSample` 字段都能被通用 codec 表达,则无需再次修改 TQ 传输层。 - -### 5.7 `verl_speco/trainer/base_trainer.py` - -不需要修改。Consumer 解码结果继续走: - -```python -trainer.prepare_training_batch_from_samples(samples, step=optimizer_step) -``` - -其中每个元素已经是完整 `DraftFeatureSample`,随后调用现有: - -```python -sample.to_training_item() -``` - -算法 backend 仍由启动配置的 `speculative_algorithm` 选择。 - -## 6. 重构后的端到端流程 - -### Producer - -1. 读取 prompt/response; -2. 请求 vLLM generate/prefill; -3. `feature_from_vllm_payload()` 按算法生成完整 `DraftFeatureSample`; -4. 生成简化 `SampleMetadata` 控制信封; -5. `encode_sample()` 无损编码 sample; -6. 将 tensor fields 和 ready tag 写入 TQ; -7. TQ put 成功后删除 vLLM safetensors 临时文件; -8. 全部样本完成后写 EOS。 - -### Consumer - -1. rank 0 使用统一 tag parser 发现当前 run 的 ready keys; -2. 按 `sequence_no` 排序并切出一个 global batch; -3. 广播各 rank 的 key/tag 分配; -4. 各 rank 从 TQ 获取自己的 tensor fields; -5. `decode_sample()` 无损恢复 `DraftFeatureSample`; -6. 调用 `prepare_training_batch_from_samples()`; -7. 由已选择的算法 backend 完成训练; -8. 所有 rank 训练成功后,rank 0 清理这一 global batch 的 keys。 - -## 7. 测试修改 - -### 7.1 `tests/unit/test_drafter_sample_protocol.py` - -替换以 DSpark shape 为中心的测试,增加完整 round-trip: - -1. 单 tensor hidden states; -2. hidden-state tensor list; -3. 所有 optional tensor 字段; -4. metadata 嵌套 dict/list/tuple; -5. metadata 中包含 tensor; -6. EAGLE3/DFLASH/DSPARK/DOMINO 的 algorithm 字符串均能往返; -7. 不支持的 metadata 对象明确报错并包含字段路径; -8. key/tag run、sequence、sample ID 不匹配时 fail closed; -9. protocol version 不匹配时 fail closed。 - -核心断言不是只比较 shape,而是逐字段比较原始 sample 与恢复 sample。 - -### 7.2 `tests/unit/test_tq_producer.py` - -增加: - -1. Producer 写入 fields 后能完整恢复其原始 `DraftFeatureSample.metadata`; -2. pending capacity 使用统一 ready parser; -3. 其他 run、错误 protocol version 和非 sample control tag 不计入 pending; -4. TQ put 失败时仍不删除 vLLM 临时文件。 -5. token IDs 不匹配时删除临时文件、增加 dropped 计数且不调用 TQ put; -6. 任一 feature position 越界时整条样本丢弃,不生成部分长度 sample; -7. 丢弃无效样本后 worker 能继续发布后续有效样本并最终写 EOS; -8. EOS 的 `total_samples` 使用实际成功发布数,不包含 dropped 样本。 - -### 7.3 `tests/unit/test_tq_consumer.py` - -增加: - -1. TQ store 恢复完整 `DraftFeatureSample`; -2. metadata tensor 完整恢复; -3. hidden-state list 完整恢复; -4. ready discovery 与 Producer pending 统计对同一组 tags 给出相同结果; -5. 多算法 sample 都能进入 `prepare_training_batch_from_samples()` 的现有入口。 - -### 7.4 回归测试 - -必须继续运行: - -```bash -pytest -q tests/unit/test_drafter_sample_protocol.py -pytest -q tests/unit/test_tq_producer.py -pytest -q tests/unit/test_tq_consumer.py -pytest -q tests/unit/test_target_feature_replay.py -pytest -q tests/unit/test_draft_feature_store.py -``` - -然后运行一组 Producer/TQ/Consumer smoke test,至少确认: - -```text -put → list → get → decode → train one step → clear → EOS -``` - -## 8. 版本与兼容策略 - -建议将协议版本提升到 2,不实现 v1/v2 混读。 - -理由: - -- standalone TQ 是在线临时队列,不是长期离线数据集; -- Producer 和 Consumer 本来就应作为同一版本部署; -- 同一 run 中混用两种 fields 格式会增加错误恢复复杂度; -- fail closed 比错误地恢复训练数据更安全。 - -启动时 Producer、Owner 和 Consumer 必须使用同一个 protocol version。旧 run 的 TQ 数据不能被新 -Consumer 接续;重启完整 pipeline 时使用新的 `run_id`。 - -## 9. 实现顺序 - -建议按以下顺序实施,每一步都能单独测试: - -1. 简化 `SampleMetadata`,确定 v2 tag/key/EOS 格式; -2. 实现 metadata tree codec; -3. 实现完整 `DraftFeatureSample` encode/decode; -4. 完成 protocol round-trip 单测; -5. 修改 Producer 的控制信封构造和统一 ready 统计; -6. 为 TQ Producer 增加严格 hidden-state 对齐和样本级丢弃; -7. 修改 Consumer store 的统一 tag 解析; -8. 修改 Producer/Consumer 单测; -9. 运行单进程 TQ round-trip smoke; -10. 运行多 rank Consumer 一步训练; -11. 分别用 EAGLE3、DFLASH、DSPARK 的合成 sample 验证通用传输; -12. 最后再为尚未支持的算法增加 vLLM-output-to-sample converter。 - -## 10. 完成标准 - -满足以下条件才算重构完成: - -- TQ codec 中不存在 `if algorithm == "DSPARK"` 一类分支; -- `decode_sample(encode_sample(sample))` 能无损恢复所有公共字段和 metadata; -- Producer 和 Consumer 使用同一个 ready tag parser; -- Consumer 解码后直接得到 `DraftFeatureSample`; -- `base_trainer.py` 的后续训练入口无需为 TQ 添加算法分支; -- 新算法只要能构造现有 `DraftFeatureSample`,传输层无需修改; -- ready backpressure 与 Consumer discovery 对同一组 tags 的计数完全一致; -- TQ put 成功前不删除 vLLM 临时文件,训练成功前不清理 TQ sample。 -- Producer 与 Consumer 对 sample fields 调用同一组结构校验; -- hidden-state 对齐不完整的样本不会以缩短后的 feature window 进入 TQ; -- dropped 样本有明确计数和限频日志,且不会阻止后续有效样本和 EOS。 diff --git a/docs/standalone_tq_consumer_implementation.md b/docs/standalone_tq_consumer_implementation.md deleted file mode 100644 index 75eec9a8..00000000 --- a/docs/standalone_tq_consumer_implementation.md +++ /dev/null @@ -1,713 +0,0 @@ -# 独立 DSpark 训练 TQ Consumer 实现说明 - -Last updated: 08/21/2026 - -## 1. 文档范围和当前结论 - -本文只说明当前仓库中已经实现的独立训练 Consumer。这里的 Consumer 是由 `torchrun` 启动的 DSpark 草稿模型训练任务:它持续从 TransferQueue(下文简称 TQ)发现样本,各训练 rank 分别取得自己负责的 Tensor,复用原有 DSpark 训练逻辑完成一次 optimizer step,然后由 rank 0 删除这一整个 global batch 对应的 TQ 记录。 - -当前已完成的能力是: - -1. `feature_store.type=tq` 可以作为独立训练的数据源,不要求磁盘 `path`。 -2. 每个训练 rank 都连接同一个 Ray 集群、同一个 TQ Controller 和同一个 partition。 -3. 只有 rank 0 调用 `kv_list` 发现 ready key,并把 key/tag 分配给各 rank。 -4. key 和 tag 通过 `torch.distributed.broadcast_object_list` 传输;hidden states 等 Tensor 不经过该广播。 -5. 每个 rank 根据分配到的 key,直接调用 TQ `kv_batch_get` 获取自己的 Tensor。 -6. TQ Tensor 被解码成原训练代码已经认识的 `DraftFeatureSample`,然后复用 `DrafterBaseTrainer` 的 batch 构造、DSpark loss、反向传播和 optimizer step。 -7. 只有当所有 rank 都成功完成该 step 后,rank 0 才调用 `kv_clear` 删除整个 global batch。 -8. Producer 发布 EOS 后,如果剩余样本不足一个 global batch,当前第一版会丢弃并清理这部分尾样本,然后正常结束训练迭代。 - -本文不会把尚未实现的 Producer 写成现有能力。Producer 后续需要复用本文第 6 节所述的公共协议,调用 `encode_sample()` 生成 fields,再使用 bridge 写入相同 TQ。 - -## 2. 本次涉及的文件 - -### 2.1 本次新增的 Consumer 核心文件 - -| 文件 | 实现的组件 | 作用 | -|---|---|---| -| `verl_speco/trainer/tq_feature_store.py` | `TQFeatureStore`、`ReadyEntry`、`EosMetadata` | 将公共 TQ bridge 包装成 Consumer 数据访问层,负责连接、发现、批量读取、解码、删除和读取 EOS | -| `verl_speco/trainer/tq_sample_source.py` | `TQFeatureDataLoader`、`TQLocalBatch`、`build_assignments()` | 实现多 rank 流式取数:rank 0 发现样本并分配 key,各 rank 自己从 TQ 取 Tensor | -| `tests/unit/test_tq_consumer.py` | Consumer 单元测试 | 覆盖 store 构建、ready 过滤排序、协议解码、rank 分配、EOS、清理和非 rank 0 行为 | - -### 2.2 本次修改的既有文件 - -| 文件 | 修改内容 | 为什么要改 | -|---|---|---| -| `verl_speco/trainer/feature_store.py` | factory 新增 `type=tq` 分支 | 让既有独立训练入口能够像选择磁盘 feature store 一样选择流式 TQ 数据源 | -| `verl_speco/trainer/draft_training_loop.py` | 接入 TQ store/loader、跨 rank 连接检查、训练成功后清理 | 将流式取数接入原训练循环,同时保留原 DSpark trainer、loss、optimizer、metric 和 checkpoint 逻辑 | -| `verl_speco/draft_train_launcher.py` | 增加 TQ 启动参数的 fail-fast 检查 | 在启动多个 torchrun 子进程前检查 `enable`、Ray address 和 `run_id`,避免各 rank 启动后才失败 | -| `verl_speco/config/speco_base.yaml` | 标注 `feature_store.type=tq` 为无路径流式数据源 | 保留统一 Hydra 配置入口;TQ 的公共配置仍位于 sibling `training.transfer_queue` | -| `tests/unit/test_draft_train_launcher.py` | 增加 TQ 参数检查测试 | 验证必要配置缺失时 launcher 直接拒绝启动 | -| `tests/unit/test_draft_training_loop.py` | 增加连接和 clear 时序测试 | 验证只由 rank 0 清理、clear 失败会报告、连接失败会传播 | - -### 2.3 直接复用的公共基础 - -以下文件不是这次 Consumer 才创造的概念,但 Consumer 直接使用它们: - -| 文件 | 被复用的能力 | -|---|---| -| `verl_speco/integration/transferqueue_bridge.py` | 屏蔽 TQ 0.1.7 API 细节,提供连接、`kv_list`、`kv_batch_get`、`kv_clear` 和本地关闭接口 | -| `verl_speco/transport/drafter_sample_protocol.py` | 定义 key、tag、fields、metadata 格式,以及 `encode_sample()` / `decode_sample()` | -| `verl_speco/trainer/feature_store.py` | 复用 `DraftFeatureSample`,使 TQ 数据进入训练侧后与磁盘 feature sample 类型一致 | -| `verl_speco/trainer/base_trainer.py` 及既有 backend | 复用 `DrafterBaseTrainer.prepare_training_batch_from_samples()` 和 `training_step_from_batch()` 等训练实现 | - -## 3. 运行时角色 - -### 3.1 TQ Owner - -TQ Owner 是单独的普通 Python 进程。它连接指定 Ray 集群,并以带配置的 `tq.init(config)` 创建任务级 named Controller 和 storage actors。Owner 持有全局 TQ 生命周期;Consumer 结束时不能关闭它。 - -Owner 不是训练 rank,也不执行 DSpark 模型。它的主要作用是让 Producer 和 Consumer 能通过同一个 Ray actor registry 找到同一个 TQ Controller。 - -### 3.2 Producer - -Producer 是后续需要实现的独立推理进程。它应并行调用 vLLM hidden-state 接口,构造一条条 `DraftFeatureSample` 和 `SampleMetadata`,再写入 TQ。 - -Producer 与 Consumer 不通过 Ray RPC 互相调用,也不通过 HTTP 直接传 Tensor。二者只需满足: - -- 连接同一个 Ray address; -- 使用同一个 Ray namespace; -- 使用同一个 TQ partition; -- 使用同一个 `run_id` 和协议版本。 - -### 3.3 Consumer launcher - -`python -m verl_speco.draft_train_launcher` 是父进程。它检查命令行 override,构造 `python -m torch.distributed.run ...` 命令,然后启动训练子进程。 - -launcher 自己不连接 TQ、不取样本、也不持有 GPU 模型。 - -### 3.4 Consumer training rank - -`torchrun --nproc_per_node=N` 会启动 N 个训练 OS 进程。每个进程有独立的: - -- global rank; -- local rank; -- GPU; -- `DrafterBaseTrainer`; -- `TQFeatureStore` 和本地 TQ client; -- DSpark 模型分片及 optimizer 状态。 - -这些 rank 共同执行一个分布式草稿模型训练任务。rank 0 额外负责发现和删除 TQ key;但所有 rank 都会取得各自的训练 Tensor,并参加模型 collective、梯度同步和 optimizer step。 - -### 3.5 Ray 和 torch.distributed 的职责不同 - -本方案仍然使用 Ray,但只因为 TQ 0.1.7 通过 Ray named actor 找 Controller。Consumer 不创建用于训练的 Ray actor,训练本身仍由 `torchrun` 和 `torch.distributed` 执行。 - -两种通信分别是: - -- Ray/TQ:Owner、Producer、每个 Consumer rank 连接共享 TQ;大 Tensor 通过 TQ backend 传输。 -- `torch.distributed`:训练 rank 之间广播小型 key/tag 命令、同步成功状态、训练模型 collective。 - -## 4. 共同配置以及“连接同一个 TQ”的实现 - -关键配置位于: - -```yaml -actor_rollout_ref: - rollout: - drafter: - training: - mode: offline - feature_store: - type: tq - path: null - transfer_queue: - enable: true - ray: - address: 127.0.0.1:6379 - namespace: speco-drafter - partition_id: speco_drafter_features - run_id: dspark-standalone-run - schema_version: 1 - poll_interval_seconds: 0.5 - drop_last: true -``` - -这些字段的含义如下: - -| 字段 | 使用者 | 含义 | -|---|---|---| -| `feature_store.type=tq` | Consumer | 选择流式 TQ source,而不是磁盘 shard/replay source | -| `feature_store.path=null` | Consumer | TQ 不从本地路径读文件,因此无需 path | -| `transfer_queue.enable` | Owner、Producer、Consumer | 开启 bridge 的 TQ 路径 | -| `ray.address` | 三端 | 连接同一个 Ray 集群 | -| `ray.namespace` | 三端 | 在同一 actor namespace 查找 named Controller | -| `partition_id` | 三端 | 对同一个 TQ KV 分区执行 put/list/get/clear | -| `run_id` | Producer、Consumer | 在共享 partition 中区分本次训练数据;Consumer 只接收匹配的样本 | -| `schema_version` | Producer、Consumer | 共同使用的数据协议版本 | -| `poll_interval_seconds` | Consumer rank 0 | ready 数量不足时的轮询间隔 | -| `drop_last` | Consumer | 第一版必须为 true;EOS 后不足 global batch 的尾样本被清理 | - -三端并不是通过共享 Python 对象得到这些配置。每个进程都各自读取相同取值,然后执行: - -```python -configure_transfer_queue(config) -connect_ray_cluster(ray_address, ray_namespace) -connect_transfer_queue_client() -``` - -`connect_ray_cluster()` 内部调用 `ray.init(address=..., namespace=...)`。`connect_transfer_queue_client()` 再使用与 Owner 相同的 native 配置调用 `tq.init(config)`;TQ 会优先在当前 Ray namespace 查找 Owner 创建的 named Controller,找到时忽略本次配置并只创建本地 Client。即使 Client 意外先于 Owner 初始化,也会使用同一份 backend/controller 配置,而不会按默认配置创建服务。之后所有 KV 操作都显式携带相同的 `partition_id`。 - -因此,“连接同一个 TQ”实际由三层身份共同决定:同一 Ray 集群、同一 namespace 下的同一 named Controller、同一 `partition_id`。 - -## 5. 一条样本在 TQ 中的实际格式 - -### 5.1 一条 key 对应一个 sample - -本协议没有把一个训练 batch 存成一个 TQ key。一条 key 对应一条独立训练样本。假设: - -```text -run_id = dspark-run-001 -sequence_no = 17 -sample_id = prompt-000017 -``` - -则 key 为: - -```text -drafter:v1:dspark-run-001:000000000017:prompt-000017 -``` - -`sequence_no` 是本次 run 内的样本顺序号,不是 batch 编号,也不是 optimizer step。Consumer 用它稳定排序,之后每次从有序 ready 列表前部取一个 global batch。 - -### 5.2 tag:用于轻量发现和过滤 - -该 key 的 tag 是普通小字典: - -```python -{ - "record_type": "sample", - "status": "ready", - "schema_version": 1, - "run_id": "dspark-run-001", - "sequence_no": 17, - "sample_id": "prompt-000017", - "algorithm": "DSPARK", -} -``` - -tag 存在 TQ 的 KV 元信息中。`kv_list(partition_id=...)` 返回 `key -> tag`,不需要先加载 hidden states。rank 0 正是依靠 tag 筛选当前 run、当前 schema、DSPARK 且状态为 ready 的记录。 - -### 5.3 fields:真正的 Tensor payload - -同一 key 的 fields 是一个 Tensor 字典: - -```python -{ - "input_ids": Tensor[int64, shape=[L]], - "loss_mask": Tensor[float32, shape=[L]], - "position_ids": Tensor[int64, shape=[L]], - "hidden_states": Tensor[dtype, shape=[L, D]], - "metadata_json": Tensor[uint8, shape=[M]], - # 以下是可选字段: - "last_hidden_states": Tensor[..., ...], - "target": Tensor[..., ...], - "target_logprobs": Tensor[..., ...], -} -``` - -这里 `L` 是 feature window 的 token 数,`D` 是目标模型 hidden size,`M` 是 metadata JSON 序列化后的 UTF-8 字节数。 - -`hidden_states` 等大 Tensor 只存于 fields,通过 TQ `kv_batch_get` 传输;不会放进 tag,也不会通过训练 rank 的 object broadcast。 - -### 5.4 metadata_json:内容丰富但仍随 fields 读取 - -TQ fields 只能承载 Tensor,因此结构化 metadata 被编码为 `uint8` Tensor。解码后的字典格式是: - -```python -{ - "schema_version": 1, - "run_id": "dspark-run-001", - "sample_id": "prompt-000017", - "sequence_no": 17, - "algorithm": "DSPARK", - "target_model_id": "/models/Qwen3-8B", - "target_model_revision": "main", - "tokenizer_fingerprint": "...", - "target_layer_ids": [35], - "hidden_states_layout": "token_major", - "hidden_dtype": "bfloat16", - "hidden_shape": [L, D], - "feature_length": L, - "full_sequence_length": 256, - "feature_start": 64, - "feature_end": 64 + L, - "use_logits": False, -} -``` - -字段分工是: - -- tag:只放发现、过滤、排序所需的小字段;`kv_list` 可直接得到。 -- fields:放训练 Tensor 和完整 metadata;只有被某个 rank 选中后才 `kv_batch_get`。 -- key:把 tag 和 fields 重新关联起来,也是清理记录时传给 `kv_clear` 的标识。 - -### 5.5 控制记录 - -控制记录与 sample 放在同一 partition,但通过 tag 的 `record_type=control` 区分。 - -Owner readiness key: - -```text -control:v1::owner-ready -``` - -EOS key: - -```text -control:v1::eos -``` - -EOS tag 包含 `status=eos` 和 `total_samples`。EOS 表示 Producer 不会再为本次 run 增加新样本;它不是一条训练样本。 - -## 6. 公共协议如何把 Producer 输出还原为训练对象 - -Producer 应调用: - -```python -fields = encode_sample(sample, metadata) -key = make_sample_key(metadata) -tag = make_ready_tag(metadata) -put_sample(key, fields, tag=tag) -``` - -`encode_sample()` 会将所有 Tensor detach、转到 CPU、整理为 contiguous,并统一 `input_ids/position_ids` 为 int64、`loss_mask` 为 float32。随后校验 token 长度、hidden shape 和 metadata 一致,再把 metadata JSON 编成 uint8 Tensor。 - -Consumer 的逆过程位于 `TQFeatureStore.get_many()`: - -```python -records = get_samples([entry.key for entry in entries]) -sample = decode_sample( - key=key, - tag=entry.tag, - fields=fields, - expected_config=self.expected_config, -) -``` - -`get_samples()` 最终调用一次 TQ `kv_batch_get(keys=[...], partition_id=...)`。bridge 将 TQ 返回的 batched TensorDict 或 mapping 拆成与请求 key 顺序一致的普通 fields 字典。 - -`decode_sample()` 随后: - -1. 检查必需 fields 是否存在。 -2. 将 `metadata_json` 从 uint8 Tensor 还原为字典和 `SampleMetadata`。 -3. 根据 metadata 重新计算 key,并与实际 key 比较。 -4. 比较 tag 与 metadata 的公共身份字段。 -5. 检查 Consumer 的 expected config。 -6. 将 Tensor detach 到 CPU,统一基础 dtype/shape。 -7. 检查所有主 Tensor 第一维等于 `feature_length`,hidden shape/dtype 与 metadata 一致。 -8. 构造 `DraftFeatureSample`。 - -输出不再是 TQ 专用对象,而是既有训练代码使用的: - -```python -DraftFeatureSample( - input_ids=..., - loss_mask=..., - position_ids=..., - hidden_states=..., - metadata=..., - ..., -) -``` - -这是能够复用原训练逻辑的关键边界:TQ 只负责上游存储和传输,`decode_sample()` 后的数据类型与磁盘 feature store 读取结果一致。 - -## 7. Consumer 从启动到结束的完整执行流程 - -### 阶段 1:launcher 检查配置并启动 torchrun - -执行者是 launcher 父进程。入口是 `verl_speco.draft_train_launcher.main()`。 - -当 override 中出现 `feature_store.type=tq`,`validate_tq_launch_config()` 会要求: - -- `training.transfer_queue.enable=true`; -- `training.transfer_queue.ray.address` 非空; -- `training.transfer_queue.run_id` 非空。 - -检查成功后构造: - -```text -python -m torch.distributed.run - --nnodes=... - --nproc_per_node=... - -m verl_speco.draft_train - <全部 Hydra overrides> -``` - -配置参数是普通子进程命令行参数。此阶段没有 TQ Tensor 传输。 - -### 阶段 2:每个 rank 初始化训练运行时 - -每个 torchrun 子进程进入 `run_standalone_draft_training()`,调用 `_init_distributed()` 得到 `rank/local_rank/world_size`,绑定本 rank GPU,然后构造原有 `DrafterBaseTrainer` 和 DSpark backend。 - -`speculative_algorithm=DSPARK` 决定 backend 和 DSpark 模型训练实现;`feature_store.type=tq` 只改变数据来源,不替换 trainer。 - -### 阶段 3:factory 创建 TQFeatureStore - -训练循环调用: - -```python -store = build_feature_store_from_config( - feature_store_cfg, - read_only=True, - transfer_queue_cfg=training_cfg.get("transfer_queue"), -) -``` - -factory 在 `type=tq` 时不读取 `feature_store.path`,而是把 sibling `training.transfer_queue` 交给 `TQFeatureStore.from_config()`。 - -TQ store 被限定为 `read_only=True`,意思是它是训练 Consumer source。这里的“read only”不表示永不修改 TQ;成功消费后仍可通过明确的 `clear_many()` 删除记录,但不会把它当作通用 feature writer。 - -### 阶段 4:所有 rank 分别连接同一个 TQ - -训练循环调用 `_connect_tq_store_across_ranks()`。每个 rank 都独立执行 `store.connect()`: - -```text -configure_transfer_queue -→ ray.init(address, namespace) -→ tq.init(same native config) 连接 named Controller -→ 本 rank 设置 _connected=True -``` - -之后 `_all_ranks_true()` 使用 `dist.all_reduce(MIN)` 汇总连接结果。只要一个 rank 连接失败,所有 rank 都停止,不允许部分 rank 进入后续 broadcast 或 FSDP collective。 - -这里没有“rank 0 建一个 client 给其他 rank 共用”。TQ client 是进程本地对象,N 个 rank 有 N 个 client,但它们指向同一 Controller/partition。 - -### 阶段 5:创建 TQFeatureDataLoader - -每个 rank 构造自己的 loader,参数包括相同的 `batch_size_per_gpu`、`world_size`、轮询间隔和 drop-last,以及不同的 `rank`。 - -假设: - -```text -world_size = 2 -batch_size_per_gpu = 2 -global_batch_size = 4 -``` - -那么只有 ready 数量至少为 4,rank 0 才发布一个 batch 命令。 - -### 阶段 6:rank 0 发现 ready key - -rank 0 首先检查 `owner_ready()`。Owner 尚未发布 readiness marker 时,rank 0 sleep 后继续轮询,不会让其他 rank 开始取数。 - -Owner ready 后,rank 0 调用 `list_ready()`,其底层是: - -```text -tq.kv_list(partition_id) -→ key -> tag -→ 按 record_type/status/run/schema/algorithm 过滤 -→ 按 (sequence_no, key) 排序 -``` - -此阶段没有读取 fields,因此 hidden states 尚未传到训练进程。 - -### 阶段 7:rank 0 切分 global batch - -若排序后的前四条是 `k0、k1、k2、k3`,`build_assignments()` 产生: - -```python -assignments = [ - [ReadyEntry(k0, tag0), ReadyEntry(k1, tag1)], # rank 0 - [ReadyEntry(k2, tag2), ReadyEntry(k3, tag3)], # rank 1 -] -``` - -每条样本只出现在一个 rank 的 assignment 中,因此各 rank 不会取得同一训练样本。这里采用连续、不重叠的切片。 - -rank 0 随后构造普通 Python 命令字典: - -```python -{ - "kind": "batch", - "global_keys": [k0, k1, k2, k3], - "assignments": [ - [{"key": k0, "tag": tag0}, {"key": k1, "tag": tag1}], - [{"key": k2, "tag": tag2}, {"key": k3, "tag": tag3}], - ], -} -``` - -`global_keys` 只用于 rank 0 在训练完成后一次清理整个 batch;`assignments` 用于每个 rank 知道自己应该 get 哪些 key。 - -### 阶段 8:小型命令通过 torch.distributed 广播 - -各 rank 同时进入: - -```python -dist.broadcast_object_list(payload, src=0) -``` - -rank 0 的 payload 中是上述字典,其他 rank 的初始值是 `None`。PyTorch 会序列化这个普通 Python 对象并广播给所有 rank。 - -这条边界只传输字符串、整数和小字典 tag。`hidden_states`、`input_ids` 等 fields 不在命令中,所以不会经 rank 0 中转,也不会随 broadcast 复制完整 global batch Tensor。 - -### 阶段 9:每个 rank 直接从 TQ 取本地 payload - -每个 rank 从 `assignments[self.rank]` 还原自己的 `ReadyEntry`: - -```python -local_entries = [_entry_from_wire(item) for item in assignments[self.rank]] -samples = self.store.get_many(local_entries) -``` - -在上述例子中: - -- rank 0 调用 `kv_batch_get(keys=[k0, k1], partition_id=...)`; -- rank 1 调用 `kv_batch_get(keys=[k2, k3], partition_id=...)`。 - -大 Tensor 的数据面因此是 TQ storage 到目标训练 rank,不经过训练 rank 0 的 Python 内存。每个 rank 得到两个 CPU `DraftFeatureSample`。 - -loader yield: - -```python -TQLocalBatch( - local_keys=[本 rank 的 key], - local_samples=[本 rank 的 DraftFeatureSample], - global_keys=[完整 global batch key] if rank == 0 else None, -) -``` - -非 rank 0 不保存 `global_keys`,避免多个 rank 都尝试 clear。 - -### 阶段 10:复用已有训练 batch 构造 - -训练循环识别 `TQLocalBatch` 后,只取: - -```python -samples = tq_local_batch.local_samples -``` - -然后调用原有接口: - -```python -batch = trainer.prepare_training_batch_from_samples( - materialized_samples, - step=optimizer_step, -) -``` - -TQ 路径禁止同时开启 `target_feature_pipeline`,因为样本已经包含目标模型 hidden states,不需要训练侧再访问 vLLM materialize 一次。 - -此时 TQ 专用的 key/tag 已不参与 DSpark 数学计算;训练接口看到的是普通 `DraftFeatureSample`,并按原逻辑整理 input ids、hidden states、mask、position ids 和 DSpark 训练所需输入。 - -### 阶段 11:所有 rank 同步 batch 是否可训练 - -每个 rank 判断 `batch is not None`,再通过 `_all_ranks_true()` 做 `all_reduce(MIN)`。 - -只有所有 rank 都成功构造 batch,才能进入训练。如果任一 rank 解码或 batch 构造失败,TQ 路径直接报错,而且这些 key 不会被删除。 - -### 阶段 12:执行原有 DSpark training step - -每个 rank 调用: - -```python -ok = await trainer.training_step_from_batch(batch, optimizer_step) -``` - -该调用复用既有模型 forward、DSpark loss(包括配置开启时的 L1 loss)、backward、梯度同步和 optimizer step。TQ 新代码没有重新实现 loss 或 optimizer。 - -之后再次以 `_all_ranks_true(ok)` 同步。只有所有 rank 都返回成功,才认为这一个 global batch 已经安全消费。 - -### 阶段 13:训练成功后由 rank 0 删除 global batch - -训练循环调用 `_clear_tq_batch_across_ranks()`: - -1. rank 0 使用 `tq_local_batch.global_keys` 调用 `loader.clear_completed_batch()`。 -2. loader 调用 `store.clear_many(global_keys)`。 -3. bridge 最终调用 `tq.kv_clear(keys=[k0,k1,k2,k3], partition_id=...)`。 -4. 所有 rank 通过 `all_reduce(MAX)` 同步 clear 是否失败。 - -删除发生在 optimizer step 全 rank 成功之后。不是“某个 rank get 完就删除”,因为 get 完只代表 Tensor 已读取,不能代表训练 step 已成功。 - -clear 成功后才增加 `successful_steps`,然后复用原有 metrics 和 checkpoint 调度。 - -### 阶段 14:EOS 和尾 batch - -当 ready 样本少于一个 global batch时,rank 0 查询 EOS: - -- 没有 EOS:说明 Producer 以后仍可能写入更多样本,sleep 后继续轮询。 -- 已有 EOS 且 ready 为空:广播 `{"kind": "stop"}`,所有 rank 结束迭代。 -- 已有 EOS 且存在不足一个 global batch 的尾样本:rank 0 先 clear 这些尾 key,再广播 stop。 - -第一版强制 `drop_last=true`,因此不会构造各 rank batch size 不一致的最后一步。 - -### 阶段 15:checkpoint 和退出清理 - -正常 step 完成后仍按原 `save_interval_steps` 保存 checkpoint;循环结束后按 `save_final_checkpoint` 决定是否保存最终 checkpoint。 - -`finally` 中每个 rank 调用 `store.close()`。对 `TQFeatureStore` 而言,这只是: - -```text -关闭本进程 TQ client -→ 如果本进程自行 ray.init,则 ray.shutdown() -``` - -它不会调用全局 `tq.close()`,不会杀死 Owner 创建的 Controller,也不会影响仍在运行的 Producer 或其他 rank。 - -## 8. 控制面和数据面的完整边界 - -| 数据 | 从哪里到哪里 | 传输机制 | 是否经过 rank 0 | -|---|---|---|---| -| 启动配置 | launcher 到 torchrun 子进程 | 命令行 Hydra overrides | 每个 rank 都收到 | -| ready key/tag | TQ Controller 到 rank 0 | `tq.kv_list` | 是,只有 rank 0 list | -| batch assignment | rank 0 到全部 rank | `dist.broadcast_object_list` | 由 rank 0 发出 | -| hidden states 等 fields | TQ storage 到被分配的 rank | `tq.kv_batch_get` | rank 1 的 Tensor 不经过 rank 0 | -| batch 准备/训练成功状态 | 全部 rank 之间 | Tensor `all_reduce` | collective,无单点 payload relay | -| clear 请求 | rank 0 到 TQ | `tq.kv_clear(global_keys)` | 只有 rank 0 发起 | -| 梯度和模型 collective | 训练 rank 之间 | 既有 PyTorch distributed/FSDP 路径 | 与 TQ 无关 | - -## 9. 当前“最简单校验”具体简单在哪里 - -`TQFeatureStore` 构造的 expected config 只固定: - -```python -ExpectedFeatureConfig( - run_id=<当前训练 run_id>, - schema_version=<当前 schema>, -) -``` - -TQ 会保留 Producer 写入的 `SampleMetadata.algorithm`,但不使用它选择训练 backend, -也不额外与启动配置比较。与原有离线 feature-store 训练一致,实际 trainer/backend 只由 -`rollout.drafter.speculative_algorithm` 和既有 backend factory 决定。 - -因此当前不会拿 Consumer 配置额外比较: - -- target model ID/revision; -- tokenizer fingerprint; -- target layer IDs; -- hidden layout; -- hidden dtype 的外部预期值。 - -但这不等于完全不校验。`decode_sample()` 仍然强制检查: - -- 必需 fields 存在; -- key、tag、metadata 三者身份一致; -- schema/run/algorithm 符合 Consumer; -- Tensor 类型正确; -- input/mask/position/hidden 长度一致; -- hidden 实际 shape/dtype 与该样本 metadata 一致; -- feature window 合法。 - -这满足“第一版少做外部模型身份检查”,同时避免把结构损坏或错 run 的数据送入训练。 - -## 10. 失败、删除和重复消费语义 - -当前实现遵循以下规则: - -1. 连接失败:所有 rank 同步停止。 -2. rank 0 list/EOS 失败:rank 0 广播 error 命令,其他 rank 不会永久等待 batch broadcast。 -3. 某 rank get/decode 失败:`_next_batch_across_ranks()` 将失败同步给全部 rank,不进入模型训练 collective。 -4. 某 rank 无法构造 batch:报错,global keys 保留在 TQ。 -5. 某 rank training step 失败:报错,global keys 保留在 TQ。 -6. 全 rank training step 成功:rank 0 clear 整个 global batch。 -7. clear 失败:错误传播到全部 rank,训练停止;不会把该 step 继续当成已正常完成。 -8. 达到 `max_steps`:循环停止;尚未选择的 ready 样本保留在 TQ。 - -第一版尚未实现完整的崩溃恢复协议。尤其是“optimizer step 已成功,但进程在 clear 前崩溃”时,key 仍存在;重新启动 Consumer 可能再次读取它。要实现严格 exactly-once,需要把 checkpoint step、已消费 sequence 或事务状态纳入协议。该能力应作为后续增强,而不是当前已实现能力。 - -## 11. 如何启动和检查 - -正式运行统一使用端到端 launcher;它负责 Ray、TQ Owner、Producer 和 Consumer 的启动与清理: - -```bash -bash examples/run_qwen3-8b_drafter_separate_training.sh -``` - -## 12. 已完成的测试 - -### 12.1 Consumer/factory/协议/launcher 单元测试 - -已执行: - -```text -python -m pytest \ - tests/unit/test_tq_consumer.py \ - tests/unit/test_draft_train_launcher.py \ - tests/unit/test_transferqueue_bridge.py \ - tests/unit/test_drafter_sample_protocol.py \ - tests/unit/test_draft_feature_store.py \ - -q -``` - -结果:`44 passed`。 - -### 12.2 真实 TQ 0.1.7 跨进程 smoke - -早期临时跨进程工具验证过以下行为,当前回归由协议、bridge、Consumer和launcher单元测试承担: - -- Owner 发布 owner-ready; -- 两条 sample 写入 TQ; -- Consumer 经 `TQFeatureStore` 和 `TQFeatureDataLoader` 读到两条样本; -- hidden shape 正确; -- Consumer clear 已完成 batch; -- EOS 后迭代停止; -- Consumer 只关闭本地 client,Owner 仍能继续观察完成标记并正常关闭。 - -实际 smoke 输出包含: - -```text -CLIENT_OK samples=2 shape=(3,4) -CLIENT_CLOSED_LOCAL_ONLY -OWNER_OBSERVED_SAMPLES_CLEARED -OWNER_CLOSED -``` - -### 12.3 当前环境未覆盖的部分 - -完整 `tests/unit/test_draft_training_loop.py` 在当前 Windows 环境无法完整收集,因为上游 `verl/ray` 依赖不齐;新增训练循环测试代码已通过 Python 编译检查,连接/clear helper 也通过针对性单元逻辑验证。真实多 GPU DSpark 训练仍需要在目标 Linux GPU 环境执行集成测试。 - -## 13. 当前限制和后续建议 - -当前第一版有意不实现以下复杂能力: - -1. Producer 本身尚未在本次 Consumer 改动中实现。 -2. TQ 公共协议和 Consumer 已不再写死 `DSPARK`;当前测试 Producer、启动脚本和已验证的 - feature 语义仍是 DSPARK。其他算法若能复用当前公共 dense fields,只需由对应 Producer - 生成正确的 `DraftFeatureSample`;若字段结构不同,则在协议模块增加对应 codec,不需要改 - TQ 的 key/tag 发现、rank 分配和 clear 流程。 -3. 只支持 `drop_last=true`。 -4. 不支持 TQ 与 `target_feature_pipeline.enabled=true` 同时开启。 -5. 不提供严格的 crash exactly-once 或 checkpoint/queue 联合恢复。 -6. rank 0 仍通过 `kv_list` 轮询整个 partition;数据量很大时可考虑 cursor/ready queue 优化。 -7. 当前外部 expected config 校验较简化,后续可把 model revision、tokenizer fingerprint、layer/layout/dtype 预期接入 Hydra 配置。 -8. 尚需在真实多机、多 GPU、Mooncake backend 环境验证吞吐、背压、Owner 生命周期和网络故障行为。 - -建议下一阶段优先完成 Producer,并严格复用 `drafter_sample_protocol.py`,不要在 Producer 另造一套 key/tag/fields 格式。完成 Producer 后,首先跑 world size 1 的端到端训练,再跑多 rank 验证每条 key 只分配给一个 rank、训练成功后只由 rank 0 clear。 - -## 14. 最终路径摘要 - -```text -Producer(待实现) - vLLM 并行 prefill - → DraftFeatureSample + SampleMetadata - → encode_sample 得到 Tensor fields - → TQ kv_put(key, fields, tag) - -Consumer rank 0 - kv_list 只取 key/tag - → 过滤并按 sequence_no 排序 - → 切出 global batch - → broadcast 每个 rank 的 key/tag assignment - -每个 Consumer rank - 取 assignments[rank] - → kv_batch_get 本 rank keys - → decode_sample 得到 DraftFeatureSample - → 原 prepare_training_batch_from_samples - → 原 DSpark training_step_from_batch - -全部 rank - 同步确认 optimizer step 成功 - → rank 0 kv_clear(global_keys) - → 原 metrics/checkpoint - → 下一批 - -Producer 发布 EOS - → rank 0 确认没有完整 global batch - → 清理不足一批的尾样本 - → broadcast stop - → 各 rank 关闭本地 TQ client 并退出 -``` diff --git a/docs/standalone_tq_drafter_resume_plan.md b/docs/standalone_tq_drafter_resume_plan.md deleted file mode 100644 index 17ad0586..00000000 --- a/docs/standalone_tq_drafter_resume_plan.md +++ /dev/null @@ -1,448 +0,0 @@ -# Standalone TQ 独立训练断点续训方案 - -## 1. 第一版目标 - -当前 DSpark 独立训练已经能够从 checkpoint 恢复模型权重、optimizer、LR scheduler、`optimizer_steps_total` 和 `training_steps`,但没有恢复数据进度。重启后 producer 会从文件开头重新生产,导致 checkpoint 之前已经训练的数据再次进入训练。 - -第一版采用最直接的方案: - -```text -consumer 记录已经成功训练的 sequence_no -→ checkpoint 保存这些 sequence_no -→ 重启时 producer 加载这个集合 -→ 读文件时跳过已消费编号 -→ 其他样本仍按原逻辑并行请求 vLLM -``` - -不修改 producer 的并发请求和完成顺序,不恢复旧 TQ,也不引入顺序发布。 - -## 2. 当前流程已经具备的条件 - -### 2.1 输入已经有编号 - -`standalone_tq_producer.py::read_inputs()` 当前执行: - -```python -record = replace(source_record, sequence_no=stats.input_count) -``` - -编号进入现有 TQ tag: - -```python -tag = { - "record_type": "sample", - "status": "ready", - "schema_version": 2, - "run_id": "...", - "sample_id": "...", - "sequence_no": 1234, -} -``` - -所以不需要新增样本身份协议,直接使用现有 `sequence_no`。前提是续训使用相同输入文件且行顺序不变。 - -### 2.2 consumer rank 0 已经知道 batch 编号 - -`TQFeatureDataLoader.__iter__()` 中 rank 0 执行: - -```python -ready = self.store.list_ready() -selected = ready[:global_batch_size] -``` - -`selected` 中每个 `ReadyEntry` 都有 `entry.tag["sequence_no"]`。因此 rank 0 已经知道本 global batch 实际使用了哪些输入编号。 - -### 2.3 消费成功边界已经明确 - -训练循环当前顺序为: - -```text -training_step_from_batch() -→ 所有 rank 确认成功 -→ rank 0 clear_completed_batch(global_keys) -→ successful_steps += 1 -→ 按间隔保存 checkpoint -``` - -规定: - -```text -训练成功 + TQ clear 成功 = 本 batch 的 sequence_no 已消费 -``` - -训练或 clear 失败时不能更新已消费集合。 - -## 3. checkpoint 增加什么 - -当前 checkpoint: - -```text -draft_step_60000/ -├── config.json -├── model.safetensors 或模型 shards -├── metadata.json -└── optimizer/ -``` - -新增: - -```text -draft_step_60000/consumed_sequence_nos.pt -``` - -内容为排序、去重的 CPU int64 tensor: - -```python -tensor([0, 1, 2, 4, 5, 8, ...], dtype=torch.int64) -``` - -允许存在间隔:alignment 失败的数据、尚未训练的 TQ 数据和仍在推理的数据都不在集合中。 - -空间开销约为: - -```text -100 万个编号:7.6 MiB -1000 万个编号:76 MiB -``` - -相比模型和 optimizer checkpoint 很小。 - -`metadata.json` 增加: - -```python -"standalone_data_progress": { - "version": 1, - "consumed_sequence_file": "consumed_sequence_nos.pt", - "consumed_sequence_count": int, - "input_fingerprint": {...}, -} -``` - -## 4. 为什么乱序推理不影响这个方案 - -假设 vLLM 完成顺序为: - -```text -5, 1, 8, 2, 4, 0, 3 -``` - -consumer 实际训练: - -```text -batch 1: [1, 5] -batch 2: [0, 2] -``` - -checkpoint 保存: - -```python -consumed_sequence_nos = tensor([0, 1, 2, 5]) -``` - -重启后 producer 只跳过 0、1、2、5,其余编号重新请求。因此不要求已消费数据连续,也不需要 producer 按顺序写 TQ。 - -## 5. consumer 修改 - -### 5.1 `TQLocalBatch` 携带 global 编号 - -文件:`verl_speco/trainer/tq_sample_source.py` - -改为: - -```python -@dataclass(frozen=True) -class TQLocalBatch: - local_keys: list[str] - local_samples: list[DraftFeatureSample] - global_keys: list[str] | None - global_sequence_nos: list[int] | None -``` - -rank 0 构造 command 时增加: - -```python -"global_sequence_nos": [ - int(entry.tag["sequence_no"]) - for entry in selected -] -``` - -只有 rank 0 需要保存 `global_sequence_nos`,其他 rank 保持 `None`。 - -### 5.2 clear 成功后更新集合 - -文件:`verl_speco/trainer/draft_training_loop.py` - -启动时: - -```python -consumed_sequence_nos = load_consumed_sequence_nos(drafter_cfg.model_path) -``` - -在 `_clear_tq_batch_across_ranks()` 成功返回以后: - -```python -if rank == 0: - consumed_sequence_nos.update( - tq_local_batch.global_sequence_nos or [] - ) -``` - -运行时使用 `set[int]` 便于去重;保存前转换为不可变 snapshot: - -```python -consumed_snapshot = torch.tensor( - sorted(consumed_sequence_nos), - dtype=torch.int64, -) -``` - -## 6. checkpoint 修改 - -不修改 `verl_speco/trainer/base_trainer.py`。该类同时被 co-train 使用,把独立训练的数据进度塞进它的公共 checkpoint 接口,会扩大影响范围。 - -独立训练仍先调用现有的: - -```python -trainer.save_checkpoint(step=step, wait=wait) -``` - -然后只在 `draft_training_loop.py` 的 `_save_standalone_checkpoint()` 中追加独立训练 sidecar: - -```text -draft_step_N/ -├── 原有模型、optimizer、scheduler 和 metadata -├── consumed_sequence_nos.pt -└── standalone_resume.json -``` - -其中 `standalone_resume.json` 记录 consumed 数量、输入 fingerprint 和 sidecar 版本。只有原 checkpoint 已成功保存后,才发布这两个文件。 - -使用临时文件原子写入: - -```python -temporary = checkpoint_path / "consumed_sequence_nos.pt.incomplete" -final = checkpoint_path / "consumed_sequence_nos.pt" -torch.save(consumed_snapshot, temporary) -os.replace(temporary, final) -``` - -异步保存时,`_save_standalone_checkpoint()` 先复制不可变 snapshot,再给现有 checkpoint future 注册 callback。callback 只在模型 checkpoint future 成功后原子写 sidecar,不能让后台 callback 直接读取仍在变化的 Python set。 - -如果进程恰好在模型 checkpoint 完成、sidecar 尚未完成时崩溃,这个目录不能用于“精确数据续训”,应回退到上一个同时具备完整模型 checkpoint 和完整 sidecar 的目录。 - -加载时校验: - -- 文件存在; -- 一维 `torch.int64`; -- 所有值非负; -- 已排序、无重复; -- tensor数量与 metadata一致。 - -旧 checkpoint没有该文件时,默认提示它只能恢复训练状态、不能精确恢复数据;严格模式下拒绝续训。 - -## 7. producer 修改 - -### 7.1 新配置 - -```yaml -speco: - standalone_tq_producer: - consumed_sequence_path: null -``` - -launcher 在续训时传入: - -```text -/consumed_sequence_nos.pt -``` - -### 7.2 启动时加载 - -文件:`verl_speco/standalone_tq_producer.py`。 - -```python -consumed_sequence_nos = load_consumed_sequence_nos( - producer_cfg.get("consumed_sequence_path") -) -``` - -### 7.3 扫描时跳过 - -需要拆开两个计数: - -```python -source_sequence_no # 所有扫描过的输入,跳过也递增 -queued_count # 本次真正送入 input_queue 的数量 -``` - -第一版保持 `iter_input_records()` 接口不变,在它产出记录后、tokenizer 和 vLLM 之前,根据 epoch 内扫描顺序算出全局 `sequence_no`: - -```python -sequence_no = source_sequence_no -source_sequence_no += 1 - -if sequence_no in consumed_sequence_nos: - continue - -record = replace(source_record, sequence_no=sequence_no) -await input_queue.put(request) -queued_count += 1 -``` - -因此恢复时仍需从输入文件开头顺序读取并解析一次,以重建稳定编号,但已消费行不会进入 tokenizer、input queue 或 vLLM。集合查询是平均 O(1),后续仍使用原来的多个 `request_worker()`,不会降低 producer 并发。 - -这里不是每个训练 step 都重新读取 checkpoint 文件。`consumed_sequence_nos.pt` 只在 producer 启动时加载一次;之后每扫描到一个源记录,只做一次内存集合查询。对百万级样本,主要额外成本是一次顺序读文件和 JSON/Parquet 行解析,通常远小于 tokenizer 和 vLLM 推理。只有实际测量发现启动扫描成为瓶颈后,才考虑给 input reader 增加解析前跳过、文件 offset 索引或 bitmap,第一版不做。 - -## 8. 多 epoch 编号 - -编号必须跨 epoch 连续: - -```text -数据集 10000 行 -epoch 0:0~9999 -epoch 1:10000~19999 -epoch 2:20000~29999 -``` - -加入 skip 后不能继续用一个 `stats.input_count` 同时表示扫描位置和排队数量,否则跳过数据后编号会改变。必须使用独立的 `source_sequence_no`。 - -## 9. 本次需要生产多少数据 - -独立训练的 `MAX_STEPS` 改为目标总 optimizer step,与首次训练配置保持一致: - -```text -MAX_STEPS = 训练完成时的总步数 -``` - -例如 checkpoint为 60000,目标总步数 930000: - -```bash -DRAFTER_PATH=/path/to/draft_step_60000 -MAX_STEPS=930000 -LR_WARMUP_STEPS=<仍使用首次训练时的原值> -``` - -checkpoint 已恢复 optimizer 和 scheduler 状态,所以 warmup 不会从头开始,当前 learning rate 和 scheduler 计数从 checkpoint 继续。用户不需要手算剩余步数,也不需要改其他训练超参;正常情况下只新增/替换 `DRAFTER_PATH`。 - -训练循环不再用本次进程的 `successful_steps < max_steps` 判断结束,而使用恢复后的总步数: - -```python -while max_steps <= 0 or trainer.optimizer_steps_total < max_steps: - ... -``` - -launcher计算 producer 配额时使用: - -```python -remaining_steps = max(max_steps - resumed_optimizer_step, 0) -max_samples = remaining_steps * batch_size_per_gpu * world_size -``` - -但 producer应使用 `queued_count` 判断本次配额。跳过的已消费编号不计入 `queued_count`。 - -该行为要求传入的是完整训练 checkpoint,里面具有 optimizer 和 scheduler 状态。只有模型权重的目录只能作为初始化权重,不能保证 learning rate、warmup 和 optimizer 状态精确续接。 - -## 10. 输入文件校验 - -编号只有在输入文件内容和顺序不变时才稳定。checkpoint至少保存: - -```python -"input_fingerprint": { - "path": str, - "size_bytes": int, - "mtime_ns": int, -} -``` - -推荐增加 SHA-256。续训时 fingerprint不一致则默认报错,避免旧 `sequence_no` 对应到新数据。 - -## 11. 崩溃语义 - -- vLLM已完成但未训练:不在 checkpoint集合,重启后重新推理。 -- 已写 TQ但未训练:不在集合,重启后重新推理。 -- optimizer成功但新 checkpoint未完成:恢复上一个模型和集合,重新训练上个 checkpoint之后的数据。 -- checkpoint完整:模型状态和 consumed集合对应同一步,已消费数据会被跳过。 -- checkpoint写到一半:恢复时忽略不完整目录,回退上一个 `complete=true` checkpoint。 - -## 12. 文件修改清单 - -### 新增 `verl_speco/trainer/standalone_resume.py` - -```python -load_consumed_sequence_nos(path) -> set[int] -save_consumed_sequence_nos(path, values) -> dict -build_input_fingerprint(path) -> dict -validate_input_fingerprint(saved, current) -> None -``` - -### 修改 `verl_speco/trainer/tq_sample_source.py` - -- `TQLocalBatch.global_sequence_nos`; -- rank 0从 selected tags提取编号。 - -### 修改 `verl_speco/trainer/draft_training_loop.py` - -- 启动时加载集合; -- clear成功后更新集合; -- checkpoint完成后写排序 int64 snapshot和 standalone sidecar; -- 使用 `optimizer_steps_total < max_steps` 作为独立训练终止条件。 - -`base_trainer.py` 及 co-train 入口不修改。standalone sidecar 的校验和加载全部收敛在 `standalone_resume.py` 与独立训练入口中。 - -### 修改 `verl_speco/standalone_tq_producer.py` - -- 加载 consumed集合; -- 拆分 source sequence和 queued count; -- tokenizer和vLLM前排除已消费编号; -- 其他并发逻辑保持不变。 - -### 修改 `verl_speco/standalone_tq_training_launcher.py` - -- 从 resume checkpoint获取 consumed文件; -- 把路径传给 producer; -- 校验输入 fingerprint。 -- 从 checkpoint读取已恢复 optimizer step,并仅生产剩余总步数需要的样本。 - -### 修改配置和 example - -- 增加 `consumed_sequence_path`; -- 说明只需设置 `DRAFTER_PATH=draft_step_N`;`MAX_STEPS`、warmup等保持首次训练配置。 - -## 13. 测试 - -1. consumer训练乱序编号 batch,clear成功后全部加入集合。 -2. 训练失败或 clear失败时不更新集合。 -3. checkpoint文件排序、去重,metadata数量一致。 -4. 输入 `0~9`,集合 `{0,2,5}`,producer只请求 `1,3,4,6,7,8,9`。 -5. 跳过的数据不占本次 `max_samples`配额。 -6. 多 epoch编号连续且重启后稳定。 -7. vLLM worker数量和请求并发与修改前一致。 -8. 端到端保存、重启后,producer不再请求 checkpoint集合中的编号。 - -## 14. 最终流程 - -首次训练: - -```text -producer并行请求 vLLM -→ TQ -→ consumer训练 batch -→ clear成功 -→ rank 0记录 sequence_no -→ checkpoint保存模型、optimizer和 consumed_sequence_nos.pt -``` - -续训: - -```text -加载 draft_step_N -→ 恢复模型、optimizer、LR和 step -→ 加载 consumed_sequence_nos.pt -→ 校验输入文件 -→ 创建新 Ray/TQ run -→ producer跳过已消费编号 -→ 其余样本继续并行请求 vLLM -``` - -该方案不改变 producer并发或 TQ消费顺序,改动集中在“consumer记录已成功数据”和“producer排除已消费数据”,适合作为第一版实现。 diff --git a/docs/standalone_tq_foundation_implementation.md b/docs/standalone_tq_foundation_implementation.md deleted file mode 100644 index c7d39fd9..00000000 --- a/docs/standalone_tq_foundation_implementation.md +++ /dev/null @@ -1,1030 +0,0 @@ -# Standalone TQ 公共基础层实现说明 - -Last updated: 08/21/2026 - -## 1. 文档范围和已验证结论 - -本文解释当前仓库已经实现并测试通过的 TQ 公共基础层: - -```text -verl_speco/transport/drafter_sample_protocol.py -verl_speco/integration/transferqueue_bridge.py -verl_speco/config/speco_base.yaml -verl_speco/tq_owner.py -tests/unit/test_drafter_sample_protocol.py -tests/unit/test_transferqueue_bridge.py -pyproject.toml -``` - -当前已经实现: - -1. Producer 和 Consumer 共用的样本 key、tag、fields 和 metadata 协议; -2. 普通进程连接 Ray 集群; -3. TQ Owner 创建 named `TransferQueueController`; -4. 独立 Client 发现并连接同一个 Controller; -5. 单样本 put、元数据 list、批量 get 和批量 clear; -6. Owner 与 Client 不同的关闭边界; -7. 独立 Owner 入口和共享 Hydra 配置; -8. mock 单元测试和真实双进程 TQ 0.1.7 smoke test。 - -当前还没有实现: - -1. Producer 读取 JSONL、并发访问多个 vLLM endpoint 的完整 pipeline; -2. `feature_store.type=tq` 工厂分支; -3. `TQFeatureStore` 和 `TQFeatureDataLoader`; -4. rank 0 选择 global keys、各 rank 读取 local keys; -5. TQ batch 接入 DSpark optimizer step; -6. optimizer step 成功后的 rank 0 clear。 - -因此当前代码已经证明“两个独立进程能通过 Ray 连接同一个 TQ,并按共享协议批量传输多个样本”,但尚未接到正式 Producer 和 DSpark Consumer 主循环。 - -## 2. 运行时角色和术语 - -### 2.1 Ray head - -Ray head 是 Ray 集群的控制节点。TQ 0.1.7 使用 Ray 管理 named Controller 和 storage actors。 - -Ray 不传输 standalone hidden-state payload。当前代码没有对这些样本调用: - -```python -ray.put(hidden_states) -``` - -### 2.2 TQ Owner - -TQ Owner 是普通 Python OS 进程,入口为: - -```text -python -m verl_speco.tq_owner -``` - -它不是 Ray actor。它先调用 `ray.init(address=...)` 加入 Ray,再调用带完整配置的 `tq.init(config)` 创建 TQ Controller 和 storage。 - -Owner 是唯一允许调用全局 `tq.close()` 的进程。 - -### 2.3 Named TransferQueueController - -TQ 0.1.7 内部创建: - -```python -TransferQueueController.options( - name="TransferQueueController" -).remote(...) -``` - -`name` 将 Controller actor 注册到 Ray actor registry。其他进程加入同一 Ray address 和 namespace 后,通过: - -```python -ray.get_actor("TransferQueueController") -``` - -取得 actor handle,再读取 TQ backend 配置。 - -Controller 保存控制信息和数据位置;使用 MooncakeStore 时,大 tensor 本身保存在 MooncakeStore。 - -### 2.4 TQ Client - -Owner、Producer、每个 Consumer rank 都在各自 OS 进程内拥有独立 TQ Client。 - -普通 Client 也传入相同的 native 配置: - -```python -tq.init(native_config) -``` - -发现 named Controller并初始化本进程的 storage manager。Client 不是 Ray actor,Producer 和 torchrun rank 也不需要改成 Ray actor。 - -### 2.5 Partition、key、tag 和 fields - -当前固定 partition: - -```text -speco_drafter_features -``` - -TQ 中一条记录逻辑上是: - -```text -partition_id -└── key - ├── tag:轻量 dict,由 kv_list 发现 - └── fields:Tensor payload,由 kv_batch_get 读取 -``` - -## 3. 共享配置如何工作 - -共享配置定义在 `verl_speco/config/speco_base.yaml`: - -```yaml -transfer_queue: - enable: false - package_version: "0.1.7" - ray: - address: null - namespace: speco-drafter - partition_id: speco_drafter_features - run_id: null - schema_version: 1 - connect_timeout_seconds: 120 - poll_interval_seconds: 0.5 - drop_last: true - controller: - polling_mode: true - backend: - storage_backend: SimpleStorage - SimpleStorage: - total_storage_size: 100000 - num_data_storage_units: 8 - MooncakeStore: - auto_init: false - metadata_server: localhost:50050 - master_server_address: localhost:50051 - local_hostname: "" - protocol: tcp - global_segment_size: 4294967296 - local_buffer_size: 1073741824 - device_name: "" -``` - -### 3.1 Ray 连接字段 - -```yaml -ray: - address: 10.0.0.1:6379 - namespace: speco-drafter -``` - -它们决定当前进程连接哪个 Ray 集群,以及在哪个 namespace 查找 `TransferQueueController`。Owner、Producer 和所有 Consumer ranks 必须使用相同值。 - -### 3.2 SPECO 协议字段 - -```yaml -partition_id: speco_drafter_features -run_id: dspark-20260819-a -schema_version: 1 -``` - -这些字段用于路由和校验,不属于 TQ 原生配置。`run_id` 用于隔离不同 pipeline 的记录。 - -### 3.3 TQ 原生字段 - -```yaml -controller: ... -backend: ... -``` - -只有这些字段应传给 `tq.init(full_config)`。bridge 的 `_native_tq_config()` 会删除: - -```text -enable -package_version -ray -partition_id -run_id -schema_version -connect_timeout_seconds -poll_interval_seconds -drop_last -``` - -对象变化为: - -```text -完整 SPECO transfer_queue dict -→ _native_tq_config() -→ controller/backend等TQ字段 -→ OmegaConf DictConfig -→ tq.init(same native config) -``` - -## 4. Bridge 的进程内状态 - -`verl_speco/integration/transferqueue_bridge.py` 在每个 OS 进程内分别维护: - -```python -_state = { - "enabled": False, - "configured": False, - "initialized": False, - "config": None, - "owner": False, - "ray_initialized_here": False, - "ray_address": None, - "ray_namespace": None, -} -``` - -该 dict 不跨进程共享。Owner、Producer、每个 rank 分别拥有自己的 `_state`。 - -| 字段 | 含义 | -|---|---| -| `enabled` | 当前进程配置是否开启 TQ | -| `configured` | 是否调用过 `configure_transfer_queue()` | -| `initialized` | 当前进程是否执行过 `tq.init()` | -| `config` | 当前进程保存的普通 dict 配置 | -| `owner` | 当前进程是否创建了全局 Controller/Storage | -| `ray_initialized_here` | bridge 是否负责调用了本进程的 `ray.init()` | -| `ray_address/namespace` | 本进程的 Ray 连接信息 | - -`_state_lock` 只保护同一进程内多个线程同时初始化,不是分布式锁。 - -## 5. Owner 的完整启动数据流 - -Owner 入口是 `verl_speco/tq_owner.py`,直接加载 Hydra 主配置 -`verl_speco/config/speco_base.yaml`。`run_owner()` 会复制其中的 -`transfer_queue` 子配置,并仅在 Owner 进程的副本中强制设置 -`enable=true`;普通训练任务看到的共享默认值仍是 `false`。 - -### 阶段 1:读取配置 - -执行者:Owner OS 进程。 - -入口: - -```python -run_owner(config) -``` - -取得: - -```python -training_cfg = config.actor_rollout_ref.rollout.drafter.training -tq_cfg = training_cfg.transfer_queue -``` - -然后调用: - -```python -configure_transfer_queue(training_cfg) -``` - -该函数只把 OmegaConf 转成普通 dict并更新当前进程 `_state`,不会连接 Ray,也不会创建 TQ。 - -### 阶段 2:连接 Ray - -Owner 调用: - -```python -connect_ray_cluster(ray_address, namespace) -``` - -内部执行: - -```python -if not ray.is_initialized(): - ray.init(address=ray_address, namespace=namespace) -``` - -边界类型是 Ray control-plane connection。此时还没有传输训练 tensor。 - -### 阶段 3:创建 Controller 和 Storage - -Owner 调用: - -```python -start_transfer_queue_owner(tq_cfg) -``` - -执行顺序: - -1. `_extract_tq_config()` 得到普通 dict; -2. 检查 `enable=true`; -3. 检查 `TransferQueue` 包可用; -4. 防止本进程重复初始化; -5. `_native_tq_config()` 删除 SPECO 字段; -6. `_as_tq_config()` 转 OmegaConf; -7. `tq.init(native_config)` 创建 Controller、Storage 和 Owner Client; -8. 设置 `_state.owner=True`、`initialized=True`。 - -Ray 中形成: - -```text -Ray cluster / namespace -├── named actor: TransferQueueController -└── storage backend - ├── SimpleStorage actors - └── 或 MooncakeStore connection/process -``` - -### 阶段 4:发布 owner-ready - -调用: - -```python -publish_owner_ready(run_id, schema_version) -``` - -生成: - -```python -key = "control:v1::owner-ready" -fields = {"marker": torch.tensor([1], dtype=torch.uint8)} -tag = { - "record_type": "control", - "status": "owner_ready", - "schema_version": 1, - "run_id": run_id, -} -``` - -这是一条控制记录,不进入训练 batch。 - -### 阶段 5:常驻和关闭 - -Owner 安装 `SIGINT/SIGTERM` handler,并等待: - -```python -stop_event.wait() -``` - -收到信号后调用: - -```python -close_transfer_queue_owner() -``` - -执行: - -```text -验证owner身份 -→ tq.close() -→ 清理Controller/Storage -→ ray.shutdown() -``` - -Owner 必须在 Producer 和 Consumer 退出后才能关闭。 - -## 6. 普通 Client 如何连接同一个 TQ - -Producer 和每个 Consumer rank 后续使用相同顺序: - -```python -configure_transfer_queue(training_cfg) -connect_ray_cluster(ray_address, namespace) -connect_transfer_queue_client() -``` - -`connect_transfer_queue_client()` 最终调用: - -```python -tq.init(same_native_config) -``` - -TQ 0.1.7 内部通过: - -```python -ray.get_actor("TransferQueueController") -``` - -找到 Owner 创建的 Controller,读取 backend 配置,然后创建当前进程的 TQ Client。 - -对象和边界变化: - -```text -actor名称字符串 -→ Ray actor registry -→ Controller actor handle -→ Controller.get_config.remote() -→ TQ DictConfig -→ 当前进程TransferQueueClient -→ 同一个SimpleStorage/MooncakeStore -``` - -## 7. 一条具体样本的初始对象 - -真实 smoke test使用: - -```python -sample = DraftFeatureSample( - algorithm="DSPARK", - input_ids=torch.tensor([1, 2, 3]), # int64[3], CPU - loss_mask=torch.tensor([0.0, 1.0, 1.0]), # float32[3], CPU - position_ids=torch.tensor([0, 1, 2]), # int64[3], CPU - hidden_states=torch.arange( - 12, dtype=torch.float32 - ).reshape(3, 4), # float32[3,4], CPU -) -``` - -同时构造: - -```python -meta = SampleMetadata( - schema_version=1, - run_id="codex-batch-smoke", - sample_id="smoke-0000", - sequence_no=0, - algorithm="DSPARK", - target_model_id="smoke-target", - target_model_revision="smoke-revision", - tokenizer_fingerprint="smoke-tokenizer", - target_layer_ids=[0], - hidden_states_layout="dflash_aux", - hidden_dtype="float32", - hidden_shape=[3, 4], - feature_length=3, - full_sequence_length=3, - feature_start=0, - feature_end=3, - use_logits=False, -) -``` - -`DraftFeatureSample` 和 `SampleMetadata` 都是进程内 Python 对象,不直接经过 TQ。 - -## 8. Key 的生成和两个同名函数 - -共享协议调用: - -```python -make_sample_key(meta) -``` - -输出: - -```text -drafter:v1:codex-batch-smoke:000000000000:smoke-0000 -``` - -字段顺序: - -```text -drafter / schema version / run_id / 12位sequence_no / sample_id -``` - -`sequence_no` 是输入文件 record 序号,不是训练 batch 编号。 - -bridge 为兼容 PR #48 还保留另一个: - -```python -transferqueue_bridge.make_sample_key( - global_step, - replica_rank, - request_id, -) -``` - -它生成: - -```text -speco::: -``` - -standalone Producer 必须从 `verl_speco.transport.drafter_sample_protocol` import `make_sample_key`,不能使用 bridge 中的 PR #48 旧函数。 - -## 9. Tag 如何生成 - -```python -tag = make_ready_tag(meta) -``` - -输出: - -```python -{ - "record_type": "sample", - "status": "ready", - "schema_version": 1, - "run_id": "codex-batch-smoke", - "sequence_no": 0, - "sample_id": "smoke-0000", - "algorithm": "DSPARK", -} -``` - -tag 只包含发现、过滤和排序需要的标量/字符串。`list_samples()` 只返回 key/tag,不读取 hidden states。 - -## 10. `encode_sample()` 如何生成 fields - -调用: - -```python -fields = encode_sample(sample, meta) -``` - -### 10.1 校验 - -执行: - -```text -SampleMetadata.validate() -DraftFeatureSample.validate(strict=True) -``` - -随后检查: - -1. hidden states 是一个 dense tensor; -2. ids/mask/position 长度等于 `feature_length`; -3. hidden 第一维等于 `feature_length`; -4. hidden shape 等于 metadata; -5. hidden dtype 等于 metadata; -6. feature window 长度正确。 - -### 10.2 Tensor 规范化 - -```text -input_ids → CPU contiguous int64[L] -loss_mask → CPU contiguous float32[L] -position_ids → CPU contiguous int64[L] -hidden_states → CPU contiguous,保持模型dtype -``` - -没有 `position_ids` 时生成 `torch.arange(L, dtype=int64)`。 - -### 10.3 Metadata JSON 编码 - -```text -SampleMetadata dataclass -→ dict -→ JSON UTF-8 bytes -→ torch.uint8[M] -``` - -实现等价于: - -```python -raw = json.dumps(metadata).encode("utf-8") -metadata_json = torch.tensor(list(raw), dtype=torch.uint8) -``` - -### 10.4 最终 fields - -```python -fields = { - "input_ids": int64[3], - "loss_mask": float32[3], - "position_ids": int64[3], - "hidden_states": float32[3,4], - "metadata_json": uint8[M], -} -``` - -如果存在,协议也保留 `last_hidden_states`、`target` 和 `target_logprobs`。 - -## 11. Bridge 如何写入 TQ - -调用: - -```python -put_sample(key, fields, tag=tag) -``` - -bridge 执行: - -1. 检查 TQ 已启用; -2. 丢弃 fields 中非 tensor 值; -3. 确保本进程已经使用相同 native 配置执行 `tq.init(config)`; -4. 取得配置中的 partition; -5. 调用: - -```python -tq.kv_put( - key=key, - partition_id="speco_drafter_features", - fields=fields, - tag=tag, -) -``` - -使用 MooncakeStore 时,大 tensor 路径是: - -```text -Producer CPU tensor -→ Producer TQ Client -→ MooncakeStore -``` - -不是 Ray `ObjectRef`。 - -## 12. Consumer 如何发现 key - -调用: - -```python -records = list_samples() -``` - -内部调用: - -```python -tq.kv_list(partition_id="speco_drafter_features") -``` - -标准化返回类型: - -```python -dict[str, dict[str, Any]] -``` - -示例: - -```python -{ - "drafter:v1:...:smoke-0000": { - "record_type": "sample", - "status": "ready", - "run_id": "codex-batch-smoke", - "sequence_no": 0, - ... - } -} -``` - -bridge 兼容 `key → tag` 和 `partition → key → tag` 两种 wrapper。这一步不读取 fields。 - -## 13. Consumer 如何批量取样本 - -输入: - -```python -keys = [key0, key1] -``` - -调用: - -```python -records = get_samples(keys) -``` - -bridge 只调用一次: - -```python -result = tq.kv_batch_get( - keys=keys, - partition_id="speco_drafter_features", -) -``` - -TQ 0.1.7 返回带 batch 维的 TensorDict。bridge 检查 `result.batch_size`,再执行: - -```python -rows = [result[index] for index in range(len(keys))] -``` - -每行转成普通 dict,最终返回: - -```python -[ - (key0, fields0), - (key1, fields1), -] -``` - -返回顺序与输入 keys 一致。重复 key 会提前报错。 - -## 14. `decode_sample()` 如何恢复训练对象 - -调用: - -```python -sample = decode_sample( - key, - tag, - fields, - expected_config, -) -``` - -### 14.1 Metadata 解码 - -```text -metadata_json uint8[M] -→ bytes -→ UTF-8 -→ json.loads -→ dict -→ SampleMetadata.from_dict -``` - -### 14.2 身份一致性 - -代码根据 metadata 重新生成 key,要求输入 key 完全相等;然后逐项校验 tag 的: - -```text -record_type/status/schema_version/run_id/sequence_no/sample_id/algorithm -``` - -所以 key、tag 和 payload metadata 不能来自不同样本。 - -### 14.3 Consumer 合同 - -Consumer 提供: - -```python -ExpectedFeatureConfig( - run_id="codex-batch-smoke", - schema_version=1, - algorithm="DSPARK", - target_model_id="smoke-target", - target_model_revision=None, - tokenizer_fingerprint=None, - target_layer_ids=None, - hidden_states_layout="dflash_aux", - hidden_dtype="float32", -) -``` - -值为 `None` 的字段不检查;其他字段必须完全一致。正式 Consumer 应填写 target checkpoint、tokenizer、layers、layout 和 dtype,避免使用错误 target 特征。 - -### 14.4 输出 - -完成 tensor 类型、长度、shape、dtype 校验后,构造: - -```python -DraftFeatureSample.from_dict(payload, strict=True) -``` - -输出可以交给现有 `trainer.prepare_training_batch_from_samples()`。当前 smoke test验证到这里,正式 `TQFeatureDataLoader` 尚未实现。 - -## 15. EOS 控制记录 - -调用: - -```python -key, fields, tag = make_eos_record(run_id, total_samples) -``` - -输出: - -```python -key = "control:v1::eos" -fields = {"marker": torch.tensor([1], dtype=torch.uint8)} -tag = { - "record_type": "control", - "status": "eos", - "schema_version": 1, - "run_id": run_id, - "total_samples": total_samples, -} -``` - -EOS 表示不会再发布新样本。Consumer 应先 drain ready samples 再退出。协议函数已实现,正式 Producer/Consumer 尚未调用。 - -## 16. Clear 和数据生命周期 - -bridge 提供: - -```python -clear_samples(keys) -``` - -内部调用: - -```python -tq.kv_clear(keys=keys, partition_id="speco_drafter_features") -``` - -基础层只执行删除,不决定删除时机。正式 Consumer 必须遵守: - -```text -rank 0选择global keys -→ 各rank读取local keys -→ 所有rank完成同一optimizer step -→ 汇总global success -→ rank 0 clear global keys -``` - -不能在 `get_samples()` 后立即 clear,因为 get 成功不代表训练 step 成功。 - -## 17. Client close 和 Owner close - -### 17.1 Client close - -Producer/rank 调用: - -```python -close_transfer_queue_client() -``` - -执行: - -```text -tq.get_client() -→ 当前进程client.close() -→ 如果bridge负责ray.init,则ray.shutdown() -``` - -它不调用全局 `tq.close()`,不会 kill Controller。调用后本进程不能继续使用 TQ。 - -### 17.2 Owner close - -Owner 调用: - -```python -close_transfer_queue_owner() -``` - -执行: - -```text -tq.close() -→ Controller/Storage全局清理 -→ ray.shutdown() -``` - -Owner 若误用 Client close,bridge 会抛 `RuntimeError`。 - -## 18. PR #48 兼容边界 - -bridge 继续保留: - -```python -init_transfer_queue(config) -get_sample(key) -close_transfer_queue() -``` - -PR #48 的 `SpecoTaskRunner` 已是 Ray actor,所以不调用 standalone 的 `connect_ray_cluster()`。旧 worker 继续使用单 key `get_sample()`。 - -standalone 后续使用新增的: - -```python -list_samples() -get_samples(keys) -clear_samples(keys) -``` - -因此没有修改 PR #48 现有调用点的函数签名。 - -## 19. 依赖和命令入口 - -`pyproject.toml` 新增: - -```toml -[project.optional-dependencies] -transfer-queue = ["TransferQueue==0.1.7"] -``` - -安装: - -```bash -pip install -e ".[transfer-queue]" -``` - -TQ 0.1.7 会安装 Ray,并要求 `numpy<2.0.0`。若现有包要求 NumPy 2,需要隔离环境或重新解决依赖。 - -Owner 命令: - -```text -verl-speco-tq-owner -``` - -## 20. 单元测试 - -协议测试 `tests/unit/test_drafter_sample_protocol.py` 覆盖: - -1. encode/decode round trip; -2. key 格式; -3. tag 身份冲突; -4. Consumer contract 冲突; -5. hidden shape 冲突; -6. EOS 格式。 - -bridge 测试 `tests/unit/test_transferqueue_bridge.py` 覆盖: - -1. Ray address/namespace 参数; -2. Owner 只向 TQ 传原生配置; -3. Client 使用相同 native 配置调用 `tq.init(config)`; -4. put/list/get-many/clear; -5. batch 返回顺序; -6. Client close 不调用全局 close; -7. Owner 不能误用 Client close。 - -运行: - -```bash -python -m pytest \ - tests/unit/test_drafter_sample_protocol.py \ - tests/unit/test_transferqueue_bridge.py \ - -q -``` - -## 21. 真实双进程 smoke test - -早期用于该验证的临时双进程工具已经移除;正式入口统一由 -`verl_speco.standalone_tq_training_launcher` 管理 Owner、Producer 和 Consumer 生命周期。 - -Owner 路径: - -```text -连接Ray -→ tq.init(full config) -→ 写sample 0和sample 1 -→ 等待client-done -→ clear done marker -→ 全局关闭 -``` - -Client 路径: - -```text -连接同一个Ray -→ tq.init(same native config) -→ kv_list发现两个key -→ 一次kv_batch_get([k0,k1]) -→ 拆成两个fields dict -→ 分别decode_sample -→ clear两个sample keys -→ 写client-done -→ 只关闭本地client -``` - -已验证输出: - -```text -OWNER_READY keys=[k0, k1] -CLIENT_OK samples=2 shape=(3, 4) -CLIENT_CLOSED_LOCAL_ONLY -OWNER_OBSERVED_SAMPLES_CLEARED -OWNER_CLOSED -``` - -这证明: - -1. 两个普通进程能连接同一个 TQ; -2. named Controller 发现有效; -3. 0.1.7 的 `kv_list/kv_batch_get/kv_clear` 参数有效; -4. TensorDict batch 能按 key 顺序拆开; -5. 共享协议能恢复 `DraftFeatureSample`; -6. Client close 不会杀掉 Owner; -7. Owner 能最终统一关闭。 - -## 22. 当前完整路径总结 - -```text -Owner -→ ray.init(address, namespace) -→ tq.init(native config) -→ named TransferQueueController - -普通Client -→ ray.init(same address, same namespace) -→ tq.init(same native config) -→ 找到同一个Controller - -DraftFeatureSample + SampleMetadata -→ make_sample_key -→ make_ready_tag -→ encode_sample -→ fields + metadata_json tensor -→ bridge.put_sample -→ tq.kv_put -→ SimpleStorage/MooncakeStore - -Consumer/测试Client -→ bridge.list_samples -→ key + tag -→ bridge.get_samples(keys) -→ tq.kv_batch_get -→ TensorDict batch -→ 每个key对应一个fields dict -→ decode_sample -→ DraftFeatureSample - -正式训练成功后(待实现) -→ bridge.clear_samples(global_keys) - -Client退出 -→ close_transfer_queue_client - -所有业务进程退出 -→ Owner close_transfer_queue_owner -→ tq.close -→ ray.shutdown -``` - -## 23. 下一阶段接入约束 - -后续代码不能重新定义协议或直接访问 TQ 私有对象。 - -Producer 应复用: - -```text -SampleMetadata -make_sample_key -make_ready_tag -encode_sample -bridge.put_sample -make_eos_record -``` - -Consumer 应复用: - -```text -bridge.list_samples -bridge.get_samples -decode_sample -bridge.clear_samples -``` - -下一阶段需要新增: - -```text -verl_speco/trainer/tq_feature_store.py -verl_speco/trainer/tq_sample_source.py -feature_store.py 的 type=tq 分支 -draft_training_loop.py 的流式训练分支 -Producer入口、输入读取和并发vLLM文件 -``` - -这些文件应建立在本文已经实现和真实验证过的连接、协议、KV 和关闭接口之上。 diff --git a/docs/standalone_tq_producer.md b/docs/standalone_tq_producer.md deleted file mode 100644 index bb766907..00000000 --- a/docs/standalone_tq_producer.md +++ /dev/null @@ -1,212 +0,0 @@ -# Standalone vLLM → TransferQueue Producer - -Last updated: 08/21/2026 - -本文解释 standalone Producer:它可以直接读取 verl 的 prompt-only Parquet(包括 -DAPO-Math-17k 的 chat-message `prompt`),也兼容已有 `prompt`/`response` 的 JSONL -或 Parquet。缺少 response 时由 target vLLM 生成,并在同一请求中提取 prompt 与 -output hidden states,之后把样本写到已存在的 TransferQueue(TQ)。 - -这条路径面向第一版 DSpark standalone 训练:Producer、TQ owner 和 Consumer -是三个独立 OS 进程;Ray 只用于让它们找到同一个 TQ Controller,hidden states -不通过 Ray object store 传输。 - -## 为什么需要这个 Producer - -此前仓库已经有两块基础能力: - -- `drafter_sample_protocol.py`:规定一条 TQ sample 的 key、tag、Tensor 字段和 - EOS record; -- `transferqueue_bridge.py` 与 `tq_owner.py`:负责连接 Ray/TQ、写读清理样本和 - owner 生命周期。 - -缺少的是把预先生成的文本变成 DSpark 训练特征并发布到 TQ 的独立进程。新增的 -Producer 补上这一段,不引入第二套协议或 feature store。 - -## 数据流 - -```text -verl prompt Parquet 或 prompt/response JSONL/Parquet - │ - │ 按文件顺序分配 sequence_no 和 sample_id - ▼ -Tokenizer - │ input_ids / loss_mask / feature window - ▼ -多个 vLLM endpoint(有界并发) - │ OpenAI completions 请求 → 临时 safetensors 文件 - ▼ -公共 hidden-state 转换函数 - │ DSpark DraftFeatureSample + SampleMetadata - ▼ -TransferQueue kv_put(一条输入记录对应一条 sample) - │ - ├─ put 成功:删除该请求的临时文件 - └─ 全部成功:写一个 EOS control record -``` - -Producer 在开始请求前会等到对应 `run_id` 的 `owner_ready` 控制记录。它不会创建 -Ray head、TQ Controller 或 storage backend;这些由 `verl-speco-tq-owner` 管理。 - -## 输入文件 - -输入可以是 JSONL 或 Parquet。`prompt` 可以是字符串,也可以是 verl 常用的 -`[{"role": ..., "content": ...}]` chat-message 列表。`response` 是可选字符串: -存在时直接 replay;不存在时由 target vLLM 生成。Parquet 通过 -`data.train_files` 直接传入,不需要转换。 - -```json -{"sample_id":"train-000017","prompt":"Question: 1 + 1 = ","response":"2"} -{"prompt":"Translate hello: ","response":"你好"} -``` - -- `sequence_no` 按非空行的文件顺序从 0 分配;并发完成顺序不会影响它。 -- `sample_id` 可选;省略时生成 `train-000000`、`train-000001` 等稳定值。 -- verl 数据的 `extra_info.index` 存在时会优先作为稳定 `sample_id`。 -- chat-message prompt 通过 target tokenizer 的 `apply_chat_template()` 编码,并加上 - generation prompt;不能把 `reward_model.ground_truth` 当作模型 response。 -- Producer tokenize `prompt` 和 `prompt + response`。后者必须以 prompt 的 token IDs - 为前缀;否则会报错,而不会猜测 response 的 loss-mask 边界。 -- `loss_mask` 中 prompt token 为 0,response token 为 1。 -- feature window 从 response 前一个 token 开始,长度由 - `max_feature_length` 限制;传给 vLLM 的 token IDs 截止于该 window 末端。 -- 其他 JSON 字段目前只作为 Producer 进程内来源元数据;第一版协议不会把它们写入 - TQ,所以 Consumer 不能读取这些字段。 - -## vLLM 与 hidden states - -对已有 response,Producer 使用 OpenAI-compatible completions API 做 prefill。 -对 prompt-only 数据,Producer 在一次请求中生成 response 并要求保存输出 hidden: - -```text -prompt= -max_tokens=<内部有界长度> -extra_body={ - "return_token_ids": true, - "kv_transfer_params": {"include_output_tokens": true} -} -``` - -响应必须同时满足: - -1. 若返回 `choices[0].prompt_token_ids`,它必须等于请求的 token IDs; -2. `kv_transfer_params.hidden_states_path` 必须存在; -3. 该文件必须含 `token_ids` 和形状为 `[seq, layers, hidden]` 的 `hidden_states`。 - -vLLM 0.23 已内置满足这个合同的 `ExampleHiddenStatesConnector`。不需要 SpeCo -Mooncake connector。在线服务必须关闭 chunked prefill,并显式配置一个 Producer -可见的临时目录。例如: - -```bash -export MODEL_PATH=/path/to/target-model -export HIDDEN_STATES_DIR=/dev/shm/speco-hidden-states -mkdir -p "${HIDDEN_STATES_DIR}" - -vllm serve "${MODEL_PATH}" \ - --host 0.0.0.0 \ - --port 8000 \ - --speculative-config \ - '{"method":"extract_hidden_states","num_speculative_tokens":1,"draft_model_config":{"hf_config":{"eagle_aux_hidden_state_layer_ids":[1,9,17,25,33,36]}}}' \ - --kv-transfer-config \ - "{\"kv_connector\":\"ExampleHiddenStatesConnector\",\"kv_role\":\"kv_producer\",\"kv_connector_extra_config\":{\"shared_storage_path\":\"${HIDDEN_STATES_DIR}\",\"use_synchronization_lock\":true}}" \ - --no-enable-chunked-prefill -``` - -上面的 layer IDs 只是 Qwen3-4B 示例。实际值必须按 target 模型和训练配置确定; -DSpark L1 开启时,vLLM 列表是 auxiliary layer IDs 加 final layer,而 Producer 的 -`TARGET_LAYER_IDS` 只填写 auxiliary 部分。 - -官方 connector 使用持久存在的 `.lock` 文件和 `flock` 协调异步落盘。Producer -读取前等待文件锁释放;TQ `put_sample` 成功后同时删除 safetensors 和 `.lock`。 - -`feature_from_vllm_payload()` 是从旧 replay 路径提取出的公共纯函数。它校验 token -对齐、选择 feature rows、拼接 auxiliary layers;DSpark L1 开启时额外拼接 final -hidden state。旧 replay 路径仍通过薄封装调用此函数,避免两套转换规则。 - -## TQ 写入和失败语义 - -每个输入 record 只写一个协议 key: - -```text -drafter:v1::<12位sequence_no>: -``` - -写入顺序是严格的: - -```text -加载临时 safetensors -→ 校验并转换 -→ TQ kv_put -→ 删除临时文件 -``` - -因此: - -- `kv_put` 失败时临时文件保留,且 Producer 不写 EOS; -- 任一请求、转换或写入失败会停止整条 Producer,不做自动重试或 endpoint 熔断; -- 只有所有 sample 都发布完成,才写 `control:v1::eos`; -- 进程退出时只调用 `close_transfer_queue_client()`,不会调用全局 `tq.close()`, - 不会销毁共享 Controller。只有 owner 可以关闭 TQ。 - -`max_pending_samples` 是简单背压:当前 run 的 ready sample 数达到该阈值时,新的 -vLLM 请求会暂停,等待 Consumer 清理已成功训练的 key。 - -## 配置与启动 - -默认 Producer 配置位于 -`speco.standalone_tq_producer`,TQ 连接配置仍位于 -`actor_rollout_ref.rollout.drafter.training.transfer_queue`。 - -必须设置的 Producer 字段: - -| 字段 | 含义 | -| --- | --- | -| `input_path` | 上述 JSONL 或 Parquet 文件 | -| `tokenizer_path` / `tokenizer_fingerprint` | 用于 tokenization 和 Consumer 合同校验 | -| `target_model_id` / `target_model_revision` | target checkpoint 身份 | -| `target_layer_ids` | auxiliary target layer IDs;DSpark L1 时 wire metadata 会额外写 `-1` 表示 final layer | -| `vllm_endpoints` / `vllm_model` | 一个或多个 OpenAI-compatible vLLM endpoint 与模型名 | - -必须与 owner/Consumer 一致的 TQ 字段: - -| 字段 | 固定要求 | -| --- | --- | -| `package_version` | `0.1.7` | -| `partition_id` | `speco_drafter_features` | -| `schema_version` | `1` | -| `run_id`、Ray address、Ray namespace | 三个进程必须相同 | - -单独调试时可以通过安装后的命令入口运行;正式训练由统一launcher启动Producer: - -```bash -verl-speco-tq-producer \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.enable=true \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.ray.address= \ - actor_rollout_ref.rollout.drafter.training.transfer_queue.run_id= \ - speco.standalone_tq_producer.input_path= \ - speco.standalone_tq_producer.tokenizer_path= \ - speco.standalone_tq_producer.tokenizer_fingerprint= \ - speco.standalone_tq_producer.target_model_id= \ - speco.standalone_tq_producer.target_model_revision= \ - speco.standalone_tq_producer.target_layer_ids='[2,8,14,20,26]' \ - speco.standalone_tq_producer.vllm_endpoints='[http://node0:8000/v1]' \ - speco.standalone_tq_producer.vllm_model= -``` - -完整生命周期顺序仍是:Ray/TQ backend → TQ owner → Consumer → Producer → Consumer -drain → owner shutdown。Producer 完成不代表训练完成,EOS 只表示不会再有新样本。 -正式独立训练入口 -`examples/run_qwen3-8b_drafter_separate_training.sh` 会通过 -`verl_speco.standalone_tq_training_launcher` 自动管理这套生命周期;上面的 Producer -脚本仅用于单独调试 Producer。 - -## 测试覆盖与未验证项 - -新增测试覆盖:JSONL/真实 Parquet 解析、DAPO chat prompt、target response generation、 -token 边界、多个 endpoint 的并发限制、ready 队列背压、 -成功时 sample 后 EOS 与临时文件删除、失败时无 EOS 且保留临时文件,以及旧 EAGLE3 -转换路径仍可复用公共函数。 - -这些测试使用 fake vLLM/TQ。真实 Ray + TransferQueue + vLLM 的多进程 -联调没有在当前环境执行;运行前仍需确认 vLLM 版本能返回上述 -`hidden_states_path` 以及 TQ 0.1.7 依赖环境可用。 diff --git a/docs/standalone_tq_training_parameters.md b/docs/standalone_tq_training_parameters.md deleted file mode 100644 index a7a48988..00000000 --- a/docs/standalone_tq_training_parameters.md +++ /dev/null @@ -1,169 +0,0 @@ -# Standalone TQ 独立训练参数说明 - -> Last updated: 08/27/2026 - -本文说明下面两个脚本暴露的参数: - -- `tools/run_qwen3-8b_drafter_hidden_state_vllm.sh`:按可见设备与 TP 自动启动一个或多个 target vLLM 服务。 -- `examples/run_qwen3-8b_drafter_separate_training.sh`:启动 Producer、TQ 和 DSpark Consumer 训练。 - -参数可以通过环境变量设置,例如: - -```bash -MAX_STEPS=1000 \ -BATCH_SIZE_PER_GPU=2 \ -DSPARK_CE_LOSS_ALPHA=0.1 \ -DSPARK_L1_LOSS_ALPHA=0.9 \ -bash examples/run_qwen3-8b_drafter_separate_training.sh -``` - -## 1. vLLM 服务参数 - -以下参数由 `run_qwen3-8b_drafter_hidden_state_vllm.sh` 使用。 - -| 参数 | 默认值 | 作用 | -| --- | --- | --- | -| `MODEL_PATH` | `/path/to/Qwen3-8B` | 所有 target vLLM 服务加载的模型路径。 | -| `DEVICE_ENV` | `ASCEND_RT_VISIBLE_DEVICES` | 控制设备可见性的环境变量。GPU 环境可设为 `CUDA_VISIBLE_DEVICES`。 | -| `VLLM_DEVICES` | `0,1,2,3,4,5` | 分配给vLLM的完整设备列表,脚本按连续的 `VLLM_TP` 张设备切成多个实例。 | -| `VLLM_TP` | `1` | 每个vLLM实例的 tensor parallel 大小;设备总数必须能被该值整除。 | -| `VLLM_HOST` | `127.0.0.1` | 所有vLLM服务监听的主机地址。 | -| `VLLM_BASE_PORT` | `8000` | 第一个实例的端口;后续实例依次使用 `8001`、`8002` 等。 | -| `VLLM_GPU_MEMORY_UTILIZATION` | `0.8` | 单个 vLLM 服务允许使用的设备显存比例。 | -| `VLLM_MAX_NUM_SEQS` | `256` | 单个 vLLM 服务最多同时调度的 sequence 数。 | -| `VLLM_HIDDEN_STATE_LAYER_IDS` | `[1,9,17,25,33,36]` | vLLM 导出的 hidden-state 层;前面的辅助层必须与训练侧 `DSPARK_TARGET_LAYER_IDS` 相同,最后一层用于构造 L1 loss 所需的 target 概率分布。 | -| `HIDDEN_STATES_DIR` | `/tmp/speco-vllm-hidden-states` | vLLM connector 临时写 hidden-state 文件的根目录,每个实例使用独立的 `service-N` 子目录。 | - -如果每个服务使用两张卡: - -```bash -MODEL_PATH=/nas/disk1/Qwen3-4B \ -VLLM_DEVICES=0,1,2,3,4,5 \ -VLLM_TP=2 \ -bash tools/run_qwen3-8b_drafter_hidden_state_vllm.sh -``` - -## 2. 基础训练参数 - -以下参数由 `run_qwen3-8b_drafter_separate_training.sh` 使用。 - -| 参数 | 默认值 | 作用 | -| --- | --- | --- | -| `MODEL_PATH` | `/path/to/Qwen3-8B` | Target model 路径,同时用于 tokenizer、模型配置和 target embedding/LM head。 | -| `TRAIN_FILE` | `/path/to/train_file.parquet` | Producer 读取的 JSONL 或 Parquet 数据文件。 | -| `DRAFTER_PATH` | 空 | 可选的已有 drafter/checkpoint 路径;为空时根据 target 配置从头初始化 DSpark。 | -| `DRAFT_CKPTS_DIR` | `/path/to/dspark_draft_checkpoints` | 保存 drafter checkpoint 的目录。 | -| `TRAIN_DEVICES` | `2,3` | Consumer 训练使用的设备。不能与任何 vLLM 实例占用的设备重叠。 | -| `TRAIN_GPUS` | `2` | 本节点启动的训练 rank 数,通常等于 `TRAIN_DEVICES` 中的设备数量。 | -| `DEVICE_ENV` | `ASCEND_RT_VISIBLE_DEVICES` | 训练侧设备可见性环境变量;GPU 环境可改为 `CUDA_VISIBLE_DEVICES`。 | -| `SPECO_VLLM_ENDPOINTS` | `[http://127.0.0.1:8000/v1,http://127.0.0.1:8001/v1]` | Producer 并行访问的 vLLM endpoint 列表。 | -| `VLLM_READY_TIMEOUT_SECONDS` | `120` | 启动训练前等待所有 vLLM endpoint 就绪的最长时间。 | -| `PYTHON_BIN` | `python3` | 启动 Python 模块所用的解释器。 | -| `PROJECT_NAME` | `verl_dspark_drafter` | 实验项目名。 | -| `EXP_NAME` | `qwen3_8b_dspark_separate_training` | 本次实验名称。 | - -## 3. Producer和请求并发参数 - -| 参数 | 默认值 | 作用 | -| --- | --- | --- | -| `VLLM_REQUEST_TIMEOUT` | `120` | 单次 generate 或 prefill HTTP 请求的超时时间,单位为秒。 | -| `VLLM_MAX_INFLIGHT_REQUESTS` | `16` | Producer 整体最多同时存在的 vLLM 请求数。 | -| `VLLM_PER_ENDPOINT_CONCURRENCY` | `4` | 每个 vLLM endpoint 独立的并发请求上限。 | -| `PRODUCER_INPUT_QUEUE_SIZE` | `32` | 已读取、等待 vLLM worker 处理的请求队列容量。 | -| `PRODUCER_PUBLISH_QUEUE_SIZE` | `16` | 已完成推理、等待写入 TQ 的样本队列容量。 | -| `PRODUCER_MAX_PENDING_SAMPLES` | `1024` | TQ 中尚未被 Consumer 训练并删除的样本数量上限,用于限制积压。 | -| `PRODUCER_PENDING_POLL_INTERVAL` | `0.5` | TQ 积压达到上限后,Producer 重新检查容量的间隔,单位为秒。 | -| `PRODUCER_MAX_SEQUENCE_LENGTH` | `8192` | prompt 和 response 处理前允许的最大总 token 长度。 | -| `PRODUCER_MAX_FEATURE_LENGTH` | `512` | 每个样本最终保留用于训练的最大 token 窗口。 | -| `PRODUCER_GENERATION_MAX_TOKENS` | `512` | 输入没有 response 时,vLLM 最多生成的 completion token 数。 | - -两个 endpoint 下的有效客户端并发近似为: - -```text -min(VLLM_MAX_INFLIGHT_REQUESTS, - endpoint数量 × VLLM_PER_ENDPOINT_CONCURRENCY) -``` - -`VLLM_MAX_NUM_SEQS` 是 vLLM 服务端调度上限;`VLLM_MAX_INFLIGHT_REQUESTS` 和 `VLLM_PER_ENDPOINT_CONCURRENCY` 是 Producer 客户端请求上限。 - -## 4. 训练过程参数 - -| 参数 | 默认值 | 作用 | -| --- | --- | --- | -| `MAX_STEPS` | `10` | 本次运行最多完成的训练 step 数。Producer 会根据它和全局 batch size 计算需要生成多少样本。 | -| `BATCH_SIZE_PER_GPU` | `2` | 每个训练 rank 每个 step 使用的样本数。全局 batch size 为该值乘训练 rank 总数。 | -| `SAVE_INTERVAL_STEPS` | `5` | 每隔多少 optimizer step 保存一次 checkpoint;设为0表示不做周期保存。 | -| `SAVE_FINAL_CHECKPOINT` | `true` | 训练结束时是否保存最终 checkpoint。 | -| `LEARNING_RATE` | `1e-6` | Drafter optimizer 的基础学习率。 | -| `LR_WARMUP_STEPS` | `0` | 学习率 warmup 的 step 数。 | -| `LR_SCHEDULER_TYPE` | `constant` | 学习率调度类型,支持 `constant`、`cosine`、`linear`、`global_cosine`。 | -| `LR_DECAY_STEPS` | `100` | 需要衰减的 scheduler 使用的衰减 step 数。 | -| `MIN_LR_RATIO` | `0.1` | 学习率衰减后的最小值与基础学习率的比例。 | -| `PARAM_OFFLOAD` | `true` | FSDP 是否把模型参数 offload 到 CPU。 | -| `OPTIMIZER_OFFLOAD` | `true` | FSDP 是否把 optimizer state offload 到 CPU。 | - -Producer 需要发布的样本数为: - -```text -MAX_STEPS × BATCH_SIZE_PER_GPU × TRAIN_GPUS × 节点数 -``` - -如果文件样本不足,Producer 会重新从文件开头读取并再次请求 vLLM,达到所需样本数后才发布 EOS。 - -## 5. DSpark模型和采样参数 - -| 参数 | 默认值 | 作用 | -| --- | --- | --- | -| `DSPARK_BLOCK_SIZE` | `7` | 每个 anchor 并行预测的 token 数。 | -| `DSPARK_NUM_ANCHORS` | `32` | 每个样本选取的 anchor 数;越大训练计算量和显存占用越高。 | -| `DSPARK_MAX_WINDOW` | `512` | DSpark 从输入样本中取出的最大训练窗口长度。 | -| `DSPARK_NUM_TARGET_LAYERS` | `5` | 输入 DSpark 的 target 辅助 hidden-state 层数量。 | -| `DSPARK_NUM_HIDDEN_LAYERS` | `5` | DSpark drafter 自身 transformer 层数。 | -| `DSPARK_TARGET_LAYER_IDS` | `[1,9,17,25,33]` | Target model 中采集的辅助 hidden-state 层编号。必须与 vLLM 服务侧配置的辅助层一致。 | -| `DSPARK_MARKOV_RANK` | `256` | Markov head 的低秩维度。 | -| `DSPARK_MARKOV_HEAD_TYPE` | `vanilla` | Markov head 类型。当前独立训练路径建议使用 `vanilla`。 | - -## 6. DSpark损失参数 - -| 参数 | 默认值 | 作用 | -| --- | --- | --- | -| `DSPARK_LOSS_MODE` | `full_vocab` | CE 计算方式,可使用 `full_vocab`、`restricted_ce` 或 `sampled_ce`。 | -| `DSPARK_SAMPLED_CE_NEGATIVES` | `0` | `sampled_ce` 模式下采样的负类数量。 | -| `DSPARK_LOSS_DECAY_GAMMA` | `7` | block 内不同预测位置的指数衰减系数。 | -| `DSPARK_CE_LOSS_ALPHA` | `0.1` | Token CE loss 在总 loss 中的权重。 | -| `DSPARK_L1_LOSS_ALPHA` | `0.45` | Draft 与 target token 概率分布之间的 L1 loss 权重。Target 概率由 final hidden state 和 LM head 计算;设为0可关闭 L1 loss。 | -| `DSPARK_L1_CHUNK_SIZE` | `0` | L1 loss 分块计算大小;0表示不主动分块。显存不足时可设置正整数。 | -| `DSPARK_CONFIDENCE_LOSS_ALPHA` | `0.0` | Confidence loss 权重。当前 standalone 协议没有 acceptance target,必须保持0。 | - -当前 DSpark 总损失为: - -```text -loss = DSPARK_CE_LOSS_ALPHA × ce_loss - + DSPARK_L1_LOSS_ALPHA × l1_loss -``` - -只训练 CE 的配置: - -```bash -DSPARK_CE_LOSS_ALPHA=1.0 -DSPARK_L1_LOSS_ALPHA=0.0 -``` - -## 7. 调试参数 - -| 参数 | 默认值 | 作用 | -| --- | --- | --- | -| `DSPARK_DEBUG_LOG` | `false` | 是否输出 DSpark forward 的详细调试日志。 | -| `DSPARK_DEBUG_LOG_FIRST_N` | `2` | 开启调试日志后,前多少次 forward 必定打印。 | -| `DSPARK_DEBUG_LOG_INTERVAL` | `100` | 前几次之后,每隔多少次 forward 打印一次调试信息。 | - -## 8. Layer ID一致性 - -默认配置为: - -```text -训练侧 DSPARK_TARGET_LAYER_IDS = [1,9,17,25,33] -服务侧 VLLM_HIDDEN_STATE_LAYER_IDS = [1,9,17,25,33,36] -``` - -服务侧前五项是辅助层,必须与训练侧完全相同。最后的 `36` 是 Qwen3-4B 对应的 final hidden-state layer,供 DSpark L1 loss 使用。更换 target model 或辅助层配置时,两边需要一起修改。 diff --git a/docs/standalone_vllm_tq_dspark_training_plan.md b/docs/standalone_vllm_tq_dspark_training_plan.md deleted file mode 100644 index 7a865b8b..00000000 --- a/docs/standalone_vllm_tq_dspark_training_plan.md +++ /dev/null @@ -1,1130 +0,0 @@ -# 独立 vLLM Producer + TQ + DSpark Consumer 第一版方案 - -Last updated: 08/21/2026 - -## 1. 第一版要实现什么 - -只实现下面这条主链路: - -```text -verl prompt-only 数据或包含 prompt + response 的输入文件 -→ Producer 并发请求 vLLM prefill -→ Producer 将每条训练样本写入 TQ -→ Consumer 从同一个 TQ 取样本 -→ 独立 torchrun/FSDP DSpark 训练 -→ 一个 optimizer step 成功后删除该 step 的 TQ 样本 -``` - -第一版允许 **TQ 基础设施内部使用 Ray**,因为实测 `TransferQueue==0.1.7` 通过 Ray named actor 发现 `TransferQueueController`。Producer 和 Consumer 仍是普通 OS 进程,不改成 Ray actor;二者不使用 Ray RPC、Ray `ObjectRef` 或 Ray object store 传输训练样本。hidden-state payload 仍通过 TQ 的 `kv_put/kv_batch_get` 和配置的 MooncakeStore backend 传输。第一版不使用 DataProto/WorkerGroup,不生成长期 hidden-state feature store,也暂不实现复杂重试、自动重启和严格 checkpoint 数据恢复。 - -需要运行的组件: - -| 组件 | 数量 | 作用 | -|---|---:|---| -| vLLM server | 一个或多个 | 加载 target model,执行并行 prefill | -| Ray head | 1 个集群 | 保存 TQ named Controller/Storage actors;不承载 Producer/Consumer 业务 RPC | -| TQ owner | 1 个普通进程 | 连接 Ray,创建并保持 TQ controller/storage,最后统一关闭 TQ | -| Producer | 1 个进程 | 读文件、并发请求 vLLM、写 TQ | -| Consumer | 1 个 torchrun 任务 | 多个 rank 从 TQ 取数并训练 DSpark | - -Producer 和 Consumer 不互相调用,也不通过 HTTP 或 Ray 传样本。二者先连接同一个 Ray 集群,再通过携带相同 native 配置的 `tq.init(config)` 找到同一个 TQ,最后使用 TQ KV API 读写样本。已有 Controller 时 TransferQueue 0.1.7 会忽略后续配置并只连接;若 Client 意外先初始化,同一配置可避免默认 backend 抢先生效。 - -## 2. 共同的数据约定 - -这部分由两位开发者共同完成并先合入。建议文件: - -```text -verl_speco/transport/drafter_sample_protocol.py -tests/unit/test_drafter_sample_protocol.py -``` - -### 2.1 一个 key 对应一条样本 - -第一版固定: - -```text -一个输入文件 record -→ 一个 sequence_no -→ 一个 sample_id -→ 一个 TQ sample_key -→ 一个单样本 payload -``` - -`sequence_no` 是输入文件中的样本序号,不是 batch 编号。多个 sample keys 到 Consumer 后才组成训练 batch。 - -例如: - -```python -run_id = "dspark-20260818-a" -sequence_no = 17 -sample_id = "train-000017" - -partition_id = "speco_drafter_features" # 第一版沿用 PR 固定值 -sample_key = ( - "drafter:v1:dspark-20260818-a:" - "000000000017:train-000017" -) -``` - -### 2.2 Partition、key、tag 和 payload 的关系 - -TQ 中逻辑上是: - -```text -TQ 实例 -└── partition_id - └── sample_key - ├── tag - └── fields/payload -``` - -- TQ 实例:Producer 和 Consumer 共同连接的 controller/storage; -- `partition_id`:第一版沿用 PR 固定的 `speco_drafter_features`;每次启动新的 TQ owner,保证该 TQ 实例初始为空; -- `sample_key`:该分区中一条训练样本的地址; -- `tag`:轻量索引,Consumer 用 `kv_list()` 看见; -- `fields/payload`:真正的 Tensor 数据,Consumer 用 `kv_batch_get()` 读取。 - -Producer 写入: - -```python -tq.kv_put( - partition_id=partition_id, - key=sample_key, - fields=fields, - tag=tag, -) -``` - -Consumer 先发现 key: - -```python -all_records = tq.kv_list() -tags_by_key = all_records[partition_id] -``` - -这一步只拿 key 和 tag,不搬运 hidden states。 - -Consumer 再取数据: - -```python -result = tq.kv_batch_get( - partition_id=partition_id, - keys=selected_keys, -) -``` - -`kv_batch_get` 中的 batch 表示“一次读取多个独立 sample keys”,不是这些样本在 Producer 写入时就属于同一个对象。 - -### 2.3 Payload 字段 - -一个 sample key 对应的 `fields`: - -```python -fields = { - "input_ids": input_ids, # CPU int64[L] - "loss_mask": loss_mask, # CPU float32[L] - "position_ids": position_ids, # CPU int64[L] - "hidden_states": hidden_states, # CPU bf16[L,D] - "metadata_json": metadata_bytes, # CPU uint8[M] -} -``` - -| field | 含义 | Consumer 中的用途 | -|---|---|---| -| `input_ids` | 经过 feature window 选择后的 token IDs | DSpark token 输入 | -| `loss_mask` | 每个 token 是否参与 loss | 排除 prompt/padding/无效位置 | -| `position_ids` | 每个 row 对应的序列位置 | 位置编码和对齐校验 | -| `hidden_states` | target 指定层 hidden 沿最后一维拼接 | DSpark target/context 特征 | -| `metadata_json` | 模型、layers、layout、shape、样本身份 | Consumer 校验并恢复 metadata | - -符号: - -- `L`:这条训练 feature 保留的 token row 数; -- `H`:target model hidden size; -- `C`:DSpark context layer 数; -- L1 关闭:`D=C*H`,layout=`dflash_aux`; -- L1 开启:`D=C*H+H`,layout=`dflash_aux_plus_last`。 - -示例:`H=4096,C=5,L=1536`,开启 L1: - -```python -input_ids.shape == [1536] -loss_mask.shape == [1536] -position_ids.shape == [1536] -hidden_states.shape == [1536, 24576] -``` - -`metadata_json` 使用 JSON UTF-8 编码为 Tensor,因为参考 PR 的 bridge 只把 Tensor 放入 TQ fields: - -```python -raw = json.dumps(metadata, sort_keys=True).encode("utf-8") -metadata_bytes = torch.tensor(list(raw), dtype=torch.uint8) -``` - -### 2.4 Tag 字段 - -```python -tag = { - "record_type": "sample", - "status": "ready", - "schema_version": 1, - "run_id": "dspark-20260818-a", - "sequence_no": 17, - "sample_id": "train-000017", - "algorithm": "DSPARK", -} -``` - -tag 只存用于发现和筛选的 flat scalar/string。Consumer 只选择: - -```text -record_type=sample -status=ready -schema_version=1 -run_id=当前 run -algorithm=DSPARK -``` - -### 2.5 Metadata 字段 - -`metadata_json` 解码后至少包含: - -```python -metadata = { - "schema_version": 1, - "run_id": "dspark-20260818-a", - "sample_id": "train-000017", - "sequence_no": 17, - "algorithm": "DSPARK", - "target_model_id": "/models/Qwen3-8B", - "target_model_revision": "revision-or-checksum", - "tokenizer_fingerprint": "sha256:...", - "target_layer_ids": [2, 8, 14, 20, 26, -1], - "hidden_states_layout": "dflash_aux_plus_last", - "hidden_dtype": "bfloat16", - "hidden_shape": [1536, 24576], - "feature_length": 1536, - "full_sequence_length": 1800, - "feature_start": 264, - "feature_end": 1800, - "use_logits": False, -} -``` - -其中: - -- `target_model_id/revision`:hidden states 来自哪个 target checkpoint; -- `tokenizer_fingerprint`:Producer 使用的 tokenizer/template 版本; -- `target_layer_ids`:vLLM 返回和参与拼接的层; -- `hidden_states_layout`:Consumer 应如何拆 hidden 最后一维; -- `feature_length`:payload 中四个主要 Tensor 的第一维; -- `full_sequence_length`:完整 prompt+response 的 token 数; -- `[feature_start,feature_end)`:feature 在完整序列中的范围。 - -### 2.6 共享协议接口 - -Producer 和 Consumer 都 import 同一个模块,不各自手写字段名: - -```python -@dataclass(frozen=True) -class SampleMetadata: - schema_version: int - run_id: str - sample_id: str - sequence_no: int - algorithm: str - target_model_id: str - target_model_revision: str - tokenizer_fingerprint: str - target_layer_ids: list[int] - hidden_states_layout: str - hidden_dtype: str - hidden_shape: list[int] - feature_length: int - full_sequence_length: int - feature_start: int - feature_end: int - use_logits: bool - -def make_sample_key(meta: SampleMetadata) -> str: ... -def make_ready_tag(meta: SampleMetadata) -> dict: ... -def encode_sample(sample, meta: SampleMetadata) -> dict[str, Tensor]: ... -def decode_sample(key, tag, fields, expected_config) -> DraftFeatureSample: ... -def make_eos_record(run_id: str, total_samples: int): ... -``` - -Producer 使用 `SampleMetadata/make_sample_key/make_ready_tag/encode_sample`;Consumer 使用 `decode_sample`。 - -`SampleMetadata` Python 对象本身不经过 TQ: - -```text -Producer SampleMetadata -→ JSON -→ uint8 Tensor -→ TQ metadata_json -→ uint8 Tensor -→ JSON -→ Consumer metadata dict -``` - -`decode_sample()` 负责: - -1. 解码 `metadata_json`; -2. 校验 key、tag、metadata 中的 sample 身份一致; -3. 校验模型、tokenizer、layers 和 layout 与 Consumer 配置一致; -4. 校验 Tensor 必需字段、dtype 和 shape; -5. 返回现有 `DraftFeatureSample`。 - -### 2.7 EOS - -Producer 完成全部输入后写一个控制 record: - -```python -eos_key = f"control:v1:{run_id}:eos" -eos_fields = {"marker": torch.tensor([1], dtype=torch.uint8)} -eos_tag = { - "record_type": "control", - "status": "eos", - "schema_version": 1, - "run_id": run_id, - "total_samples": total_samples, -} -``` - -EOS 不进入训练 batch。Consumer 看到 EOS 后继续处理剩余 ready samples;`EOS 已出现且 ready 为空` 时结束。 - -## 3. TQ 怎么启动,Producer 和 Consumer 怎么连接同一个 TQ - -### 3.1 已验证的 TQ 0.1.7 连接机制 - -`TransferQueue==0.1.7` 没有提供“把 Controller 地址直接传给第二个进程”的高层连接接口。`tq.init(config)` 会先尝试以下已有 Controller 连接逻辑;存在时忽略传入配置,不存在时才用配置创建服务: - -```python -_TQ_CONTROLLER = ray.get_actor("TransferQueueController") -conf = ray.get(_TQ_CONTROLLER.get_config.remote()) -_maybe_create_tq_client(conf) -``` - -因此所有进程必须先加入同一个 Ray 集群。Ray 在本方案中只承担 TQ 控制面:保存 named `TransferQueueController`、返回 TQ 配置、管理 TQ 创建的 actor。业务进程不通过 Ray 发送 sample key 或 hidden-state tensor。 - -实际连接链路是: - -```text -TQ owner:ray.init(address) → tq.init(full_tq_config) → 创建 named Controller -Producer:ray.init(address) → tq.init(same config) → ray.get_actor() → 创建本地 TQ client -Consumer rank 0..N:ray.init(address) → tq.init(same config) → ray.get_actor() → 创建各自 TQ client -``` - -### 3.2 直接移植并扩展 PR #48 的 bridge - -参考文件: - -```text -C:/Users/xxyyrr/Desktop/上班/verl-SpeCo/ - verl_speco/integration/transferqueue_bridge.py -``` - -第一版不新增 `standalone_tq.py`,直接把该文件移植到目标项目并扩展。保留 PR 已有的 import 检查、进程内幂等状态、`kv_put`、`kv_batch_get`、TensorDict 解包和 owner-only shutdown;新增 Ray 显式连接、批量 list/get/clear 和本地 client 关闭接口。 - -目标文件必须实现以下函数,而不只是提供一个笼统的 transport class: - -```python -def configure_transfer_queue(config: Mapping[str, Any]) -> bool: ... -def connect_ray_cluster(ray_address: str, namespace: str | None = None) -> None: ... -def start_transfer_queue_owner(tq_config: Mapping[str, Any]) -> None: ... -def connect_transfer_queue_client() -> None: ... -def make_sample_key(run_id: str, sequence_no: int, sample_id: str) -> str: ... -def put_sample(key: str, fields: dict[str, Tensor], tag: dict[str, Any]) -> None: ... -def list_samples() -> dict[str, dict[str, Any]]: ... -def get_samples(keys: list[str]) -> list[tuple[str, dict[str, Tensor]]]: ... -def clear_samples(keys: list[str]) -> None: ... -def close_transfer_queue_client() -> None: ... -def close_transfer_queue_owner() -> None: ... -``` - -逐个函数的责任如下。 - -#### `configure_transfer_queue(config)` - -- 从 Hydra/OmegaConf 中读取 Ray address、可选 namespace、固定 partition、run ID、schema version 和 TQ native backend 配置; -- 转成普通 Python dict,保存在进程内 `_state`; -- 校验 `TransferQueue==0.1.7` 可 import; -- 不连接 Ray,不创建 TQ,不产生跨进程副作用; -- 返回该进程是否启用了 TQ。 - -#### `connect_ray_cluster(ray_address, namespace)` - -- 若 `ray.is_initialized()` 为 false,调用 `ray.init(address=ray_address, namespace=namespace)`; -- 若已经初始化,校验现有 Ray context 指向期望集群,不能悄悄连到另一个本地 Ray; -- Owner、Producer 和所有 torchrun ranks 都调用它; -- 此函数只建立当前 OS 进程到 Ray control plane 的连接。 - -#### `start_transfer_queue_owner(tq_config)` - -- 仅由 `tq_owner.py` 调用; -- 前置条件是 `connect_ray_cluster()` 已成功; -- 调用一次 `tq.init(OmegaConf.create(tq_config))`; -- 将 `_state.owner=True`、`_state.initialized=True`; -- TQ 0.1.7 会创建名为 `TransferQueueController` 的 Ray actor,并创建所选 storage backend; -- 重复调用必须报错,不能启动第二套同名 Controller。 - -#### `connect_transfer_queue_client()` - -- 由 Producer 和每个 Consumer rank 调用; -- 前置条件是当前进程已经连接 Ray; -- 调用 `tq.init(same native config)`,通过 `ray.get_actor("TransferQueueController")` 发现 owner;已有 Controller 时配置会被忽略,意外抢先时则以相同配置创建; -- 只创建当前进程的 TQ client,不创建新的 Controller; -- 成功后设置 `_state.initialized=True`;重复调用直接返回。 - -#### `put_sample/list_samples/get_samples/clear_samples` - -- 全部固定使用 `_SPECO_TQ_PARTITION = "speco_drafter_features"`; -- `put_sample()` 调用单样本 `tq.kv_put()`; -- `list_samples()` 调用 `tq.kv_list(partition_id=...)`,只返回 key/tag 元数据; -- `get_samples()` 一次调用 `tq.kv_batch_get(keys=...)`,再按输入 key 顺序解包为普通 dict; -- `clear_samples()` 调用 `tq.kv_clear(keys=..., partition_id=...)`; -- 这些函数不调用 Ray RPC,不把 tensor 放入 Ray object store。 - -#### `close_transfer_queue_client()` 与 `close_transfer_queue_owner()` - -TQ 0.1.7 的公共 `tq.close()` 会 kill 共享 Controller,所以两者不能写成同一个实现: - -- `close_transfer_queue_client()`:通过 0.1.7 已公开的 `tq.get_client()` 取得当前进程 client并调用它的 `close()`,然后 `ray.shutdown()`;绝不能调用会 kill Controller 的全局 `tq.close()`。这个 client close 只用于进程退出阶段,关闭后本进程不能再次调用 TQ; -- `close_transfer_queue_owner()`:仅当 `_state.owner=True` 时调用 `tq.close()`,清理 Controller/Storage,最后 `ray.shutdown()`; -- Producer 或任一训练 rank 提前退出都不能关闭全局 TQ。 - -### 3.3 共享配置 - -Owner、Producer 和 Consumer 必须使用相同的 Ray address/namespace,并通过同一个 named Controller 获得 backend 配置。第一版 partition 固定,不按任务动态创建: - -```yaml -transfer_queue: - enable: true - package_version: "0.1.7" - ray: - address: "ray-head-node:6379" - namespace: "speco-drafter" - partition_id: "speco_drafter_features" - run_id: "dspark-20260819-a" - schema_version: 1 - backend: - storage_backend: MooncakeStore - MooncakeStore: - auto_init: false - metadata_server: "node0:50050" - master_server_address: "node0:50051" - local_hostname: "" - protocol: tcp - global_segment_size: 4294967296 - local_buffer_size: 1073741824 - device_name: "" -``` - -`partition_id/run_id/schema_version` 用于过滤和校验样本;它们不能帮助进程发现 TQ。真正让三个任务连接到同一 TQ 的是“连接同一个 Ray 集群和 namespace,然后找到同名 Controller”。 - -依赖也必须锁定并单独验证。实机安装 `TransferQueue==0.1.7` 会安装 Ray,并要求 `numpy<2.0.0`;当前测试环境中的 `twinkle-kit` 要求 `numpy>=2.0.0`,两者冲突。开发时应使用专门的 TQ/训练环境或重新确认整套依赖约束,不能直接把 0.1.7 安装进已有生产环境后忽略 resolver warning。 - -### 3.4 `tq_owner.py` 要实现的入口和函数 - -新增: - -```text -verl_speco/tq_owner.py -``` - -`tq_owner.py` 建议明确实现: - -```python -def install_signal_handlers(stop_event: threading.Event) -> None: ... -def publish_owner_ready(run_id: str, schema_version: int) -> None: ... -def wait_until_stopped(stop_event: threading.Event) -> None: ... -def run_owner(config: DictConfig) -> int: ... -def main() -> None: ... -``` - -`run_owner()` 的执行顺序必须是: - -```text -configure_transfer_queue(config) -→ connect_ray_cluster(ray.address, ray.namespace) -→ start_transfer_queue_owner(full TQ native config) -→ put owner_ready 控制 record -→ 安装 SIGINT/SIGTERM handler -→ 保持 owner 进程存活 -→ 收到停止信号 -→ close_transfer_queue_owner() -``` - -Owner 必须常驻。TQ 0.1.7 创建 Controller 时没有设置 `lifetime="detached"`,不能在初始化后立即退出。 - -### 3.5 启动和关闭顺序 - -第一版由外部脚本管理全生命周期: - -```text -1. ray start --head,记录 Ray address -2. 启动 Mooncake metadata/master(若 auto_init=false) -3. 启动 TQ owner;owner 连接 Ray并调用 tq.init(full config) -4. 等待 owner_ready -5. 启动一个或多个 vLLM servers -6. 启动 Consumer;每个 torchrun rank 连接 Ray,然后 tq.init(same native config) -7. 启动 Producer;连接 Ray,然后 tq.init(same native config) -8. Producer 写 EOS,关闭本地 client并退出 -9. Consumer drain、保存 final checkpoint,所有 ranks 关闭本地 client并退出 -10. 给 TQ owner 发送 SIGTERM;仅 owner 执行 tq.close() -11. 等 owner 退出后执行 ray stop -12. 停止 Mooncake 服务 -``` - -外部脚本要用 `trap` 保证异常退出也按“Producer/Consumer → owner → Ray → Mooncake”的顺序清理。不能在 Consumer rank 的 `finally` 中调用全局 `tq.close()`。 - -## 4. Producer 要实现什么 - -### 4.1 Producer 完整顺序 - -```text -读取共享配置 -→ 连接 TQ并校验 owner_ready -→ 初始化 tokenizer -→ 初始化多个 vLLM endpoint clients -→ 流式读取输入文件 -→ 为每条输入分配 sequence_no/sample_id -→ 缺少 response 时由 target vLLM 生成;构造 input_ids/loss_mask -→ 并发请求 vLLM prefill -→ 读取 vLLM hidden-state 临时结果 -→ 转换成 DSpark DraftFeatureSample -→ 构造 SampleMetadata -→ encode_sample 得到 fields/tag/key -→ TQ kv_put 一条 sample -→ 删除该请求临时文件 -→ 所有输入完成后写 EOS -→ close_transfer_queue_client()并退出 -``` - -### 4.2 并发模型 - -Producer 是一个进程,内部并发请求多个 endpoint: - -```text -InputReader -→ bounded asyncio input_queue -→ N 个 RequestWorker -→ bounded publish_queue -→ TQ Publisher -``` - -- `vllm_endpoints` 是列表; -- 每个 endpoint 有独立 semaphore; -- 总并发由 `max_inflight_requests` 限制; -- input/publish queue 必须有上限; -- TQ `kv_put` 如果是同步 API,用 `asyncio.to_thread()` 调用; -- TQ ready 数达到 `max_pending_samples` 时暂停继续请求,形成简单背压。 - -`sequence_no` 在 InputReader 中分配,不按 vLLM 完成顺序分配。因此并发乱序不会改变 sample key。 - -### 4.3 vLLM 结果转换 - -复用现有 `TargetFeatureReplayer` 的: - -- OpenAI-compatible vLLM 请求; -- `prompt_token_ids` 校验; -- `kv_transfer_params.hidden_states_path`; -- safetensors 加载; -- `[seq,layers,hidden]` 校验; -- feature positions 选择; -- aux layers flatten; -- DSpark L1 时拼 final hidden。 - -不要复制 `_feature_from_vllm_payload()`;把纯转换逻辑提取成公共函数,让旧 file backend 和新 Producer 共同调用。 - -临时文件顺序: - -```text -加载 -→ 校验/转换 -→ TQ put 成功 -→ 删除 -``` - -第一版仍允许每个并发请求产生短期临时文件,但不生成全量 hidden-state 数据集。 - -### 4.4 Producer 文件分工 - -| 文件 | 实现内容 | -|---|---| -| `verl_speco/standalone_tq_producer.py` | CLI/Hydra 入口,连接 TQ,启动 asyncio pipeline,写 EOS和汇总指标 | -| `verl_speco/producer/input_reader.py` | 流式读输入、分配 `sequence_no/sample_id`、tokenize、构造 loss mask | -| `verl_speco/producer/vllm_feature_client.py` | endpoint pool、并发控制、HTTP 请求、临时结果读取和删除 | -| `verl_speco/trainer/target_feature_replay.py` | 提取可复用的 vLLM payload → `DraftFeatureSample` 纯转换函数 | -| `tests/unit/test_drafter_sample_protocol.py` | 协议编码、字段和 shape 测试(两人共同) | -| `tests/integration/test_tq_producer_smoke.py` | fake/短 vLLM → TQ sample + EOS | - -Producer 开发者同时负责移植/扩展 `integration/transferqueue_bridge.py` 和实现 `tq_owner.py`,因为这部分与 TQ 写入和连接直接相关。 - -### 4.5 Producer 各文件的函数级实现规格 - -#### `verl_speco/standalone_tq_producer.py` - -需要实现: - -```python -@dataclass -class ProducerStats: - input_count: int - published_count: int - failed_count: int - pending_bytes: int - -async def publish_one(result: PreparedFeature, transport) -> str: ... -async def run_producer(config: DictConfig) -> ProducerStats: ... -def validate_producer_config(config: DictConfig) -> None: ... -def main() -> None: ... -``` - -`main()` 只负责 Hydra/日志/退出码。`run_producer()` 是可测试的业务入口,执行:连接 Ray → 连接 TQ client → 创建 InputReader 和 vLLM client pool → 启动有界 asyncio pipeline → 等待所有请求及 `kv_put` 完成 → 发布 EOS → 关闭本地 client。它不能启动 Ray head、不能创建 TQ owner、不能调用全局 `tq.close()`。 - -`publish_one()` 接收已经完成转换的一条 `PreparedFeature`,调用共享协议的 `encode_sample()` 得到 `(key, fields, tag)`,再通过 bridge 的 `put_sample()` 发布。只有 `put_sample()` 成功返回,`published_count` 才增加,vLLM 临时文件才允许删除。 - -#### `verl_speco/producer/input_reader.py` - -需要实现: - -```python -@dataclass(frozen=True) -class InputRecord: - sequence_no: int - sample_id: str - prompt: str - response: str | None - source_metadata: dict[str, Any] - -def iter_input_records(path: str) -> Iterator[InputRecord]: ... -def tokenize_record(record: InputRecord, tokenizer, config) -> TokenizedRequest: ... -def build_loss_mask(input_ids: Tensor, prompt_length: int) -> Tensor: ... -``` - -`iter_input_records()` 流式读取 JSONL/Parquet,不把全文件载入内存,并按文件顺序分配稳定的 `sequence_no`。已有 response 时 `tokenize_record()` 直接拼接;prompt-only verl 数据通过 chat template 编码后由 target vLLM 生成 response,并设置 `include_output_tokens=true` 同步提取输出 hidden states。 - -#### `verl_speco/producer/vllm_feature_client.py` - -需要实现: - -```python -@dataclass(frozen=True) -class VllmEndpoint: - base_url: str - max_concurrency: int - -class VllmFeatureClientPool: - async def start(self) -> None: ... - async def prefill(self, request: TokenizedRequest) -> RawVllmFeature: ... - async def close(self) -> None: ... - -async def request_prefill(endpoint, request) -> VllmResponse: ... -def choose_endpoint(endpoints, state) -> VllmEndpoint: ... -def load_hidden_state_result(response) -> RawVllmFeature: ... -def delete_temporary_result(raw: RawVllmFeature) -> None: ... -``` - -`prefill()` 必须允许多个 coroutine 同时运行;全局 semaphore 限制总并发,每个 endpoint 另有独立 semaphore。`request_prefill()` 只负责 HTTP 请求和响应校验;`load_hidden_state_result()` 负责读取 vLLM 返回的临时 safetensors/path。临时结果的删除不放在 `load_hidden_state_result()`,而由 `publish_one()` 成功后触发。 - -#### `verl_speco/trainer/target_feature_replay.py` - -把当前类内部的纯转换部分抽成: - -```python -def feature_from_vllm_payload( - payload: RawVllmFeature, - request: TokenizedRequest, - feature_config: FeatureContract, -) -> DraftFeatureSample: ... -``` - -它不进行 HTTP、不访问 TQ、不删除文件,只完成 shape/layout/layer 校验、feature-position 选择、aux layer flatten 和可选 L1 final hidden 拼接。现有 file replay backend 和新 Producer 都调用这一函数,避免两套转换规则。 - -#### Producer 启动配置 - -统一 launcher 向 `verl_speco.standalone_tq_producer` 提供同一套: - -```text -RAY_ADDRESS / Ray namespace -run_id / schema_version / 固定 partition -Mooncake/TQ backend 配置 -输入文件和 tokenizer/model 配置 -vLLM endpoint 列表 -max_inflight_requests / per_endpoint_concurrency -``` - -正式运行不再保留单独的角色 shell wrapper。 - -## 5. Consumer 要实现什么 - -### 5.1 不新写另一套训练器 - -继续使用现有入口: - -```text -draft_train_launcher.py -→ draft_train.py -→ trainer/draft_training_loop.py -→ DrafterBaseTrainer -→ DSparkTrainerBackend -``` - -训练模式继续使用现有 standalone `training.mode=offline`,用 `feature_store.type=tq` 选择 TQ 数据源。不复制 FSDP、loss、optimizer 或 checkpoint 代码。 - -当前 `feature_store.type` 的作用是告诉 `build_feature_store_from_config()` 应创建哪一种训练数据来源: - -| 当前 type | 对象 | 数据来源 | -|---|---|---| -| `torch_shard` | `TorchShardFeatureStore` | 本地 `.pt` shard,当前 separate-training 示例使用它 | -| `token_replay` | `TokenReplayFeatureStore` | 保存 token replay 的 shard | -| `vllm_safetensors` / `safetensors` | `VllmSafetensorsFeatureStore` | 已生成的 safetensors feature | -| `jsonl_token_replay` / `jsonl` | `JsonlTokenReplayFeatureStore` | JSONL token/text replay | - -第一版新增: - -```yaml -actor_rollout_ref: - rollout: - drafter: - training: - mode: offline - feature_store: - type: tq - path: null - shuffle: false - repeat: false -``` - -这里 `offline` 的含义是“独立于 RL 的 standalone drafter training”,不等于数据必须来自磁盘。 - -不能只给现有工厂增加一个 `type=tq` 然后继续使用通用 `DraftFeatureDataLoader`。当前 loader 会: - -```python -keys = list(store.iter_keys(...)) -``` - -它假设数据集 keys 是一个静态快照;空 store 会直接结束,并且各 rank 独立枚举时可能在 Producer 持续写入的过程中看到不同快照。TQ 是流式、会新增并删除 keys 的数据源,所以 `type=tq` 必须选择专用的 `TQFeatureDataLoader`,由 rank 0 统一发现和分配 keys。 - -#### 5.1.1 当前磁盘 feature store 为什么每个 rank 都会取 keys - -`draft_train_launcher.py` 通过 `torchrun` 启动多个训练进程。每个进程都是一个 rank,并且每个 rank 都会独立进入 `run_standalone_draft_training()`、创建 `DraftFeatureDataLoader`、执行它的 `__iter__()`。当前 loader 的核心逻辑是: - -```python -keys = list( - self.store.iter_keys( - shuffle=self.shuffle, - seed=self.seed + epoch, - ) -) -rank_keys = keys[rank::world_size] - -for key in rank_keys: - batch.append(self.store.read(key)) -``` - -因此当前并不是 rank 0 枚举 keys 后再发送给其他 rank,而是所有 rank 都访问相同的静态 feature-store 路径: - -```text -rank 0:iter_keys() 得到完整静态列表 → keys[0::world_size] → read 自己的 samples -rank 1:iter_keys() 得到完整静态列表 → keys[1::world_size] → read 自己的 samples -... -``` - -例如 store 中固定存在: - -```python -keys = ["k0", "k1", "k2", "k3", "k4", "k5", "k6", "k7"] -``` - -当 `world_size=2` 时,两个 rank 都先得到上述完整列表,然后分别计算: - -```python -# rank 0 -rank_keys = keys[0::2] # ["k0", "k2", "k4", "k6"] - -# rank 1 -rank_keys = keys[1::2] # ["k1", "k3", "k5", "k7"] -``` - -这种实现能够成立,是因为 `.pt`、safetensors 或 replay 文件在训练期间是静态数据集。只要所有 rank 使用同一路径和相同 shuffle seed,`iter_keys()` 就会产生相同顺序的 key 快照,各 rank 可以无通信地算出互不重叠的子集。 - -#### 5.1.2 TQ 为什么必须改成 rank 0 发现 keys - -TQ 中的 keys 会在训练期间动态变化:Producer 持续 `kv_put`,训练成功后 Consumer 又执行 `kv_clear`。如果所有 rank 仍然各自调用 `kv_list`,不同调用时刻可能看到不同快照: - -```python -# rank 0 较早调用 -rank0_keys = ["k0", "k1", "k2", "k3"] - -# Producer 随后写入 k4、k5,rank 1 较晚调用 -rank1_keys = ["k0", "k1", "k2", "k3", "k4", "k5"] -``` - -各 rank 再独立切片后,可能得到不同数量的 local samples,进而无法保证它们以相同顺序进入 forward、backward 和梯度 collective。为此,TQ 专用 loader 必须把控制面和数据面分开: - -```text -控制面:rank 0 执行 kv_list,固定本 step 的 global_keys,并向各 rank 分发 local_keys -数据面:每个 rank 使用自己的 local_keys 直接执行 kv_batch_get,从 TQ/Mooncake 读取 hidden-state payload -``` - -rank 0 只发送较小的 key/tag 元数据,不读取并转发其他 rank 的 hidden-state tensor。所有 rank 仍然都会连接同一个 TQ,也都会调用批量读取接口;只有动态 key 的发现和本 step 的 global batch 决策集中在 rank 0。 - -因此 `feature_store.type=tq` 不是只替换底层 `read()`:它同时改变了 key 的发现、分配、等待和删除语义,需要专用 `TQFeatureDataLoader`,并要求训练循环在所有 rank 成功完成 optimizer step 后,由 rank 0 对本 step 的 `global_keys` 执行 `kv_clear`。 - -### 5.2 Consumer 完整顺序 - -```text -torchrun 启动多个 ranks -→ 每个 rank 初始化 torch.distributed -→ 每个 rank 连接同一个 TQ -→ rank 0 校验 owner_ready,并 broadcast 结果 -→ 初始化现有 DSpark trainer -→ rank 0 kv_list 查找 ready sample keys -→ rank 0 选一个 global batch并分给各 rank -→ 每个 rank kv_batch_get 自己的 local keys -→ decode_sample 得到 list[DraftFeatureSample] -→ prepare_training_batch_from_samples() -→ training_step_from_batch() -→ 所有 rank 汇总 success -→ 成功后 rank 0 kv_clear 这个 global batch 的 keys -→ 继续下一批 -→ 看到 EOS 且 ready 为空 -→ 保存 final checkpoint -→ 所有 ranks close_transfer_queue_client()并退出 -``` - -### 5.3 多 rank 如何分 key - -例如: - -```text -world_size=2 -batch_size_per_gpu=2 -global batch size=4 -``` - -rank 0 选出: - -```python -global_keys = ["k10", "k11", "k12", "k13"] -assignments = [ - ["k10", "k11"], # rank 0 - ["k12", "k13"], # rank 1 -] -``` - -通过 `dist.scatter_object_list` 或 broadcast 发送短字符串列表。每个 rank 直接从 TQ 读取自己的 payload,不通过 rank 0 转发 hidden states。 - -### 5.4 从 TQ 到训练 batch - -每个 rank: - -```python -records = tq_transport.get_samples(local_keys) - -samples = [ - decode_sample( - key=key, - tag=tags_by_key[key], - fields=fields, - expected_config=expected_contract, - ) - for key, fields in records -] - -batch = trainer.prepare_training_batch_from_samples( - samples, - step=optimizer_step, -) - -ok = await trainer.training_step_from_batch( - batch, - optimizer_step, -) -``` - -`records` 是多个独立单样本 payload;`samples` 是变长 `DraftFeatureSample` 列表。Consumer transport 层不要直接 `torch.stack()`,现有 trainer/backend 负责对齐和组 batch。 - -### 5.5 删除与结束 - -第一版采用简单逻辑: - -```text -所有 rank get/decode/train 都成功 -→ all_reduce(global_success)=True -→ rank 0 kv_clear(global_batch_keys) -``` - -任何 rank 失败都不 clear。程序报错退出,第一版不自动恢复。 - -EOS 后不足一个 global batch 的尾部,第一版使用 `drop_last=True`:记录数量,rank 0 clear 这些尾部 keys,然后正常结束。 - -checkpoint 仍使用现有 `save_interval_steps` 和 final checkpoint;第一版不保证崩溃后已 clear 数据能够严格重放。 - -### 5.6 Consumer 文件分工 - -| 文件 | 实现内容 | -|---|---| -| `verl_speco/trainer/feature_store.py` | 工厂增加 `type=tq`,返回 `TQFeatureStore`;TQ 类型不要求 `path` | -| `verl_speco/trainer/tq_feature_store.py` | 实现固定 partition 上的 list/get-many/clear/EOS,不伪装静态 `iter_keys()` | -| `verl_speco/trainer/tq_sample_source.py` | 实现专用 `TQFeatureDataLoader`:rank 0 动态 list/filter/sort、key 分配、各 rank get/decode、EOS/tail | -| `verl_speco/trainer/draft_training_loop.py` | 保持 `mode=offline`;`type=tq` 时跳过 `feature_store.path` 必填检查并选择专用 loader;调用现有 prepare/train,global success 后 clear,final checkpoint | -| `verl_speco/draft_train_launcher.py` | `feature_store.type=tq` 时不要求 feature-store path,透传 torchrun/TQ 配置 | -| `verl_speco/config/speco_base.yaml` | 新增 TQ connection、run、poll、batch 等默认配置 | -| `tests/unit/test_tq_sample_source.py` | key 过滤/排序/分 rank、EOS、decode 调用测试 | -| `tests/integration/test_dspark_tq_consumer_smoke.py` | 两 rank 取不同样本,完成训练并 clear | - -### 5.7 Consumer 各文件的函数级实现规格 - -#### `verl_speco/trainer/feature_store.py` - -修改现有工厂: - -```python -def build_feature_store_from_config(feature_store_cfg, read_only=False): - store_type = str(feature_store_cfg.get("type", "torch_shard")).lower() - if store_type == "tq": - return TQFeatureStore.from_config(feature_store_cfg) - ... -``` - -要求: - -- 保留 `torch_shard/token_replay/vllm_safetensors/jsonl_token_replay` 的当前行为; -- `type=tq` 时不读取 `path`; -- 工厂只创建对象,不在 import 阶段连接 Ray/TQ; -- `TQFeatureStore` 是流式数据源适配器,不强行实现有误导性的静态 `iter_keys()` 和逐条 `read()`。 - -#### `verl_speco/trainer/tq_feature_store.py` - -需要实现: - -```python -@dataclass(frozen=True) -class ReadyEntry: - key: str - tag: dict[str, Any] - -class TQFeatureStore: - @classmethod - def from_config(cls, cfg) -> "TQFeatureStore": ... - def connect(self) -> None: ... - def list_ready(self, run_id: str) -> list[ReadyEntry]: ... - def get_many(self, entries: list[ReadyEntry]) -> list[DraftFeatureSample]: ... - def clear_many(self, keys: list[str]) -> None: ... - def read_eos(self, run_id: str) -> EosMetadata | None: ... - def close_local(self) -> None: ... -``` - -`connect()` 调用共享 bridge 的 `connect_ray_cluster()` 和 `connect_transfer_queue_client()`。`list_ready()` 调用 `list_samples()` 后只保留 `tag.status=ready`、匹配 `run_id/schema_version` 的数据 key,并按 `(sequence_no,key)` 排序。`get_many()` 一次批量读取 fields,然后逐条调用共享 `decode_sample()`;返回顺序必须和 entries 相同。`clear_many()` 只允许 rank 0 在全局 step 成功后调用。`close_local()` 只关闭本 rank client,不能关闭 Controller。 - -#### `verl_speco/trainer/tq_sample_source.py` - -需要实现: - -```python -@dataclass -class TQLocalBatch: - local_keys: list[str] - local_samples: list[DraftFeatureSample] - global_keys: list[str] | None - -class TQFeatureDataLoader: - def __iter__(self) -> Iterator[TQLocalBatch]: ... - def _select_global_batch(self) -> tuple[list[str], dict[str, ReadyEntry]]: ... - def _broadcast_assignments(self, assignments) -> list[ReadyEntry]: ... - def _handle_eos_and_tail(self) -> bool: ... - def clear_completed_batch(self, global_keys: list[str] | None) -> None: ... -``` - -执行责任必须明确: - -- 所有 rank 创建 loader 并调用 `store.connect()`; -- 只有 rank 0 执行 `_select_global_batch()` 和 `kv_list`; -- rank 0 生成 `list[list[ReadyEntry]]` assignments,使用 `torch.distributed.scatter_object_list` 或 broadcast 分发小型 key/tag; -- 每个 rank 对自己的 local entries 调用 `store.get_many()`,hidden states 从 TQ/Mooncake 直接进入该 rank; -- loader yield `TQLocalBatch`,不能只 yield samples,因为训练后清理还需要 `global_keys`; -- 暂时为空时轮询等待,不能像静态 loader 一样结束;只有 EOS 已出现并且 ready 数据 drain 完才停止; -- `drop_last=True` 时由 rank 0 记录并清理不足 global batch 的尾部 key。 - -#### `verl_speco/trainer/draft_training_loop.py` - -需要新增或调整: - -```python -def build_training_source(config, rank, world_size): ... -def all_ranks_succeeded(local_ok: bool, device) -> bool: ... -async def run_tq_training_loop(trainer, loader, config) -> dict[str, Any]: ... -``` - -`build_training_source()` 根据 `feature_store.type` 分支:静态类型继续创建 `DraftFeatureDataLoader`;`tq` 创建 `TQFeatureDataLoader`,跳过 `feature_store.path` 必填检查,并禁止 `shuffle/repeat`。`run_tq_training_loop()` 对每个 `TQLocalBatch` 调用现有 `prepare_training_batch_from_samples()` 和 `training_step_from_batch()`;所有 rank 通过 collective 得到 global success 后,才让 rank 0 调用 `clear_completed_batch(global_keys)`。异常路径不 clear,finally 只执行 `store.close_local()`。 - -#### `verl_speco/draft_train_launcher.py` - -保留当前 `torch.distributed.run` 启动方式,只增加配置校验和环境透传: - -```python -def validate_tq_launch_config(overrides, launch_config) -> None: ... -def build_child_env(config) -> dict[str, str]: ... -``` - -它不启动 Ray head、不调用 `tq.init()`。职责是确认 `type=tq` 时提供了 Ray address/namespace,且不要求 feature-store path;随后把同一 Ray connection 配置传给每个 torchrun 子进程。每个子 rank 自己建立 Ray/TQ client,不能在 launcher 父进程建一个 client后期待 fork 继承。 - -#### `verl_speco/config/speco_base.yaml` - -增加默认字段: - -```yaml -feature_store: - type: torch_shard - path: null - shuffle: true - repeat: true - tq: - ray_address: null - ray_namespace: speco-drafter - partition_id: speco_drafter_features - run_id: null - schema_version: 1 - poll_interval_seconds: 0.5 - connect_timeout_seconds: 120 - drop_last: true -``` - -当 `type=tq` 时运行期覆盖为 `shuffle=false/repeat=false`。backend 的完整 owner 配置只需要 TQ owner 使用;Producer/Consumer client 从 named Controller 获取它,不各自重新创建 backend。 - -#### Consumer 测试必须覆盖的函数边界 - -- `test_feature_store_factory_builds_tq_without_path()`; -- `test_rank0_filters_and_sorts_ready_entries()`; -- `test_nonzero_rank_never_calls_kv_list()`; -- `test_assignments_are_disjoint_and_global_batch_complete()`; -- `test_each_rank_gets_only_local_keys()`; -- `test_decode_preserves_hidden_states_layout()`; -- `test_clear_only_after_all_ranks_success()`; -- `test_failure_does_not_clear()`; -- `test_eos_drains_ready_then_stops()`; -- `test_client_close_does_not_kill_owner()`。 - -## 6. 两个人怎么分工 - -### 共同先完成 - -1. `drafter_sample_protocol.py`; -2. Ray/TQ connection 配置字段; -3. 一个小型 golden sample; -4. 启动一个测试 Ray head 和 TQ owner,独立进程 A put、独立进程 B list/get/clear 的 smoke test; -5. 验证 client 退出不会 kill owner,只有 owner shutdown 才销毁 Controller。 - -### 开发者 A:Producer/TQ - -负责: - -```text -integration/transferqueue_bridge.py -tq_owner.py -standalone_tq_producer.py -producer/input_reader.py -producer/vllm_feature_client.py -target_feature_replay.py 的公共转换函数 -owner/producer 启动脚本 -Producer/TQ 测试 -``` - -开发者 A 的可交付接口不是“提供一个 TQ 类”,而是: - -```text -bridge:connect_ray_cluster/start_owner/connect_client/put/list/get_many/clear/client_close/owner_close -owner:run_owner/main/signal handler/owner_ready -producer:run_producer/publish_one/统计与 EOS -input reader:iter_input_records/tokenize_record/build_loss_mask -vLLM client:endpoint pool/request_prefill/load/delete -feature conversion:feature_from_vllm_payload -``` - -开发者 B 可以先针对这些接口写 fake transport,不需要等待真实 vLLM 和 Mooncake 联通。 - -### 开发者 B:Consumer/训练 - -负责: - -```text -feature_store.py 的 type=tq 工厂分支 -tq_feature_store.py -tq_sample_source.py / TQFeatureDataLoader -draft_training_loop.py 的 offline + type=tq 分支 -draft_train_launcher.py 配置适配 -speco_base.yaml Consumer 配置 -Consumer 启动脚本 -Consumer/DSpark 测试 -``` - -开发者 B 的可交付接口是: - -```text -feature-store factory:type=tq 分支 -TQFeatureStore:connect/list_ready/get_many/clear_many/read_eos/close_local -TQFeatureDataLoader:select global batch/distribute local keys/get/yield/EOS-tail -training loop:build source/train/global success/clear/final checkpoint -launcher:TQ 配置校验和 torchrun 子进程环境透传 -``` - -### 联调入口 - -建议再提供: - -```text -examples/run_dspark_tq_pipeline_local.sh -``` - -只用于单机联调,顺序启动: - -```text -ray start --head -→ Mooncake metadata/master -→ TQ owner(ray.init + tq.init(full config)) -→ owner_ready -→ vLLM health check -→ Consumer -→ Producer -→ 等 Producer/Consumer 退出 -→ SIGTERM TQ owner(owner 执行 tq.close) -→ ray stop -→ 停止 Mooncake -``` - -最小联调:Producer 发布 8 条样本,2 个 Consumer ranks、每 rank batch size 2,完成 2 个 optimizer steps,8 个 sample keys 被清理,EOS 后保存 final checkpoint。 - -## 7. 第一版验收标准 - -1. TQ owner、Producer 和所有 Consumer ranks 日志显示相同 Ray address/namespace、TQ Controller、partition 和 run ID。 -2. 在同一 Ray 集群中,独立 owner 创建 TQ 后,独立进程 A put,独立进程 B 能 list/get/clear。 -3. Producer 对多个 vLLM endpoints 并发请求,不串行访问。 -4. 一个输入 record 只生成一个 sample key 和一个 payload。 -5. `kv_list` 只拿 tag;`kv_batch_get` 才拿 hidden states。 -6. Consumer 各 rank 读取不同 local keys,不通过 rank 0 搬运 Tensor。 -7. `decode_sample` 能拒绝模型、layer、layout、dtype 或 shape 不匹配的数据。 -8. 所有 rank 训练成功后才 clear 当前 global batch。 -9. Producer 先完成时,Consumer 能 drain 后再退出。 -10. 不产生长期 hidden-state feature store。 -11. Producer 或任一 Consumer rank 退出不会销毁 TQ Controller;只有 owner 调用全局 `tq.close()`。 -12. Ray object store 中不承载 hidden-state payload,训练 tensor 通过 TQ/Mooncake 路径读取。 - -## 8. 后续建议:第一版跑通后再做 - -以下内容不进入第一版开发: - -- Producer HTTP/TQ 复杂重试和 endpoint 熔断; -- Producer 发布 journal,避免重启后重复生成已 clear 样本; -- 自动重启;第一版失败后人工停止整条 pipeline并使用新 `run_id` 重跑; -- Consumer 从最新 checkpoint 自动恢复; -- checkpoint 成功后再 clear 的严格提交窗口; -- TQ owner/storage 整体丢失后的数据重建; -- 多个独立 Consumer 竞争同一 partition; -- lease、ack、超时回收和 exactly-once; -- 动态扩缩容; -- vLLM server 直接写 TQ。 - -第一版先保证:同一个 TQ 能连通、Producer 能并发生产、Consumer 能正确取数训练、每步成功后能及时清理。 diff --git a/docs/transferqueue_integration_plan.md b/docs/transferqueue_integration_plan.md deleted file mode 100644 index cac16ac1..00000000 --- a/docs/transferqueue_integration_plan.md +++ /dev/null @@ -1,145 +0,0 @@ -# verl-SpeCo TransferQueue 落地方案 - -Last updated: 08/21/2026 - -> 目标:在**不修改上游 verl**的前提下,把 SpeCo online 训练里的逐样本特征流 -> 从「`SpecoRayPPOTrainer` driver 中转 + Ray object store」改为「TransferQueue -> 直传」,干掉 driver 这个数据瓶颈,并解锁流式消费与跨副本负载均衡。 -> -> 约束:仅 hook,与 SpeCo 现有 hook 模式一致;TQ 作为独立库使用,**不复用** verl -> 的 `main_ppo_sync` TQ 集成。 - ---- - -## 0. 现状:SpeCo online 特征流的 controller 瓶颈 - -SpeCo online 路径以 `SpecoRayPPOTrainer` 为 hub,所有跨进程大张量都被 driver 串行 -中转,介质是 Ray object store(`ray.put`/`ray.get`/`parallel_put`)。这与 verl 引入 -TQ 想解决的痛点 1:1 对应,只是 verl 干掉的是 `RayPPOTrainer`,我们要干掉的是 -SpeCo 在它之上加的 drafter 管线中转。 - -| # | 流向 | 当前机制 | 是否逐样本 | hook 位置(SpeCo 侧) | -|---|---|---|---|---| -| **a1** | target hidden states(SGLang 采集)-> drafter | `drafter_sample` 塞进 `DataProto.non_tensor_batch` -> driver pop/bucket -> `parallel_put` -> drafter `ray.get` | ✅ | `speco_ray_trainer.py` `generate_sequences_with_speco`;`sglang_adapter.py` `pop_drafter_samples`/`bucket_drafter_samples_by_replica`;`sglang_runtime.py` 组装 `drafter_sample` | -| **a2** | target hidden states(old-logprob hook)-> drafter | actor 前向 hook 截行 -> `ray.put` chunk -> driver 重打包 -> 分发 | ✅ | `oldlogprob_runtime.py` `_install_oldlogprob_hidden_hooks`/`_put_oldlogprob_hidden_refs`;`speco_ray_trainer.py` `_speco_collect_oldlogprob_features` | -| **b2** | target top-logprobs -> drafter(`use_logits=true`) | 随 a1 同一 side-channel | ✅ | `sglang_runtime.py` `target_logprobs`/`hidden_raw_target_logprobs` | -| **d** | rollout tokens -> drafter 训练集 | 随 a1 同一 side-channel(online)/`torch.save` 分片(offline) | ✅ | `sglang_runtime.py`;`speco_worker.py` `_store_rollout_sample` | -| b1 | target **lm_head 权重**(行)-> drafter `TargetHead` | ONE_TO_ALL Ray 分发 | ❌ 参数广播 | `rollout_publish.py` `export_actor_lm_head_weight`/`get_actor_lm_head_weight`;`speco_ray_trainer.py` `_speco_sync_target_lm_head_weight` | -| c | drafter 权重 -> rollout 引擎 | `ray.put` -> driver -> actor;vLLM 末段 ZMQ+SHM,SGLang 进程内 | ❌ 参数广播 | `speco_worker.py` `maybe_publish`;`rollout_publish.py` `update_draft_weights`;`vllm_runtime.py` `BucketedWeightSender` | - -**关键事实**:hidden states 跨进程前一律 CPU 物化(`oldlogprob_runtime.py`、 -`sglang_runtime.py`、`feature_store.py` 均 `.cpu()`),a1 路径下 driver 进程的 -host memory 会真正承载整批 hidden states 并做一次 Ray store 往返。这正是 TQ 要 -消除的往返。 - ---- - -## 1. 为什么不把"替换 feature_store"作为第一刀 - -`TorchShardFeatureStore`(`feature_store.py`)是 `torch.save` 分片 + JSONL manifest, -**只服务于 `collect_only`/`offline`**,不参与 online 热路径。替换它能统一离线存储 -抽象、换更快的分布式后端,但**不解决 controller 瓶颈**,性能收益有限。降级为 -可选尾项(见 §5 P3)。 - ---- - -## 2. 目标方案:TQ 直传逐样本特征流(a1 / a2 / b2 / d) - -### 2.1 角色映射 - -| TQ 角色 | SpeCo 对应 | -|---|---| -| Producer(写) | rollout worker(SGLang 路径,a1/b2/d)/ actor worker(old-logprob 路径,a2)——均在 SpeCo 既有 hook 内 | -| Consumer(读) | drafter worker `collect_rollout_features`(SpeCo 侧) | -| TransferQueueController(control plane) | SpeCo launcher 启动一个 Ray actor;drafter 经 `Sampler`/`StreamingDataLoader` 拉取 | -| Storage backend | `SimpleStorage`(ZMQ,跨节点 CPU 内存);进阶可切 `MooncakeStore`(RDMA,GPU-DRAM) | - -### 2.2 partition / key / 字段设计 - -- `partition_id`:`speco_train`(验证集用 `speco_val`)。 -- `key`:`{uid}_{session_id}_{index}`,与 verl TQ 一致;`uid` SpeCo 已有。 -- `tags`:`global_steps`、`source`∈{`rollout`,`oldlogprob`}、`replica_rank`/`owner_rank`、`status`、`prompt_len`/`response_len`/`seq_len`。ReplayBuffer/负载均衡按 tag 匹配。 -- `fields`(列):`input_ids`、`loss_mask`、`position_ids`、`hidden_states`、`last_hidden_states`/`target`、`target_logprobs`、`hidden_positions`、`prompts`、`responses`。与 `DraftFeatureSample`(`feature_store.py`)字段对齐,便于 online/offline 复用。 - -### 2.3 数据流(目标) - -``` -rollout/actor worker (SpeCo hook) - │ 生成/截取 hidden states 后,就地 tq.kv_batch_put(samples) - ▼ -TransferQueue (SimpleStorage, 跨节点 CPU 内存;可选 MooncakeStore RDMA) - │ control plane 按 sample 粒度追踪 ready 状态,Sampler 跨 drafter 副本均衡 - ▼ -drafter worker - │ tq.kv_batch_get / StreamingDataLoader 消费 → 喂入既有 DataBuffer / collect_online_data - ▼ -drafter 训练 (不变) -``` - -driver 只下发触发与轻量 key/meta,**不再承载 hidden states**。 - ---- - -## 3. 落地改动点(全部在 SpeCo 侧,hook-only) - -### 3.1 启动与配置 -- `draft_train_launcher.py` / `main.py`:`tq.init(config.transfer_queue)`;起 `TransferQueueController.remote(Sampler)`。 -- `config/speco_base.yaml`:新增 `drafter.transfer_queue` 块(backend、partition、enable 开关)。参考 verl `ppo_trainer.yaml` 的 `transfer_queue:` 结构,但**独立配置**,不复用 verl 的。 - -### 3.2 Producer 侧 -- **a1/b2/d(SGLang)**:`sglang_runtime.py` 组装 `drafter_sample` 处(~1594-1648),增加 `tq.kv_batch_put`;返回给 driver 的 `drafter_sample` 只保留 key/meta(或整段不再走 DataProto side-channel,driver 仅触发)。 -- **a2(old-logprob)**:`oldlogprob_runtime.py` `_put_oldlogprob_hidden_refs`(~216),把 `ray.put(hidden_chunk)` 换成 `tq.kv_batch_put`;`OLD_LOGPROB_HIDDEN_CHUNK_REFS_KEY` 改为 TQ key 列表。 - -### 3.3 Consumer 侧 -- `speco_worker.py` `collect_rollout_features`(~665):把 `_resolve_ray_object_ref`/`_resolve_hidden_state_chunks`(`ray.get`)换成 `tq.kv_batch_get`;`_dispatch_nd_compute`(~159)的 `parallel_put` 退化为只传 key(或 drafter 直接从 TQ Sampler 拉,driver 不参与分发)。 -- drafter 内部 `DataBuffer`/`collect_online_data`(`base_trainer.py`)保持不变,只是数据来源由 `ray.get` 改为 TQ get。 - -### 3.4 Driver 侧 -- `speco_ray_trainer.py`:`_speco_collect_rollout_features_rpc`/`speco_collect_rollout_features`(~351)、`_speco_collect_oldlogprob_features`(~1114)不再搬数据,只做触发/传 key;`bucket_drafter_samples_by_replica` 可由 TQ `Sampler` 替代(逐步迁移,先保留作回退)。 - -### 3.5 不改动 -- **b1(lm_head 权重)、c(drafter 权重)**:保持现状。与 verl 上游一致(权重不走 TQ),且 c 的 vLLM 末段已有专用 ZMQ+SHM 通道。 -- verl 本体:零改动。 - ---- - -## 4. 收益与边界(诚实评估) - -### 收益 -1. **去掉 driver 对 hidden states 的 host-memory 中转 + Ray store 往返**:producer 直存 TQ,consumer 直取,driver 不再承载整批特征。 -2. **流式消费**:drafter 在样本 ready 时即可消费,不必等整批 `generate_sequences` 返回,采集与训练可重叠。 -3. **跨 drafter 副本负载均衡**:TQ `Sampler`/`RankAwareSampler` 替代手写 `bucket_drafter_samples_by_replica`/`owner_rank` 分配。 -4. **(若采纳 P3)统一 online/collect_only/offline 存储**:同一 TQ partition,`collect_only` 写、`offline` 读,消掉 on-disk 分片层。 - -### 边界 / 不解决的事 -- 只优化**特征采集**这一子阶段,**不加速** rollout 本身、actor update、reward;e2e 增益取决于该子阶段在 step 中的占比。 SpeCo README 的 20% rollout / 11% e2e 提升来自 acceptance length,与本方案是不同机制,不要混为一谈。 -- **权重同步(b1/c)不放进 TQ**,与 verl 上游保持一致。 -- hidden states 跨进程前**仍需 CPU 物化**(现状如此);要避免物化需切 `MooncakeStore` RDMA,属进阶项。 -- 引入 TQ 依赖与一个 control-plane Ray actor,增加少量运维面。 - -### 风险 -- TQ 与 SpeCo 现有 `owner_rank`/`replica_rank` 路由语义需对齐(Sampler 要复刻「按 owner 分桶」语义,否则样本会错配 drafter 副本)。 -- old-logprob 的 chunk 拆分(`hidden_states_ref_chunks`)映射到 TQ 列式存储时,需保证 chunk meta 与 key 的一致性。 -- 回退路径:保留 `enable_transfer_queue=False` 时走原 Ray 路径,渐进切换。 - ---- - -## 5. 分阶段实施 - -| 阶段 | 范围 | 产出 | -|---|---|---| -| **P0** | a1(SGLang hidden states)走 TQ 直传;drafter `kv_batch_get` 消费;driver 仅触发 | 验证 controller-bypass 闭环 + 正确性 | -| **P1** | a2(old-logprob hidden states)走 TQ;chunk 拆分映射 TQ 列 | 覆盖第二条采集路径 | -| **P2** | b2(top-logprobs)+ d(tokens)随 a1 同 partition 传输;Sampler 替代手写 bucket | 完整特征流 + 跨副本均衡 | -| **P3(可选)** | `TorchShardFeatureStore` → TQ partition,统一 online/collect_only/offline | 离线工作流统一 | - -每个阶段保留 `enable_transfer_queue` 开关与原 Ray 路径回退。 - ---- - -## 6. 待确认决策 - -1. **TQ backend**:`SimpleStorage`(CPU 内存,默认)起步,还是直接上 `MooncakeStore`(RDMA,省 CPU 物化)?后者依赖 RDMA 网络,建议 P0 用 SimpleStorage。 -2. **drafter 消费模式**:`kv_batch_get`(主动拉,改动小)还是 `StreamingDataLoader`(全自动流式,改动大、收益高)?建议 P0 用前者,P2 再考虑后者。 -3. **driver 角色**:P0 先保留 driver 传 key(最小改动),还是直接让 drafter 从 TQ Sampler 自取(driver 彻底退出数据路径)?前者风险低,建议 P0 用前者。 -4. **是否做 P3**:离线统一是否在本次范围内,还是单独立项。 diff --git a/docs/vllm_direct_hidden_state_cotrain_plan.md b/docs/vllm_direct_hidden_state_cotrain_plan.md deleted file mode 100644 index 6c8731f4..00000000 --- a/docs/vllm_direct_hidden_state_cotrain_plan.md +++ /dev/null @@ -1,666 +0,0 @@ -# Co-train 复用 Producer 直接从 vLLM 获取 Hidden State 的实施方案 - -## 1. 目标与边界 - -本文方案面向 **RL 与草稿模型共同训练(co-train)**,目标是把 target hidden state 的来源从 actor 的 old-logprob 前向切换为 vLLM: - -```text -原方案:vLLM 生成 response - -> actor 为 PPO 计算 old_log_probs - -> 在 actor forward 内通过 hook/output_hidden_states 抓 hidden state - -> 草稿模型训练 - -新方案:vLLM 生成 response - -> 使用 prompt + response 再向 hidden-state vLLM 发起 prefill 请求 - -> vLLM 返回 hidden-state 文件位置,客户端读取并归一化 - -> 现有 scheduler -> SpecoWorker -> drafter buffer -> 草稿模型训练 -``` - -需要特别区分两个动作: - -- actor 的 `compute_log_prob()` **仍然保留**,因为 PPO 更新 actor 需要 `old_log_probs`。 -- 删除的是 old-logprob 前向中专为草稿训练增加的 hidden-state 捕获、拼接、CPU copy 和 Ray put 逻辑。 - -第一版不修改独立训练的 TQ producer/consumer 行为,不让 co-train 使用独立训练的文件读取、EOS、TQ owner 和离线 dataloader。复用的是 producer 中已经验证过的 **vLLM 请求、并发、重试、文件读取、token 对齐和样本归一化能力**。 - -## 2. 当前 co-train 全流程 - -### 2.1 运行时角色 - -| 角色 | 所在位置 | 当前职责 | -|---|---|---| -| `RayPPOTrainer` driver | `verl_speco/trainer/speco_ray_trainer.py` | 驱动 rollout、old-logprob、actor update,并通过 drafter scheduler 决定采样和训练时机 | -| vLLM rollout workers | verl rollout worker group | 根据 prompt 生成 response;当前 vLLM 路径不直接给 drafter hidden state | -| actor workers | `actor_rollout_wg` | 计算 PPO 所需 old log-prob;当前还承担 hidden-state 捕获 | -| drafter workers | `drafter_wg` 中的 `SpecoWorker` | 接收按 owner 分桶的样本,写入在线 buffer,执行 drafter train/publish/checkpoint | -| `DrafterScheduler` | driver 进程内普通 Python 对象 | 决定本 step 是否采集、是否训练,并组织 collection 的 stage/commit/finalize | - -这里的 scheduler 不是另一个服务,也不读取数据;它只负责控制顺序和事务状态。 - -### 2.2 rollout 生成的数据 - -rollout 后,driver 持有 `DataProto batch`。与本方案直接相关的 tensor 通常为: - -```python -batch.batch["prompts"] # [B, Pmax],左 padding -batch.batch["responses"] # [B, Rmax],右 padding -batch.batch["attention_mask"] # [B, Pmax + Rmax] -batch.batch["response_mask"] # [B, Rmax],若存在则表示有效 response token -``` - -单个有效样本会还原为: - -```python -prompt_ids: Tensor[P] -response_ids: Tensor[R] -input_ids = torch.cat([prompt_ids, response_ids]) # Tensor[P + R] -``` - -padding token 不能发给 hidden-state vLLM;必须通过 mask 去掉。 - -### 2.3 当前 old-logprob hidden-state 路径 - -入口位于 `SpecoRayPPOTrainer._speco_online_fit_hooks()` 安装的 `compute_old_log_prob_with_speco()` 包装函数: - -1. `_speco_plan_drafter_collection(OLD_LOGPROB)` 调 scheduler,决定当前 `global_step` 是否采集。 -2. `_speco_build_oldlogprob_collect_plan(batch)` 选择样本、hidden positions 和 owner rank。 -3. driver 把以下控制数据写进 `batch_td`: - - ```python - OLD_LOGPROB_COLLECT_MASK_KEY - OLD_LOGPROB_HIDDEN_POSITIONS_KEY - OLD_LOGPROB_HIDDEN_POSITION_MASK_KEY - OLD_LOGPROB_OWNER_RANK_KEY - OLD_LOGPROB_AUX_LAYER_IDS_KEY - OLD_LOGPROB_HIDDEN_CAPTURE_IMPL_KEY - OLD_LOGPROB_HIDDEN_LAYOUT_KEY - ``` - -4. `actor_rollout_wg.compute_log_prob(batch_td)` 远程执行 actor 前向。 -5. `verl_speco/integration/oldlogprob_runtime.py` 根据 `forward_hook` 或 `output_hidden_states` 捕获指定层,并把结果作为 tensor、Ray object ref 或分块 ref 返回。 -6. driver 的 `_speco_collect_oldlogprob_features()` 把 actor 输出还原为逐样本字典: - - ```python - sample = { - "input_ids": Tensor[1, P + R], - "prompts": Tensor[1, P], - "responses": Tensor[1, R], - "hidden_positions": Tensor[1, Hrows], - "hidden_states": Tensor[1, Hrows, Hdim], # 或 *_ref / *_ref_chunks - "hidden_states_layout": "dflash_aux" | "eagle3_aux_plus_last", - "hidden_position_start": int, - "hidden_position_end": int, - "global_step": int, - "replica_rank": int, - } - ``` - -7. `OldLogProbCollectionAdapter.prepare_payload()` 根据显式 `owners` 把样本分到 drafter owner buckets。 -8. scheduler 执行 collection transaction,Ray RPC 参数本质上是每个 owner 对应的 `list[dict]`。 -9. `SpecoWorker._commit_rollout_features(collection_id, samples)` 解析 hidden tensor/ref,调用 `_store_rollout_sample()`。 -10. `_store_rollout_sample()` 调 `DrafterBaseTrainer.collect_online_data()`,把 CPU 数据写入当前 step 或跨 step buffer。 -11. `update_actor_with_speco()` 调 `_speco_on_before_actor_update()`;scheduler 根据已收集数据产生 training plan,然后执行 actor update 和 drafter training。 - -因此,现有 worker、buffer 和训练后半段并不关心 hidden state 是由 actor 还是 vLLM 产生。需要替换的主要是第 2~6 步的数据生产方式。 - -## 3. 当前 standalone producer 的详细流程 - -### 3.1 哪些部分可以复用 - -`verl_speco/standalone_tq_producer.py` 当前是一个三段式异步流水线: - -```text -read_inputs - -> request_queue - -> N 个 request_worker - -> publish_queue - -> publish_results - -> TQ -``` - -其中只有最后的 TQ publish 和最前面的文件 reader 是 standalone 专属。中间部分已经包含 co-train 需要的核心能力。 - -#### `VllmFeatureClientPool` - -文件:`verl_speco/producer/vllm_feature_client.py` - -职责: - -- 解析多个 `VllmEndpoint`; -- 维护全局并发 semaphore 和每 endpoint semaphore; -- 优先选择当前 inflight 较少的 endpoint; -- 通过 OpenAI completions 接口发送 token IDs; -- 对连接错误、read error、超时等执行指数退避重试; -- 从响应的 `kv_transfer_params.hidden_states_path` 取得 safetensors 路径; -- 等待文件完成,读取 hidden state、token IDs 和相关字段; -- 读取完成后删除临时 safetensors 与 lock 文件。 - -请求不是让 vLLM 再生成一段文本,而是一次 prefill 请求: - -```python -await client_pool.request_prefill( - prompt_token_ids=request.vllm_prompt_token_ids, - sample_id=request.sample_id, -) -``` - -返回的 `RawVllmFeature` 仍是 vLLM 原始坐标系下的数据,例如: - -```python -RawVllmFeature( - token_ids=Tensor[Tprefill], - hidden_states=Tensor[Tprefill, L, D], - hidden_position_start=..., - hidden_position_end=..., - ..., -) -``` - -其中 `L` 是导出的 target layer 数量,`D` 是 target hidden size。 - -#### `prepare_generated_prefill_request` - -文件:`verl_speco/producer/input_reader.py` - -standalone 遇到仅有 prompt 的数据时,先生成 response,再调用该函数把 prompt 和生成结果组成训练请求。核心规则是: - -```python -full_ids = prompt_ids + response_ids -vllm_prompt_token_ids = full_ids[:-1] -``` - -去掉最后一个 token 的原因是:位置 `i` 的 target hidden state 用于预测后续 token,最后一个 token 后面没有本样本内的监督 token。co-train 已经拥有 rollout response,所以只需要直接执行这一步,不需要再次生成 response。 - -当前 `TokenizedRequest` 包含: - -```python -TokenizedRequest( - sequence_no: int, - sample_id: str, - input_ids: list[int], # prompt + response - loss_mask: list[float], # prompt 为 0,有效 response 为 1 - position_ids: list[int], - feature_positions: list[int], # 选中的 target hidden 绝对位置 - draft_position_ids: list[int], - source_metadata: dict, - vllm_prompt_token_ids: list[int], # 发给 vLLM 的 full_ids[:-1] -) -``` - -#### `feature_from_vllm_payload` - -文件:`verl_speco/trainer/target_feature_replay.py` - -该函数把 `RawVllmFeature + TokenizedRequest + FeatureContract` 转成算法训练侧统一使用的 `DraftFeatureSample`。它负责: - -- 检查 vLLM 返回的 token IDs 是否与请求一致; -- 检查 hidden rows、层数、hidden size 和 layout; -- 按 `feature_positions` 选择训练窗口; -- 对齐 `input_ids`、`loss_mask`、positions 与 hidden state; -- hidden state 不完整或位置对不上时拒绝该样本,不把错误样本交给训练。 - -典型结果: - -```python -DraftFeatureSample( - input_ids=Tensor[T], - loss_mask=Tensor[T], - hidden_states=Tensor[Hrows, L * D], - target_logprobs=None, - position_ids=Tensor[T], - feature_positions=Tensor[Hrows], - draft_position_ids=Tensor[Hrows], - metadata={...}, -) -``` - -这里的协议和算法处理应继续由已有 `DraftFeatureSample`、backend 和 contract 决定,不能在新 co-train 组件里再次硬编码 DSpark。 - -### 3.2 哪些部分不能直接搬入 co-train - -以下 standalone 逻辑不能原样调用: - -- 从 JSONL 循环读 epoch;co-train 的输入来自当前 rollout `DataProto`。 -- prompt-only 时调用生成接口;co-train 的 response 已经生成。 -- `sequence_no/run_id/tag/EOS/max_pending_samples`;这些用于 TQ 流式生产消费,不属于单个 RL step。 -- `publish_results()` 和 TQ clear;co-train 已有 scheduler collection transaction 和 worker buffer。 -- standalone owner/consumer 生命周期;co-train 由 Ray trainer 和 worker group 管理。 - -正确的复用方式是抽取“给定 tokenized rollout sample,异步取得并归一化 hidden state”的核心,而不是在 co-train 内启动一个 `standalone_tq_producer` 进程。 - -## 4. 建议的新数据流 - -### 4.1 完整顺序 - -```text -1. rollout vLLM 生成 response -2. driver 得到 DataProto(prompts, responses, masks) -3. scheduler 判断本 step 是否需要采集 VLLM_PREFILL -4. driver 按采样计划选择样本、去 padding、构造 TokenizedRequest -5. CotrainVllmFeatureProducer.submit_batch() 提交并发 prefill -6. hidden-state vLLM endpoint 执行 prompt+response[:-1] prefill -7. client 读取 safetensors,执行 token/shape/position 对齐 -8. 得到 DraftFeatureSample;失败或不完整样本在此处过滤 -9. 将 DraftFeatureSample 转为现有 SpecoWorker collection sample -10. VllmPrefillCollectionAdapter 按 replica owner 分 buckets -11. scheduler stage -> Ray commit RPC -> finalize -12. SpecoWorker._store_rollout_sample() -> collect_online_data() -> buffer -13. scheduler 产生 training plan -14. actor update 与 drafter training 按现有顺序执行 -``` - -步骤 5 提交后不应立刻阻塞等待。driver 可以继续 reward、reference log-prob、advantage、actor old-logprob 等工作;在 drafter collection 必须完成的边界再 `await/result()`。这样 vLLM prefill 与 RL 侧计算重叠。 - -### 4.2 vLLM 请求的具体 token 对齐 - -给定一个 rollout 样本: - -```python -prompt_ids = prompts[i][prompt_mask] # [P] -response_ids = responses[i][response_mask] # [R] -full_ids = cat(prompt_ids, response_ids) # [P + R] -prefill_ids = full_ids[:-1] # [P + R - 1] -``` - -构造: - -```python -loss_mask = zeros(P + R) -loss_mask[P:P + R] = 1 -``` - -然后再应用现有 collection plan 的窗口限制。必须保证: - -```text -返回 token_ids == prefill_ids -hidden rows 能覆盖 feature_positions -feature_positions 非空 -选中区域对应的 loss_mask 中存在有效训练 token -``` - -只含 prompt、有效 response 长度为 0、hidden rows 为 0、token 不一致或窗口为空的样本,都在 producer 转换阶段丢弃,不进入 scheduler payload。这样 producer 的“发布成功数”和 consumer 的“可接收数”天然一致,不会把无效条目带入 collection transaction。 - -### 4.3 scheduler 到 worker 的样本格式 - -建议保留 worker 当前已经支持的 collection sample 外形,不大改训练后半段: - -```python -worker_sample = { - "input_ids": Tensor[1, T], - "prompts": Tensor[1, P], - "responses": Tensor[1, R], - "hidden_positions": Tensor[1, Hrows], - "hidden_states": Tensor[1, Hrows, HiddenWidth], - "hidden_states_layout": str, - "hidden_position_start": int, - "hidden_position_end": int, - "global_step": int, - "replica_rank": int, -} -``` - -`HiddenWidth` 取决于现有 backend/layout。例如多个层已经按最后一维拼接时为 `L * D`。该转换必须调用 `DraftFeatureSample` 已有字段和 metadata,不在 adapter 中按算法猜测。 - -`replica_rank` 不是 vLLM 返回的数据,而是 driver 根据现有 drafter owner 路由计划为样本分配的控制字段。scheduler 只用它决定该样本发给哪个 `SpecoWorker` owner。 - -## 5. 代码修改方案 - -### 5.1 新增 co-train producer 核心 - -新增:`verl_speco/producer/cotrain_vllm_feature_producer.py` - -建议接口: - -```python -@dataclass -class CotrainFeatureRequest: - batch_index: int - owner_rank: int - request: TokenizedRequest - prompt_ids: torch.Tensor - response_ids: torch.Tensor - - -@dataclass -class CotrainFeatureResult: - batch_index: int - owner_rank: int - sample: DraftFeatureSample - - -class CotrainVllmFeatureProducer: - def __init__(self, config, *, contract: FeatureContract): ... - - def submit_batch( - self, - requests: list[CotrainFeatureRequest], - ) -> Future[list[CotrainFeatureResult]]: ... - - async def _produce_one( - self, - request: CotrainFeatureRequest, - ) -> CotrainFeatureResult | None: ... - - def close(self) -> None: ... -``` - -内部直接复用: - -```python -raw = await self.client_pool.request_prefill(...) -sample = feature_from_vllm_payload(raw, request.request, self.contract) -``` - -组件应持有一个长期存在的 `VllmFeatureClientPool`,不能每 step 重建 HTTP client、semaphore 和线程池。由于 PPO driver 主流程通常是同步代码,第一版可让组件内部持有一个后台 asyncio event loop thread,`submit_batch()` 返回 `concurrent.futures.Future`。训练结束时统一 `close()`,取消未完成任务并关闭 HTTP client。 - -### 5.2 增加从 rollout tensor 构造请求的函数 - -修改:`verl_speco/producer/input_reader.py` - -新增纯函数,复用现有长度截断、feature window、position 和 loss-mask 规则: - -```python -def build_rollout_prefill_request( - *, - sample_id: str, - sequence_no: int, - prompt_ids: Sequence[int], - response_ids: Sequence[int], - producer_cfg, - source_metadata: dict, -) -> TokenizedRequest: - ... -``` - -它不接收文本、不调用 tokenizer、不生成 response,只做: - -1. 拼接有效 prompt/response token; -2. 按已有 `max_feature_length` 等规则截取 response; -3. 建 loss mask、positions; -4. 设置 `vllm_prompt_token_ids=full_ids[:-1]`。 - -必须把 standalone 与 co-train 的公共构造逻辑下沉到同一个私有 helper,避免两个路径以后出现 off-by-one 或截断规则差异。 - -### 5.3 扩展 scheduler 的 collection source - -修改: - -- `verl_speco/trainer/scheduler/schedule_types.py` -- `verl_speco/trainer/scheduler/collection_adapter.py` -- `verl_speco/trainer/scheduler/drafter_scheduler.py` - -新增: - -```python -class DrafterCollectionSource(str, Enum): - SGLANG = "sglang" - OLD_LOGPROB = "oldlogprob" - VLLM_PREFILL = "vllm_prefill" -``` - -新增 `VllmPrefillCollectionAdapter`。它只负责: - -- 校验每个 sample 有 `replica_rank`; -- 使用 `_build_payload()` 按 owner 分桶; -- 设置 `CollectionPayload.source=VLLM_PREFILL`。 - -它不负责请求 vLLM、不解码 hidden state、不实现算法逻辑。 - -同时更新 collection source 的稳定排序值、adapter registry 和 metrics source label。 - -### 5.4 在 `SpecoRayPPOTrainer` 接入异步生产 - -修改:`verl_speco/trainer/speco_ray_trainer.py` - -新增或调整以下职责: - -```python -def _speco_vllm_prefill_collection_requested(self) -> bool: ... -def _speco_vllm_prefill_collection_enabled(self) -> bool: ... -def _speco_get_cotrain_vllm_producer(self) -> CotrainVllmFeatureProducer: ... -def _speco_build_vllm_prefill_requests(self, batch, collection_plan): ... -def _speco_submit_vllm_prefill_collection(self, batch): ... -def _speco_finish_vllm_prefill_collection(self) -> int: ... -def _speco_close_vllm_prefill_producer(self) -> None: ... -``` - -接入点建议如下: - -1. rollout 返回 `gen_batch_output` 并合并成训练 batch 后,调用 `_speco_submit_vllm_prefill_collection(batch)`。 -2. 提交函数先调用 scheduler 的 `plan_collection(VLLM_PREFILL)`;未命中 interval 时不发 HTTP 请求。 -3. 继续执行 reward、ref、old-logprob 和 advantage。 -4. 在 `_speco_on_before_actor_update()` 生成 training plan 之前调用 `_speco_finish_vllm_prefill_collection()`: - - 等待 Future; - - 过滤失败样本; - - 转成 worker sample; - - adapter 分桶; - - `_speco_execute_collection()`。 -5. 原 `compute_old_log_prob_with_speco()` 在该模式下走普通 `original_compute_old_log_prob()`,不再注入 hidden capture keys。 -6. fit 的 `finally` 中关闭 producer。 - -这里“等待点必须在 training plan 之前”是必要条件。否则 scheduler 看到的 buffer version 仍是旧值,本 step 可能错误判断没有可训练样本。 - -### 5.5 worker 和训练侧尽量不改 - -`verl_speco/workers/speco_worker.py` 的以下路径可以直接复用: - -```text -collect_rollout_features / collection transaction RPC - -> _commit_rollout_features - -> _store_rollout_sample - -> DrafterBaseTrainer.collect_online_data -``` - -第一版只在必要时增加一个从 `DraftFeatureSample` 转现有 sample dict 的小 helper;不要新增另一套 buffer,也不要让 worker 连接 standalone TQ。 - -如果 `DraftFeatureSample.to_training_item()` 与 `collect_online_data()` 的 metadata 表达存在差异,应在一个公共转换函数中补齐,而不是在 driver、adapter、worker 分别写一套字段映射。 - -### 5.6 禁用 old-logprob hidden capture,但保留 PPO old-logprob - -修改配置判定和 hook 分支: - -```yaml -collect_hidden_states_from_sgl: false -collect_hidden_states_from_old_logprob: false -collect_hidden_states_from_vllm: true -``` - -当 `collect_hidden_states_from_vllm=true` 时: - -- 不设置 `OLD_LOGPROB_HIDDEN_*` keys; -- 不安装/启用 `oldlogprob_runtime` hidden hooks; -- 不调用 `_speco_collect_oldlogprob_features()`; -- 仍执行标准 `actor_rollout_wg.compute_log_prob()`,得到 PPO 的 log-prob 和 entropy。 - -三种来源第一版必须互斥: - -```python -sum([ - collect_hidden_states_from_sgl, - collect_hidden_states_from_old_logprob, - collect_hidden_states_from_vllm, -]) <= 1 -``` - -## 6. 配置建议 - -修改:`verl_speco/config/actor/actor.yaml` 或本仓实际承载 drafter training 默认值的配置文件,并在示例脚本暴露关键参数。 - -建议结构: - -```yaml -actor_rollout_ref: - rollout: - drafter: - training: - mode: online - collect_hidden_states_from_sgl: false - collect_hidden_states_from_old_logprob: false - collect_hidden_states_from_vllm: true - - vllm_feature_source: - endpoints: - - http://127.0.0.1:8000/v1 - - http://127.0.0.1:8001/v1 - model: /path/to/target-model - max_inflight_requests: 128 - per_endpoint_concurrency: 64 - request_timeout_seconds: 600 - max_retries: 3 - retry_base_delay_seconds: 1.0 - max_sequence_length: 8192 -``` - -`target_layer_ids`、`max_feature_length`、hidden layout、算法类型等应继续读取现有 drafter 配置,不在 `vllm_feature_source` 重复定义。`FeatureContract` 也从同一份运行配置构建,从而保证 vLLM 导出层和 trainer 预期一致。 - -### vLLM 服务要求 - -当前 `VllmFeatureClientPool` 使用 OpenAI HTTP endpoint 和 `kv_transfer_params.hidden_states_path`。因此第一版要求: - -- co-train 可访问一个或多个已启动的 hidden-state vLLM 服务; -- 服务加载的 target model/tokenizer 与 rollout/actor 使用的版本一致; -- `extract_hidden_states` 的 layer IDs 与 drafter contract 一致; -- vLLM hidden 文件目录对 driver 可见; -- prefix caching 对 hidden-state 导出必须关闭或已验证能返回所有所需 rows。 - -verl 内部 rollout vLLM worker 不一定天然暴露当前 client 所需的 OpenAI 地址和共享 hidden 文件路径。第一版建议使用独立启动的 hidden-state vLLM endpoints。后续若 rollout 服务能够暴露同等接口,再把 endpoint discovery 接入 worker group,producer 核心无需改变。 - -## 7. 是否在 co-train 中使用 TQ - -### 7.1 第一版建议:不使用 TQ - -第一版直接把 CPU hidden tensor 放入现有 scheduler payload,必要时沿用 Ray object ref/chunk ref 机制。理由: - -- co-train 已经有 scheduler collection transaction、owner 路由和 worker buffer; -- 当前 standalone TQ 的 run ID、pending、EOS、clear 语义针对跨进程无限流,不适合直接套在单个 RL step 上; -- 少改 worker 和生命周期,能先验证 vLLM hidden 与 actor hidden 的数值/训练等价性。 - -此时的控制面和数据面是: - -```text -控制面:driver -> scheduler -> Ray RPC(collection_id, owner bucket) -数据面:CPU tensor 随 Ray 参数,或先 ray.put 后传 ObjectRef -``` - -### 7.2 第二阶段可选:TQ 承载大 tensor - -若 Ray object store 压力明显,再让 producer 将 `DraftFeatureSample` 编码后写 TQ,而 scheduler sample 只传: - -```python -{ - "feature_key": str, - "replica_rank": int, - "global_step": int, -} -``` - -worker commit 时按 key `get + decode`,commit 成功后 clear,rollback 时保留或清理。这需要定义 co-train 专属的 step-scoped key 和事务清理规则,不能复用 standalone EOS。该阶段会增加失败恢复复杂度,不建议和第一版一起提交。 - -## 8. 错误处理与一致性 - -### 8.1 单样本错误 - -以下错误在 `_produce_one()` 内记录 sample ID、batch index、endpoint 和原因,然后丢弃该样本: - -- response 为空; -- token IDs 不一致; -- hidden rows 为 0 或覆盖不了选择位置; -- layer/hidden size/layout 不符合 contract; -- 截断后没有有效训练 token。 - -只有成功转换成 `DraftFeatureSample` 的样本才计入 `CollectionPayload.collected_samples`。 - -### 8.2 请求级错误 - -连接/读取错误先使用现有 client pool 重试。超过最大重试后,第一版建议默认让当前 collection 失败并终止本 step,而不是静默用不完整 batch 训练;可以后续增加 `failure_policy=fail_step|skip_sample|skip_collection`。 - -### 8.3 多 rank 一致性 - -vLLM 请求和样本过滤都发生在 driver;driver 形成最终成功样本列表后才按 owner 分桶并发 RPC。因此每个 drafter owner 收到的数量是 scheduler 已知的,不让各训练 rank 自行请求 vLLM、各自过滤。这避免某 rank 接受、另一个 rank 拒绝后进入不同 collective 顺序。 - -## 9. 指标与日志 - -建议增加: - -```text -drafter/vllm_prefill/candidate_samples -drafter/vllm_prefill/submitted_samples -drafter/vllm_prefill/succeeded_samples -drafter/vllm_prefill/dropped_samples -drafter/vllm_prefill/request_elapsed_sec -drafter/vllm_prefill/wait_elapsed_sec -drafter/vllm_prefill/overlap_elapsed_sec -drafter/vllm_prefill/payload_mib -drafter/vllm_prefill/retry_count -drafter/vllm_prefill/per_endpoint_inflight -``` - -每次 collection 至少记录:`global_step`、`collection_id`、候选数、提交数、成功数、按 owner 分桶数量、hidden rows、payload bytes 和等待时间。单样本拒绝日志记录 sample ID 和失败检查项,但不要打印完整 token 或 hidden tensor。 - -## 10. 测试计划 - -### 10.1 单元测试 - -1. padded `prompts/responses` 能还原正确有效 token。 -2. `prefill_ids == prompt_ids + response_ids[:-1]` 的边界测试,包括 response 长度 0/1。 -3. co-train request builder 与 standalone builder 对同一 token 输入产生一致的 loss mask、feature positions 和截断结果。 -4. 多 endpoint 并发、重试和 endpoint 选择沿用现有 client pool 测试。 -5. token 不一致、hidden rows=0、缺层、空窗口均被 producer 拒绝,且不进入 payload。 -6. `VllmPrefillCollectionAdapter` 能按 replica owner 正确分桶。 -7. vLLM source 开启时 old-logprob batch 不包含任何 hidden capture key。 -8. 标准 old-logprob 结果仍正确返回给 PPO。 -9. producer Future 在训练计划生成前完成 collection;关闭时无残留线程和 HTTP client。 - -### 10.2 集成测试 - -用小模型和两个 hidden-state endpoints 运行数个 co-train steps,对比: - -- old-logprob capture 与 vLLM prefill 的 token IDs、positions、hidden shape; -- 相同权重和样本下 drafter loss/metrics 是否接近; -- actor PPO metrics 是否不变; -- collect interval 未命中时没有 vLLM hidden 请求; -- endpoint 临时失败时重试后能继续; -- 多 drafter owner 下每个 owner 收到预期样本数。 - -## 11. 建议实施顺序与文件清单 - -### 阶段 A:抽取公共 producer 能力 - -- 修改 `verl_speco/producer/input_reader.py`:增加 rollout-token request builder,共享截断/对齐 helper。 -- 新增 `verl_speco/producer/cotrain_vllm_feature_producer.py`:长期 client pool、异步 batch submit、转换与过滤、关闭逻辑。 -- 不改 standalone TQ publish 行为。 - -### 阶段 B:接入 scheduler 和 driver - -- 修改 `verl_speco/trainer/scheduler/schedule_types.py`:增加 `VLLM_PREFILL`。 -- 修改 `verl_speco/trainer/scheduler/collection_adapter.py`:增加 owner 分桶 adapter。 -- 修改 `verl_speco/trainer/scheduler/drafter_scheduler.py`:注册 adapter。 -- 修改 `verl_speco/trainer/speco_ray_trainer.py`:提交 future、等待、构造 payload、执行 collection、关闭 producer。 - -### 阶段 C:配置、示例和验证 - -- 修改默认 drafter training 配置:增加 source 开关与 vLLM client 参数。 -- 修改 co-train example:关闭 old-logprob hidden capture,填写 hidden-state endpoints。 -- 增加 request builder、adapter、driver hook 和端到端测试。 -- 用同一批固定 token 对照 actor-captured hidden 与 vLLM hidden,再进行正式性能测试。 - -### 阶段 D:可选 TQ 数据面 - -- 仅在 Ray object store 成为瓶颈后实施。 -- 新增 co-train TQ key/ref adapter、worker decode 和 collection finalize/rollback 清理。 -- 不改变独立训练现有 TQ 协议。 - -## 12. 最终推荐 - -推荐第一版采用: - -```text -独立 hidden-state vLLM endpoints - + 复用 VllmFeatureClientPool - + 复用 TokenizedRequest / FeatureContract / feature_from_vllm_payload - + 新增 VLLM_PREFILL scheduler source - + 复用现有 SpecoWorker collection/buffer/train - + 暂不在 co-train 中引入 TQ -``` - -这样修改范围集中在“hidden-state 来源”和“异步接入点”,不会重写已经稳定的 drafter worker 与训练逻辑,也不会影响 PPO 必需的 actor old-logprob 计算。等这一版验证 hidden 数值、loss 和吞吐后,再决定是否把 Ray 中的大 tensor 数据面替换为 TQ。 From c3faa4627d25247e91c4f76f89bac818a502dce7 Mon Sep 17 00:00:00 2001 From: xyr <1615681145@qq.com> Date: Tue, 1 Sep 2026 10:45:59 +0800 Subject: [PATCH 47/50] style: apply ruff format to standalone_resume.py --- verl_speco/trainer/standalone_resume.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/verl_speco/trainer/standalone_resume.py b/verl_speco/trainer/standalone_resume.py index a8c274aa..206388e9 100644 --- a/verl_speco/trainer/standalone_resume.py +++ b/verl_speco/trainer/standalone_resume.py @@ -71,9 +71,7 @@ def save_standalone_resume( "input_fingerprint": build_input_fingerprint(input_path), } metadata_path = checkpoint_dir / RESUME_METADATA_NAME - metadata_temporary = metadata_path.with_suffix( - metadata_path.suffix + ".incomplete" - ) + metadata_temporary = metadata_path.with_suffix(metadata_path.suffix + ".incomplete") metadata_temporary.write_text( json.dumps(metadata, ensure_ascii=False, indent=2) + "\n", encoding="utf-8", From 11722c61805d0ad55f6ae1b10af950aca91c934f Mon Sep 17 00:00:00 2001 From: xyr <1615681145@qq.com> Date: Thu, 3 Sep 2026 15:46:39 +0800 Subject: [PATCH 48/50] feat: normalize vLLM plus_last features with target final norm extract_hidden_states collects layer residuals before the target final norm. For *_plus_last layouts, load only the frozen target's final norm from the checkpoint (meta-built architecture, norm weights read into CPU RAM) and apply it exactly once to the final supervision block, in both the standalone TQ producer and vLLM-file replay. Extract checkpoint tensor loading into verl_speco/checkpoint_tensor.py so target_head can share it, and bump the replay cache key to invalidate stale pre-norm entries. --- tests/unit/test_target_feature_replay.py | 204 +++++++++++++++++++- tests/unit/test_tq_producer.py | 79 +++++++- verl_speco/checkpoint_tensor.py | 58 ++++++ verl_speco/models/target/target_head.py | 42 +--- verl_speco/standalone_tq_producer.py | 13 +- verl_speco/trainer/target_feature_replay.py | 64 +++++- 6 files changed, 412 insertions(+), 48 deletions(-) create mode 100644 verl_speco/checkpoint_tensor.py diff --git a/tests/unit/test_target_feature_replay.py b/tests/unit/test_target_feature_replay.py index a608611a..8b7730c0 100644 --- a/tests/unit/test_target_feature_replay.py +++ b/tests/unit/test_target_feature_replay.py @@ -13,6 +13,7 @@ # limitations under the License. import threading +from dataclasses import replace from types import SimpleNamespace import pytest @@ -22,10 +23,13 @@ from verl_speco.trainer.feature_store import DraftFeatureSample, DraftReplaySample # noqa: E402 from verl_speco.trainer.target_feature_replay import ( # noqa: E402 BoundedReplayCache, + FeatureContract, TargetFeatureReplayer, _VllmEndpointState, _hidden_capture_target, _normalize_vllm_endpoints, + feature_from_vllm_payload, + load_vllm_final_norm, ) @@ -95,7 +99,9 @@ def create(self, **kwargs): model="target", ), ] - monkeypatch.setattr("verl_speco.trainer.target_feature_replay.time.sleep", lambda _: None) + monkeypatch.setattr( + "verl_speco.trainer.target_feature_replay.time.sleep", lambda _: None + ) response = replayer._request_vllm_response([1, 2, 3]) @@ -176,6 +182,7 @@ def test_vllm_payload_maps_suffix_hidden_rows_to_absolute_positions(): replayer.target_revision = None replayer.target_config_fingerprint = "unit" replayer.use_logits = False + replayer.vllm_final_norm = torch.nn.RMSNorm(4, eps=1e-6) sample = DraftReplaySample( algorithm="DSPARK", @@ -206,3 +213,198 @@ def test_vllm_payload_maps_suffix_hidden_rows_to_absolute_positions(): assert feature.metadata["feature_start"] == 5 assert feature.metadata["feature_end"] == 10 assert feature.metadata["vllm_hidden_position_offset"] == 5 + torch.testing.assert_close(feature.hidden_states[:, :8], hidden[:, :2].flatten(1)) + torch.testing.assert_close( + feature.hidden_states[:, 8:], replayer.vllm_final_norm(hidden[:, 2]) + ) + + +@pytest.mark.parametrize("model_type", ["llama", "qwen2", "qwen3", "qwen3_moe"]) +@pytest.mark.parametrize("sharded", [False, True]) +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +def test_vllm_final_norm_matches_target_forward( + tmp_path, monkeypatch, model_type, sharded, dtype +): + transformers = pytest.importorskip("transformers") + from transformers.models.auto.configuration_auto import CONFIG_MAPPING + from verl_speco import checkpoint_tensor + + if model_type not in CONFIG_MAPPING: + pytest.skip(f"Installed Transformers does not include {model_type}") + config = transformers.AutoConfig.for_model( + model_type, + hidden_size=8, + intermediate_size=16, + num_hidden_layers=2, + num_attention_heads=2, + num_key_value_heads=2, + vocab_size=32, + rms_norm_eps=1e-5, + head_dim=4, + moe_intermediate_size=8, + num_experts=2, + num_experts_per_tok=1, + ) + model = transformers.AutoModelForCausalLM.from_config(config).to(dtype).eval() + with torch.no_grad(): + model.model.norm.weight.copy_(torch.linspace(0.5, 2.0, 8)) + model.save_pretrained(tmp_path, max_shard_size="1KB" if sharded else "1GB") + loaded_keys = [] + original_load = checkpoint_tensor._load_checkpoint_tensor + + def load_one(path, key): + loaded_keys.append(key) + return original_load(path, key) + + monkeypatch.setattr(checkpoint_tensor, "_load_checkpoint_tensor", load_one) + norm = load_vllm_final_norm(str(tmp_path), dtype=dtype) + assert loaded_keys == ["model.norm.weight"] + assert all( + p.device.type == "cpu" and not p.requires_grad for p in norm.parameters() + ) + + captured = {} + handle = model.model.norm.register_forward_pre_hook( + lambda module, args: captured.update(final_input=args[0].detach().clone()) + ) + ids = torch.tensor([[1, 2, 3, 4]]) + with torch.no_grad(): + output = model(ids, output_hidden_states=True) + handle.remove() + # A connector-style payload: auxiliary layer output + final PRE-norm output. + raw = torch.stack([output.hidden_states[1][0], captured["final_input"][0]], dim=1) + original_raw = raw.clone() + request = DraftReplaySample( + input_ids=ids[0], + loss_mask=torch.ones(4), + attention_mask=torch.ones(4), + position_ids=torch.arange(4), + feature_positions=torch.arange(4), + draft_position_ids=torch.arange(1, 5), + ) + for algorithm, layout in [ + ("DSPARK", "dflash_aux_plus_last"), + ("EAGLE3", "eagle3_aux_plus_last"), + ]: + contract = FeatureContract( + algorithm=algorithm, + target_layer_ids=[0], + hidden_states_layout=layout, + dtype=dtype, + target_model_id=str(tmp_path), + target_model_revision=None, + tokenizer_fingerprint="test", + ) + payload = {"token_ids": ids[0], "hidden_states": raw} + feature = feature_from_vllm_payload(payload, request, contract, final_norm=norm) + torch.testing.assert_close( + feature.hidden_states[:, :8], raw[:, 0], rtol=0, atol=0 + ) + torch.testing.assert_close( + feature.hidden_states[:, 8:], output.hidden_states[-1][0] + ) + torch.testing.assert_close(raw, original_raw, rtol=0, atol=0) + assert not feature.hidden_states.requires_grad + # Re-reading the same raw payload must not apply norm to an already-mutated tensor. + again = feature_from_vllm_payload(payload, request, contract, final_norm=norm) + torch.testing.assert_close( + again.hidden_states, feature.hidden_states, rtol=0, atol=0 + ) + with pytest.raises(ValueError, match="require the target final norm"): + feature_from_vllm_payload(payload, request, contract) + aux_only = feature_from_vllm_payload( + payload, + request, + replace(contract, hidden_states_layout="dflash_aux", algorithm="DFLASH"), + ) + torch.testing.assert_close(aux_only.hidden_states, raw[:, 0], rtol=0, atol=0) + + +def test_final_norm_loader_requires_checkpoint_weight(tmp_path): + transformers = pytest.importorskip("transformers") + from safetensors import SafetensorError + from safetensors.torch import save_file + + config = transformers.LlamaConfig( + hidden_size=8, + intermediate_size=16, + num_hidden_layers=2, + num_attention_heads=2, + num_key_value_heads=2, + vocab_size=32, + ) + config.save_pretrained(tmp_path) + save_file({"unrelated.weight": torch.ones(8)}, str(tmp_path / "model.safetensors")) + with pytest.raises(SafetensorError, match="model.norm.weight"): + load_vllm_final_norm(str(tmp_path), dtype=torch.float32) + + +def test_vllm_replay_initializes_norm_and_invalidates_old_cache(tmp_path, monkeypatch): + from omegaconf import OmegaConf + import transformers + + transformers.LlamaConfig( + hidden_size=8, + intermediate_size=16, + num_hidden_layers=2, + num_attention_heads=2, + num_key_value_heads=2, + vocab_size=32, + ).save_pretrained(tmp_path) + calls = [] + norm = torch.nn.RMSNorm(8) + + def loader(*args, **kwargs): + calls.append(args) + return norm + + monkeypatch.setattr( + "verl_speco.trainer.target_feature_replay.load_vllm_final_norm", loader + ) + config = OmegaConf.create( + { + "actor_rollout_ref": { + "model": {"path": str(tmp_path)}, + "rollout": { + "drafter": { + "speculative_algorithm": "DSPARK", + "target_layer_ids": [0], + "training": { + "use_logits": False, + "dspark_l1_loss_alpha": 0.9, + "target_feature_replay": { + "backend": "vllm_file", + "dtype": "float32", + }, + }, + } + }, + } + } + ) + replayer = TargetFeatureReplayer( + config, rank=0, world_size=1, device=torch.device("cpu") + ) + assert calls == [(str(tmp_path),)] + assert replayer.model is None # No full target model was loaded for replay. + assert replayer.vllm_final_norm is norm + sample = DraftReplaySample( + input_ids=torch.arange(4), + loss_mask=torch.ones(4), + attention_mask=torch.ones(4), + position_ids=torch.arange(4), + feature_positions=torch.arange(4), + draft_position_ids=torch.arange(1, 5), + ) + new_key = replayer._cache_key(sample) + replayer.backend = "torch" + assert new_key != replayer._cache_key(sample) + config.actor_rollout_ref.rollout.drafter.training.target_feature_replay.backend = ( + "torch" + ) + calls.clear() + torch_replayer = TargetFeatureReplayer( + config, rank=0, world_size=1, device=torch.device("cpu") + ) + assert calls == [] + assert torch_replayer.vllm_final_norm is None diff --git a/tests/unit/test_tq_producer.py b/tests/unit/test_tq_producer.py index d6873397..b8d3543d 100644 --- a/tests/unit/test_tq_producer.py +++ b/tests/unit/test_tq_producer.py @@ -26,6 +26,21 @@ from verl_speco.standalone_tq_producer import run_producer, validate_producer_config from verl_speco.trainer.standalone_resume import save_standalone_resume from verl_speco.transport.drafter_sample_protocol import PROTOCOL_SCHEMA_VERSION +from verl_speco.transport.drafter_sample_protocol import decode_sample + + +@pytest.fixture(autouse=True) +def target_final_norm(monkeypatch): + # These pipeline tests use a fake /target checkpoint. Loader accuracy is + # covered separately with real tiny HF checkpoints. + norm = torch.nn.RMSNorm(2, eps=1e-6).requires_grad_(False) + with torch.no_grad(): + norm.weight.copy_(torch.tensor([2.0, 3.0])) + monkeypatch.setattr( + "verl_speco.standalone_tq_producer.load_vllm_final_norm", + lambda *args, **kwargs: norm, + ) + return norm def _config(input_path: Path) -> dict[str, Any]: @@ -214,7 +229,9 @@ def _write_input(path: Path) -> None: ) -def test_run_producer_publishes_samples_then_eos(tmp_path: Path) -> None: +def test_run_producer_publishes_samples_then_eos( + tmp_path: Path, target_final_norm +) -> None: input_path = tmp_path / "input.jsonl" _write_input(input_path) transport = _Transport() @@ -251,6 +268,62 @@ def test_run_producer_publishes_samples_then_eos(tmp_path: Path) -> None: assert pool.started and pool.closed and transport.closed first_fields = transport.payloads[sorted(sample_keys)[0]] assert tuple(first_fields["sample__hidden_states"].shape) == (3, 6) + # The existing response feature window starts at prompt_length - 1 = 1. + raw = torch.arange(24, dtype=torch.float32).reshape(4, 3, 2)[1:] + torch.testing.assert_close( + first_fields["sample__hidden_states"][:, :4], raw[:, :2].flatten(1) + ) + torch.testing.assert_close( + first_fields["sample__hidden_states"][:, 4:], target_final_norm(raw[:, 2]) + ) + first_key = sorted(sample_keys)[0] + sample = decode_sample( + first_key, transport.records[first_key], first_fields, {"run_id": "run-a"} + ) + torch.testing.assert_close( + sample.hidden_states[:, 4:], target_final_norm(raw[:, 2]) + ) + assert sample.metadata["last_hidden_state_norm"] == "target_final_norm" + + +@pytest.mark.parametrize("algorithm", ["DFLASH", "DSPARK"]) +def test_aux_only_producer_does_not_load_or_apply_final_norm( + tmp_path, monkeypatch, algorithm +): + def unexpected_load(*args, **kwargs): + raise AssertionError("aux-only features must not load a target final norm") + + monkeypatch.setattr( + "verl_speco.standalone_tq_producer.load_vllm_final_norm", unexpected_load + ) + input_path = tmp_path / "input.jsonl" + _write_input(input_path) + config = _config(input_path) + drafter = config["actor_rollout_ref"]["rollout"]["drafter"] + drafter["speculative_algorithm"] = algorithm + drafter["training"]["dspark_l1_loss_alpha"] = 0.0 + transport = _Transport() + asyncio.run( + run_producer( + config, + transport=transport, + tokenizer=_Tokenizer(), + client_pool=_Pool(tmp_path), + ) + ) + sample_key = next( + key + for key, tag in transport.records.items() + if tag.get("record_type") == "sample" + ) + sample = decode_sample( + sample_key, + transport.records[sample_key], + transport.payloads[sample_key], + {"run_id": "run-a"}, + ) + raw_aux = torch.arange(24, dtype=torch.float32).reshape(4, 3, 2)[1:, :2].flatten(1) + torch.testing.assert_close(sample.hidden_states, raw_aux) def test_run_producer_restarts_input_until_max_samples(tmp_path: Path) -> None: @@ -271,9 +344,7 @@ def test_run_producer_restarts_input_until_max_samples(tmp_path: Path) -> None: ) sample_tags = [ - tag - for tag in transport.records.values() - if tag.get("record_type") == "sample" + tag for tag in transport.records.values() if tag.get("record_type") == "sample" ] eos_tags = [tag for tag in transport.records.values() if tag.get("status") == "eos"] assert stats.input_count == stats.published_count == 5 diff --git a/verl_speco/checkpoint_tensor.py b/verl_speco/checkpoint_tensor.py new file mode 100644 index 00000000..83cd373e --- /dev/null +++ b/verl_speco/checkpoint_tensor.py @@ -0,0 +1,58 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Read a target checkpoint tensor without importing drafter model implementations.""" + +import glob +import json +import os + +import torch +from huggingface_hub import snapshot_download +from safetensors import safe_open + + +def _load_checkpoint_tensor(model_path: str, key: str) -> torch.Tensor: + if not os.path.exists(model_path): + model_path = snapshot_download(repo_id=model_path) + + index_paths = glob.glob(os.path.join(model_path, "*.index.json")) + if len(index_paths) > 1: + raise FileNotFoundError(f"Multiple index.json files found in {model_path}") + + if index_paths: + with open(index_paths[0], encoding="utf-8") as f: + index_json = json.load(f) + weight_map = index_json.get("weight_map", {}) + if key not in weight_map: + raise KeyError( + f"Tensor {key!r} is not present in checkpoint index for {model_path}" + ) + ckpt_file = os.path.join(model_path, weight_map[key]) + if ckpt_file.endswith(".safetensors"): + with safe_open(ckpt_file, framework="pt", device="cpu") as f: + return f.get_tensor(key) + return torch.load(ckpt_file, map_location="cpu", weights_only=True)[key] + + safetensors_path = os.path.join(model_path, "model.safetensors") + if os.path.exists(safetensors_path): + with safe_open(safetensors_path, framework="pt", device="cpu") as f: + return f.get_tensor(key) + + pytorch_path = os.path.join(model_path, "pytorch_model.bin") + if os.path.exists(pytorch_path): + return torch.load(pytorch_path, map_location="cpu", weights_only=True)[key] + + raise FileNotFoundError( + f"No index.json, model.safetensors or pytorch_model.bin found in {model_path}" + ) diff --git a/verl_speco/models/target/target_head.py b/verl_speco/models/target/target_head.py index abc0830d..990d679e 100644 --- a/verl_speco/models/target/target_head.py +++ b/verl_speco/models/target/target_head.py @@ -13,50 +13,10 @@ # limitations under the License. """Minimal target lm-head loader for SPECO drafter training.""" -import glob -import json -import os - import torch -from huggingface_hub import snapshot_download -from safetensors import safe_open from torch import nn - -def _load_checkpoint_tensor(model_path: str, key: str) -> torch.Tensor: - if not os.path.exists(model_path): - model_path = snapshot_download(repo_id=model_path) - - index_paths = glob.glob(os.path.join(model_path, "*.index.json")) - if len(index_paths) > 1: - raise FileNotFoundError(f"Multiple index.json files found in {model_path}") - - if index_paths: - with open(index_paths[0], encoding="utf-8") as f: - index_json = json.load(f) - weight_map = index_json.get("weight_map", {}) - if key not in weight_map: - raise KeyError( - f"Tensor {key!r} is not present in checkpoint index for {model_path}" - ) - ckpt_file = os.path.join(model_path, weight_map[key]) - if ckpt_file.endswith(".safetensors"): - with safe_open(ckpt_file, framework="pt", device="cpu") as f: - return f.get_tensor(key) - return torch.load(ckpt_file, map_location="cpu", weights_only=True)[key] - - safetensors_path = os.path.join(model_path, "model.safetensors") - if os.path.exists(safetensors_path): - with safe_open(safetensors_path, framework="pt", device="cpu") as f: - return f.get_tensor(key) - - pytorch_path = os.path.join(model_path, "pytorch_model.bin") - if os.path.exists(pytorch_path): - return torch.load(pytorch_path, map_location="cpu", weights_only=True)[key] - - raise FileNotFoundError( - f"No index.json, model.safetensors or pytorch_model.bin found in {model_path}" - ) +from verl_speco.checkpoint_tensor import _load_checkpoint_tensor class TargetHead(nn.Module): diff --git a/verl_speco/standalone_tq_producer.py b/verl_speco/standalone_tq_producer.py index e8df7170..0a22542f 100644 --- a/verl_speco/standalone_tq_producer.py +++ b/verl_speco/standalone_tq_producer.py @@ -46,6 +46,7 @@ FeatureContract, HiddenStateAlignmentError, feature_from_vllm_payload, + load_vllm_final_norm, ) from verl_speco.transport.drafter_sample_protocol import ( DRAFTER_TQ_PARTITION, @@ -246,6 +247,14 @@ async def run_producer( use_logits=False, require_full_alignment=True, ) + final_norm = None + if feature_contract.hidden_states_layout.endswith("_plus_last"): + final_norm = await asyncio.to_thread( + load_vllm_final_norm, + feature_contract.target_model_id, + dtype=feature_contract.dtype, + trust_remote_code=bool(producer_cfg.get("trust_remote_code", False)), + ) worker_count = int(producer_cfg["max_inflight_requests"]) input_queue: asyncio.Queue[Any] = asyncio.Queue( maxsize=int(producer_cfg["input_queue_size"]) @@ -353,7 +362,9 @@ async def request_worker() -> None: raw = await pool.prefill(request) stats.pending_bytes += int(raw.byte_size) try: - sample = feature_from_vllm_payload(raw, request, feature_contract) + sample = feature_from_vllm_payload( + raw, request, feature_contract, final_norm=final_norm + ) except HiddenStateAlignmentError as exc: stats.dropped_count += 1 stats.pending_bytes = max( diff --git a/verl_speco/trainer/target_feature_replay.py b/verl_speco/trainer/target_feature_replay.py index 43ab361e..b5d4b9a8 100644 --- a/verl_speco/trainer/target_feature_replay.py +++ b/verl_speco/trainer/target_feature_replay.py @@ -198,6 +198,46 @@ def _hidden_capture_target(layer_id: int, num_layers: int) -> tuple[str, int | N return "layer", hidden_state_index - 1 +def load_vllm_final_norm( + model_path: str, + *, + dtype: torch.dtype, + trust_remote_code: bool = False, + target_config: Any = None, +) -> nn.Module: + """Load only the target's final norm, using its actual HF implementation. + + extract_hidden_states collects layer residuals BEFORE the target final norm. + Constructing the architecture on meta discovers the norm without allocating + target weights; only that module's checkpoint tensors are read into CPU RAM. + The checkpoint must be the same frozen target served by the vLLM endpoint. + """ + from transformers import AutoConfig, AutoModelForCausalLM + + from verl_speco.checkpoint_tensor import _load_checkpoint_tensor + + if target_config is None: + target_config = AutoConfig.from_pretrained( + model_path, trust_remote_code=trust_remote_code + ) + with torch.device("meta"): + model = AutoModelForCausalLM.from_config( + target_config, + trust_remote_code=trust_remote_code, + attn_implementation="eager", + ) + _, norm = _find_layers_and_final_norm(model) + norm_name = next(name for name, module in model.named_modules() if module is norm) + state = { + key: _load_checkpoint_tensor(model_path, f"{norm_name}.{key}") + for key in norm.state_dict() + } + norm.load_state_dict(state, strict=True, assign=True) + norm = norm.to(device="cpu", dtype=dtype).eval().requires_grad_(False) + logger.info("Loaded vLLM target final norm %s from %s", norm_name, model_path) + return norm + + def _load_json_config(path: Any) -> dict[str, Any] | None: if not path: return None @@ -373,6 +413,8 @@ def feature_from_vllm_payload( payload: Mapping[str, Any] | Any, request: DraftReplaySample | Any, feature_config: FeatureContract, + *, + final_norm: nn.Module | None = None, ) -> DraftFeatureSample: """Pure vLLM payload conversion shared by replay and standalone Producer.""" @@ -469,7 +511,12 @@ def feature_from_vllm_payload( selected = hidden.index_select(0, relative_positions).to(dtype=feature_config.dtype) aux_hidden = selected[:, : len(target_layer_ids), :].flatten(1) if include_final: - final_hidden = selected[:, required_layers - 1, :] + if final_norm is None: + raise ValueError("vLLM plus_last features require the target final norm") + # Auxiliary layers stay raw. Only the final supervision block goes + # through the frozen target norm, exactly once, before storage/transport. + with torch.no_grad(): + final_hidden = final_norm(selected[:, required_layers - 1, :]) output_hidden = torch.cat([aux_hidden, final_hidden], dim=-1) else: output_hidden = aux_hidden @@ -507,6 +554,8 @@ def feature_from_vllm_payload( "use_logits": feature_config.use_logits, } ) + if include_final: + metadata["last_hidden_state_norm"] = "target_final_norm" return DraftFeatureSample( algorithm=algorithm, input_ids=selected_input_ids, @@ -641,6 +690,14 @@ def __init__( if self.algorithm in {"DFLASH", "DSPARK"} else "eagle3_aux_plus_last" ) + self.vllm_final_norm = None + if self.backend == "vllm_file" and self.hidden_layout.endswith("_plus_last"): + self.vllm_final_norm = load_vllm_final_norm( + self.model_path, + dtype=self.dtype, + trust_remote_code=self.trust_remote_code, + target_config=self.target_config, + ) config_json = json.dumps( self.target_config.to_dict(), sort_keys=True, default=str ).encode() @@ -809,6 +866,9 @@ def _cache_key(self, sample: DraftReplaySample) -> str: "use_logits": self.use_logits, "logits_topk": self.logits_topk, } + if self.backend == "vllm_file" and self.hidden_layout.endswith("_plus_last"): + # Old cache entries contain pre-norm final hidden; never reuse them. + contract["vllm_last_hidden_state_norm"] = "target_final_norm_v1" digest.update(json.dumps(contract, sort_keys=True).encode()) for tensor in ( sample.input_ids, @@ -1379,6 +1439,7 @@ def _feature_from_vllm_payload( target_config_fingerprint=self.target_config_fingerprint, source=source, ), + final_norm=self.vllm_final_norm, ) def _build_sparse_target_logprobs( @@ -1436,6 +1497,7 @@ def metrics(self) -> dict[str, float]: return metrics def close(self) -> None: + self.vllm_final_norm = None for state in self._vllm_endpoint_states: client = state.client close = getattr(client, "close", None) From b43e4cfa9fbb5f501bb0bc455010723b7f91b4f3 Mon Sep 17 00:00:00 2001 From: xyr <1615681145@qq.com> Date: Fri, 4 Sep 2026 16:43:17 +0800 Subject: [PATCH 49/50] perf(standalone_tq_producer): parallelize vllm feature conversion - offload feature_from_vllm_payload to an 8-worker ThreadPoolExecutor, bounded by an asyncio.Semaphore, with executor shutdown in cleanup - add tools/benchmark_standalone_producer.py to time the producer pipeline - add docs on external vLLM weight sync and fsdp2 weight sync smoke test --- docs/external_vllm_weight_sync.md | 91 +++++ docs/fsdp2_vllm_weight_sync_smoke.md | 168 ++++++++ tools/benchmark_standalone_producer.py | 541 +++++++++++++++++++++++++ verl_speco/standalone_tq_producer.py | 32 +- 4 files changed, 827 insertions(+), 5 deletions(-) create mode 100644 docs/external_vllm_weight_sync.md create mode 100644 docs/fsdp2_vllm_weight_sync_smoke.md create mode 100644 tools/benchmark_standalone_producer.py diff --git a/docs/external_vllm_weight_sync.md b/docs/external_vllm_weight_sync.md new file mode 100644 index 00000000..4e65e2d0 --- /dev/null +++ b/docs/external_vllm_weight_sync.md @@ -0,0 +1,91 @@ +# co-train 外部 vLLM 权重热更新 + +Last updated: 2026-09-03 + +## 范围 + +依据 `bench_weight_transfer.py` 和本机 py311 中实际安装的 vLLM 0.23.0 源码实现。控制消息通过 HTTP,权重 Tensor 由 actor rank 0 通过 NCCL 直接发送,不经过 driver、Producer Actor 或 TQ。 + +当前支持 CUDA 上的 FSDP/FSDP2/VeOmni actor 导出的完整 HF 权重,包括引擎导出的 merged 权重。不支持未合并 LoRA adapter;Megatron 等其他引擎的单 rank 完整导出语义未验证,暂时明确拒绝。当前 vLLM NCCL API 使用 CUDA,不能直接用于 NPU/HCCL。 + +## 启动配置 + +外部 hidden-state vLLM 保留原来的 hidden-state connector 等启动参数,另外增加: + +```bash +export VLLM_SERVER_DEV_MODE=1 +# 在原 vllm serve 命令上添加: +# --weight-transfer-config '{"backend":"nccl"}' +``` + +不必使用 `--load-format dummy`;如果使用,必须确认首次权重同步成功后才发送推理请求。Dev API 只应暴露在可信网络中。 + +co-train 配置示例(放在原 drafter.training 下): + +```yaml +collect_hidden_states_from_old_logprob: false +collect_hidden_states_from_sgl: false +collect_hidden_states_from_vllm: true +vllm_feature_source: + endpoints: + - http://10.0.0.10:8000/v1 + - http://10.0.0.11:8000/v1 + weight_hot_update: + enabled: true + master_address: null + timeout_seconds: 600 + bucket_size_mb: 256 + packed: true + packed_num_buffers: 2 +``` + +- `endpoints`:复用 Producer 的列表,不再单独填一份。程序去掉末尾 `/v1` 后访问 `/get_world_size` 等管理接口。 +- `master_address`:默认使用 actor rank 0 所在机器的地址;手动填写时必须是 vLLM workers 能访问到的 actor rank 0 地址,不是 driver 地址。端口在该 actor 进程上自动选择。 +- `timeout_seconds`:HTTP 等待超时,不是所有底层 NCCL 故障的统一退出期限。 +- `bucket_size_mb`:发送端分批暂存权重的目标大小,避免自己额外保存一整份完整模型。单个 Tensor 超过这个值时独占一批,packed buffer 同步扩容;不代表总显存上限,也不改变训练引擎自身导出过程的显存需求。 +- `packed` / `packed_num_buffers`:使用 vLLM 原生打包发送及缓冲数量,两端接收相同配置。 + +actor rank 0 和 Producer Actor 都必须能访问所有 endpoint。跨机器不要使用 `127.0.0.1`。外部 vLLM 应由这个训练任务独占更新,不要让其他任务同时写权重或发送请求。每个 endpoint 建立单独 NCCL 组:一个发送 rank,加该 endpoint 的 TP × PP × DP workers。 + +## 执行顺序 + +1. 初始化 actor worker 后,driver 调用所有 actor ranks 的初始化 RPC。仅 rank 0 查询 vLLM worker 数量、在后台发 HTTP 初始化请求,同时在本地调用 `NCCLWeightTransferEngine.trainer_init()` 完成握手。 +2. `fit()` 恢复 actor checkpoint 后,在原来的 rollout 权重同步边界先同步外部 vLLM。因此首次发送的是当前训练模型(含恢复的 checkpoint),不是另加载的模型文件。 +3. 当前轮 Producer 使用这份权重 prefill 当前 rollout 数据,仍按原 scheduler 分桶、提交给 drafter。 +4. actor 更新、drafter 训练后,在下一次 rollout 权重同步边界再次同步外部 vLLM,再执行原 rollout target/drafter 权重发布。没有延后一轮的逻辑。先发送外部权重,是为了尽量在同卡 rollout 恢复显存前释放 actor 导出的暂存 Tensor。 +5. 每次发送的 HTTP 顺序为 `pause(mode=wait, clear_cache=true)`、`start_weight_update`、一批或多批 `update_weights`、`finish_weight_update`、`resume`。后台 HTTP 等待接收时,主线程调用 `trainer_send_weights()`。所有 endpoint 都 finish 成功后才开始 resume。 +6. 所有 actor ranks 消费 `get_per_tensor_param()` 返回的 iterator,完成其内部的分片汇集操作;只有 rank 0 通过 NCCL 发送。需要参数 offload 的引擎在完成后恢复 CPU offload。 +7. 退出时关闭 sender 的 HTTP Session,调用本地 communicator 的 `destroy()`。不关闭用户自己启动的 vLLM 服务。 + +`pause` 清理旧权重的 KV/prefix cache,不等于永久禁用 prefix caching。原 hidden-state 服务为保证完整 hidden rows 所需的禁用 prefix cache 等配置仍须保留。 + +## 改了哪些代码 + +- `verl_speco/integration/external_vllm_weight_sync.py`:替换原占位接口,实现 HTTP/NCCL sender、分批发送、driver 和 actor 侧生命周期。 +- `verl_speco/integration/rollout_publish.py`:在已有 `DraftWeightPublishMixin` 上新增初始化、更新、关闭三个 `ONE_TO_ALL` RPC,复用已有 actor worker,不新加载目标模型。 +- `verl_speco/trainer/speco_ray_trainer.py`:连接初始化、checkpoint 恢复后的首次发送、每轮权重同步边界和最终关闭。未修改上游 verl 源码。 +- `verl_speco/config/speco_base.yaml`:补充上述配置;默认 `enabled: false`,原有模式不发起权重传输。 +- `tests/unit/test_external_vllm_weight_sync.py`:模拟控制面、sender、分批、失败处理和 driver hook 的测试。 +- `tests/special_sanity/check_device_api_usage.py`:声明这一个集成文件使用的是 CUDA 专属 vLLM API。 + +## 日志与失败处理 + +### Producer 的 final norm 同步 + +EAGLE3 和开启 L1 的 DSpark 使用辅助层加最终 hidden;vLLM connector 的最后一块是 final norm 前的 residual,Producer 在公共转换函数内补目标模型 final norm,辅助层不变。DFlash 和关闭 L1 的 DSpark 不需要这一步。 + +Producer 初始化时复用独立训练的 norm loader,只加载 final norm 的 checkpoint 权重。开启 `weight_hot_update.enabled` 后,不能一直使用初始化时的权重:Producer 报告 norm 对应的 HF 参数名;actor rank 0 在同一次完整权重导出中复制这几个小 tensor 到 CPU。发送 vLLM 成功后,现有 Ray RPC 把这些 tensor 返回 driver,再更新 Producer 的 norm。只有两边都成功才记录 `last_synced_step`,之后才进入下一轮取数。不是重新读磁盘,也没有新增一个大模型副本或额外 FSDP 汇集。 + +未开启热更新时仍使用 checkpoint 的 norm,因此 endpoint 必须服务同一份固定目标模型。此模式不能在外部自行更新 endpoint 权重而不更新 Producer。 + +握手成功:`[external vLLM weights] connected endpoint=... workers=...` + +一轮所有 endpoint 更新成功:`[external vLLM weights] updated step=... tensors=... endpoints=...` + +失败会使当前训练同步抛错,不标记成功,也不自动恢复可能只更新了部分权重的服务。后台 HTTP 异常会通过 future.result() 传回;真正 NCCL 故障仍可能等待底层超时。第一版不做失败重试或自动重连,失败后应检查日志并重启服务/任务。 + +## 验证状态 + +已完成 CPU stand-in/mock 测试与相关回归测试(104 passed,8 skipped)。跳过项为本机旧 Transformers 缺少 Qwen3 类的数值测试;包括 Producer norm 更新及失败传播测试,但这不是多卡端到端证明。 + +未启动真实 vLLM 服务。用户将自行在服务器测试实际 HTTP/NCCL 更新、训练与推理结果。 diff --git a/docs/fsdp2_vllm_weight_sync_smoke.md b/docs/fsdp2_vllm_weight_sync_smoke.md new file mode 100644 index 00000000..58e25c8e --- /dev/null +++ b/docs/fsdp2_vllm_weight_sync_smoke.md @@ -0,0 +1,168 @@ +# 四卡独立测试:FSDP2 → vLLM TP=2 权重热更新 + +Last updated: 2026-09-03 + +## 这个测试做什么 + +不启动 co-train、Ray、Producer 或 TQ,只测试实际 FSDP2 模型的权重导出、HTTP/NCCL 传输和 vLLM 加载。 + +测试会真实执行一次或几次 SGD 更新,但不保存模型、optimizer 或 checkpoint,不修改原模型目录。测试结束后内存中的更新权重随进程释放。 + +| 物理 GPU | 进程 | 用途 | +|---|---|---| +| 0、1 | 一个 vLLM 服务的两个 TP workers | 接收并加载权重,执行推理 | +| 2、3 | torchrun 的两个 FSDP2 ranks | 加载分片模型、训练、导出参数 | + +使用有空闲显存的四张 NVIDIA GPU。当前代码不是 NPU/HCCL 测试。先选普通 dense Qwen3-0.6B、Qwen3-4B 或 Llama 类本地模型;不要先用 MoE、量化、LoRA 或多模态模型。该脚本的模型需具有 `model.model.layers`。 + +## 0. 准备代码和环境 + +把当前分支的代码同步到服务器。教程假设仓库在 `/model/xyr/verl-SpeCo-ls`,模型在 `/nas/disk1/Qwen3-4B`;按实际位置修改。 + +激活服务器现有的、能运行 vLLM 0.23.0 的 CUDA Python 环境。两个终端使用同一个环境。 + +```bash +cd /model/xyr/verl-SpeCo-ls +nvidia-smi +python -c 'import torch; from importlib.metadata import version; print("torch:", torch.__version__, "CUDA:", torch.version.cuda, "vllm:", version("vllm")); assert torch.cuda.is_available(), "CUDA unavailable"' +``` + +不要为此随意升级 PyTorch。已核对的 vLLM 0.23.0 包要求 `torch==2.11.0`,实际 CUDA wheel 还需要与服务器驱动兼容。报驱动过旧、CUDA unavailable 或动态库错误时先修环境,尚未进入本测试逻辑。 + +测试脚本会自动优先导入当前仓库,不必重新 `pip install verl-speco`;运行时会打印 `SOURCE=.../verl_speco/integration/external_vllm_weight_sync.py`,确认没有使用旧安装包。 + +## 1. 终端 A:启动 vLLM,使用卡 0、1 + +确认这两个 GPU 没有被其他训练占用,并确认端口 8000 空闲。 + +```bash +cd /model/xyr/verl-SpeCo-ls +MODEL_PATH=/nas/disk1/Qwen3-4B + +CUDA_VISIBLE_DEVICES=0,1 VLLM_SERVER_DEV_MODE=1 \ +vllm serve "$MODEL_PATH" \ + --host 127.0.0.1 \ + --port 8000 \ + --served-model-name weight-sync-smoke \ + --tensor-parallel-size 2 \ + --dtype bfloat16 \ + --max-model-len 1024 \ + --gpu-memory-utilization 0.5 \ + --enforce-eager \ + --generation-config vllm \ + --no-enable-prefix-caching \ + --load-format dummy \ + --weight-transfer-config '{"backend":"nccl"}' +``` + +保持这个终端运行,等服务就绪。 + +- `VLLM_SERVER_DEV_MODE=1`:启用权重管理接口。只在可信环境使用,这里仅监听本机回环地址。 +- `TP=2`:vLLM 内部把目标模型按 Tensor Parallel 分到两张卡。 +- `--load-format dummy`:不先加载 checkpoint 权重,首次有效推理依赖发送端同步成功。不要在首次同步前用输出判断模型质量。 +- `--generation-config vllm`:避免模型目录的采样默认值影响概率对比。 +- 不需要 hidden-state connector,也不需要启动草稿模型,这一步隔离验证目标模型权重传输。 + +本机四卡测试使用 `127.0.0.1`。如果两端不在同一台服务器,需要另外配置监听地址、网络和 actor rank 0 的可达地址,本教程不覆盖跨机部署。 + +## 2. 终端 B:先检查管理接口 + +```bash +curl --noproxy '*' -fsS http://127.0.0.1:8000/get_world_size +``` + +预期: + +```json +{"world_size":2} +``` + +返回 404 通常表示没有启用 dev API,或者访问了错误的 vLLM 版本/端口。连接拒绝说明服务尚未起来。不是 2 则先核对 vLLM 的 TP/PP/DP 配置。 + +这个接口只是读取服务配置,不会执行推理或修改权重。 + +## 3. 终端 B:用卡 2、3 跑 FSDP2 测试 + +```bash +cd /model/xyr/verl-SpeCo-ls +MODEL_PATH=/nas/disk1/Qwen3-4B + +CUDA_VISIBLE_DEVICES=2,3 \ +python -m torch.distributed.run \ + --standalone \ + --nproc_per_node=2 \ + tools/fsdp2_vllm_weight_sync_smoke.py \ + --model "$MODEL_PATH" \ + --endpoint http://127.0.0.1:8000/v1 \ + --served-model-name weight-sync-smoke \ + --bucket-size-mb 256 \ + --train-steps 1 \ + --lr 0.01 \ + --logprob-atol 0.1 +``` + +不需要 `ray start`。 + +参数含义: + +- `--model`:必须与 vLLM 是同一架构、同一份模型;脚本只从本地目录加载,不下载。 +- `--endpoint`:带 `/v1` 的推理地址。生产代码会自动去掉 `/v1` 来访问管理接口。 +- `--served-model-name`:与终端 A 一致,是 HTTP 请求里的模型名,不是文件路径。 +- `--bucket-size-mb`:沿用实际发送代码的分批目标大小;超大单 Tensor 会独占一批并扩展 packed buffer。 +- `--train-steps`:首次对比后真实训练几步,默认 1。 +- `--lr`:本测试的临时 SGD 学习率,不是正式训练的推荐学习率。 +- `--logprob-atol`:不同推理实现下,同一 token 的 logprob 允许的绝对误差。 +- `--no-packed`:可选,用逐 Tensor 广播替代默认 packed。不是绕过 NCCL。 +- `--timeout`:HTTP 超时秒数,默认 600;NCCL 的底层故障可能仍需要等待其自身超时。 + +模型首先在每个进程的 CPU 内存中加载完整 checkpoint,再按 transformer layer 进行 FSDP2 分片并放到设备。因此 CPU 内存也要足够,不适合直接从超大模型开始。 + +## 4. 应看到哪些日志 + +```text +SOURCE=.../verl_speco/integration/external_vllm_weight_sync.py +FSDP_READY rank=0 global_shape=(...) local_shape=(...) +FSDP_READY rank=1 global_shape=(...) local_shape=(...) +CONNECT_OK +INITIAL_SYNC_OK +COMPARE prompt=0 hf_top1=... vllm_top1=... max_logprob_error=... +... +INITIAL_COMPARE_OK +TRAIN_STEP_OK step=1 loss=... +UPDATED_SYNC_OK +COMPARE prompt=0 hf_top1=... vllm_top1=... max_logprob_error=... +... +CHECKED_TOKEN_REFERENCE_CHANGE=... +UPDATED_COMPARE_OK +PASS: initial and trained FSDP2 weights match external vLLM; no checkpoint written +``` + +`FSDP_READY` 会同时打印全局参数形状与本 rank 的分片形状,确认不是两个完整的普通模型。 + +脚本用三个固定 prompt 的 token IDs,分别对比 FSDP2 模型和 vLLM 的下一 token top-5 logprob。不会用“生成的文字是否一样”作为唯一依据,也不会只凭 HTTP 200 判成功。浮点差异可能让非常接近的 top-1 token 换位,脚本同时检查其概率差距。 + +第二轮还检查被比较 token 的参考概率是否相对训练前充分变化。如果变化不足以排除旧权重,打印 `INCONCLUSIVE` 并退出失败,而不声称热更新已验证。遇到这种情况:在终端 A 停止并重新启动测试服务,再把终端 B 改为 `--train-steps 3 --lr 0.05` 重试。不建议直接放宽误差掩盖问题。 + +## 5. 结束与排查 + +发送端正常完成后会自行关闭 NCCL sender、HTTP Session 和 FSDP2 进程组,不保存文件。vLLM 是你单独启动的,因此到终端 A 按 Ctrl+C 关闭。 + +发送失败后不要继续使用这个测试 vLLM:部分权重可能已更新,服务可能保持暂停。先停止发送端,再停止并重新启动终端 A 的测试服务,然后重试。不要使用 `pkill -f ray` 或其他会误杀无关任务的命令。 + +常见定位: + +- 没到 `CONNECT_OK`:检查 dev API、两端网络可达性和 NCCL 握手。传输 rendezvous 默认使用发送 rank 0 的自动检测 IP,可用 `--master-address <该机器可达IP>` 指定。 +- 卡在 `INITIAL_SYNC`:结合发送端和 vLLM 日志,检查 `update_weights` 错误、NCCL 错误及显存。发送端两张卡都必须参与 FSDP2 导出。 +- 首次 `COMPARE` 超阈值:核对模型、dtype、是否量化、参数名、两端包版本;先不要提高阈值。 +- 初次正确、训练后不正确:重点检查后续权重更新和旧缓存。发送代码在 pause 时清缓存,本教程还禁用了 prefix caching。 +- OOM:先换小模型,或减少 bucket 大小。普通模型加载、FSDP2 导出、vLLM TP 权重和 packed 接收缓冲都有显存需求,`bucket_size_mb` 不是总显存上限。 + +## 6. 与正式 co-train 的关系 + +这份脚本复用的是现有 `initialize_worker_weight_sync()`、`update_worker_weights()`、`close_worker_weight_sync()` 和 `ExternalVllmWeightSender`,不是另写一套 NCCL。 + +只有测试引擎适配部分不同:脚本用原生 FSDP2 的 `state_dict()` / `DTensor.full_tensor()`;正式 co-train 使用实际 verl engine 的 `get_per_tensor_param()`。因此测试通过能验证真实分片汇集和传输,但不能证明 MoE 参数转换、LoRA、offload、Ray 调度或整个 co-train 都已经通过。 + +两个 FSDP ranks 都参与汇集;只有 FSDP rank 0 发给两个 vLLM TP workers。FSDP 训练通信组为 2 个进程,外部权重传输组为 3 个进程(sender + 两个 TP workers),不是 4。 + +本机未运行此 GPU/vLLM 测试,脚本和步骤供服务器手动验证。 diff --git a/tools/benchmark_standalone_producer.py b/tools/benchmark_standalone_producer.py new file mode 100644 index 00000000..7d69e604 --- /dev/null +++ b/tools/benchmark_standalone_producer.py @@ -0,0 +1,541 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Benchmark the standalone Producer path without starting Ray or TQ. + +The benchmark deliberately reuses the production input reader, asynchronous +vLLM client pool, hidden-state alignment, and feature conversion. In ``both`` +mode, each vLLM result is converted twice: once with the real target final norm +and once with ``torch.nn.Identity``. This keeps the HTTP response and hidden +tensor identical, so the conversion-time difference isolates final-norm cost. +""" + +from __future__ import annotations + +import argparse +import asyncio +import json +import logging +import math +import statistics +import sys +import time +from collections import defaultdict +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +import torch +from torch import nn + + +REPO_ROOT = Path(__file__).resolve().parents[1] +if str(REPO_ROOT) not in sys.path: + sys.path.insert(0, str(REPO_ROOT)) + +from verl_speco.integration.oldlogprob_layer_ids import ( # noqa: E402 + resolve_drafter_hidden_states_layout, +) +from verl_speco.producer.input_reader import ( # noqa: E402 + GenerationRequest, + TokenizedRequest, + iter_input_records, + prepare_generated_prefill_request, + prepare_generation_request, + tokenize_record, +) +from verl_speco.producer.vllm_feature_client import ( # noqa: E402 + VllmEndpoint, + VllmFeatureClientPool, + delete_temporary_result, +) +from verl_speco.trainer.target_feature_replay import ( # noqa: E402 + FeatureContract, + HiddenStateAlignmentError, + feature_from_vllm_payload, + load_vllm_final_norm, +) + + +logger = logging.getLogger("producer_benchmark") +_INPUT_DONE = object() + + +@dataclass(frozen=True) +class QueuedRequest: + request: GenerationRequest | TokenizedRequest + queued_at: float + + +class Timings: + def __init__(self) -> None: + self.values: dict[str, list[float]] = defaultdict(list) + self.completed = 0 + self.dropped = 0 + self.generated = 0 + self.hidden_bytes = 0 + self.feature_tokens = 0 + self._lock = asyncio.Lock() + + async def add(self, **values: float) -> None: + async with self._lock: + for name, value in values.items(): + self.values[name].append(float(value)) + + async def complete( + self, *, generated: bool, hidden_bytes: int, feature_tokens: int + ) -> int: + async with self._lock: + self.completed += 1 + self.generated += int(generated) + self.hidden_bytes += int(hidden_bytes) + self.feature_tokens += int(feature_tokens) + return self.completed + + async def drop(self) -> None: + async with self._lock: + self.dropped += 1 + + +def _parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Time the production standalone vLLM Producer path without TQ." + ) + parser.add_argument("--input-path", required=True) + parser.add_argument("--tokenizer-path", required=True) + parser.add_argument( + "--target-model-path", + required=True, + help="HF checkpoint used to load only the target final-norm parameters.", + ) + parser.add_argument( + "--vllm-model", + default=None, + help="Model name sent to vLLM; defaults to --target-model-path.", + ) + parser.add_argument( + "--endpoints", + nargs="+", + default=["http://127.0.0.1:8000/v1"], + help="One or more OpenAI-compatible vLLM base URLs.", + ) + parser.add_argument("--algorithm", default="DSPARK") + parser.add_argument( + "--target-layer-ids", + required=True, + help="Comma-separated auxiliary layer IDs, excluding the final layer.", + ) + parser.add_argument( + "--norm-mode", + choices=("both", "norm", "no-norm"), + default="both", + help="both reuses each response and converts it through both paths.", + ) + parser.add_argument( + "--hidden-dtype", choices=("bf16", "fp16", "fp32"), default="bf16" + ) + parser.add_argument("--dspark-l1-loss-alpha", type=float, default=0.9) + parser.add_argument("--max-samples", type=int, default=100) + parser.add_argument("--max-inflight-requests", type=int, default=64) + parser.add_argument("--per-endpoint-concurrency", type=int, default=64) + parser.add_argument("--input-queue-size", type=int, default=128) + parser.add_argument("--request-timeout", type=float, default=600.0) + parser.add_argument("--generation-max-tokens", type=int, default=512) + parser.add_argument("--max-sequence-length", type=int, default=8192) + parser.add_argument("--max-feature-length", type=int, default=512) + parser.add_argument("--progress-interval", type=int, default=10) + parser.add_argument("--trust-remote-code", action="store_true") + parser.add_argument("--json-output", default=None) + return parser.parse_args() + + +def _dtype(name: str) -> torch.dtype: + return { + "bf16": torch.bfloat16, + "fp16": torch.float16, + "fp32": torch.float32, + }[name] + + +def _producer_config(args: argparse.Namespace) -> dict[str, Any]: + return { + "max_sequence_length": args.max_sequence_length, + "max_feature_length": args.max_feature_length, + "generation_max_tokens": args.generation_max_tokens, + } + + +def _load_tokenizer(args: argparse.Namespace) -> Any: + from transformers import AutoTokenizer + + return AutoTokenizer.from_pretrained( + args.tokenizer_path, + trust_remote_code=args.trust_remote_code, + ) + + +def _percentile(sorted_values: list[float], percentile: float) -> float: + if not sorted_values: + return 0.0 + index = max(math.ceil(percentile * len(sorted_values)) - 1, 0) + return sorted_values[min(index, len(sorted_values) - 1)] + + +def _summary(values: list[float]) -> dict[str, float | int]: + ordered = sorted(values) + return { + "count": len(ordered), + "total_s": sum(ordered), + "mean_ms": statistics.fmean(ordered) * 1000 if ordered else 0.0, + "p50_ms": _percentile(ordered, 0.50) * 1000, + "p95_ms": _percentile(ordered, 0.95) * 1000, + "max_ms": ordered[-1] * 1000 if ordered else 0.0, + } + + +def _optional_ms(value: float | None) -> str: + return "n/a" if value is None else f"{value * 1000:.3f}" + + +def _print_report(report: dict[str, Any]) -> None: + totals = report["totals"] + print("\nProducer benchmark result") + print( + f"completed={totals['completed']} dropped={totals['dropped']} " + f"generated={totals['generated']} wall={totals['wall_seconds']:.3f}s " + f"samples/s={totals['samples_per_second']:.3f} " + f"feature_tokens/s={totals['feature_tokens_per_second']:.1f}" + ) + print( + "stage".ljust(30) + + "count".rjust(8) + + "mean_ms".rjust(12) + + "p50_ms".rjust(12) + + "p95_ms".rjust(12) + + "max_ms".rjust(12) + + "sum_s".rjust(12) + ) + for name, item in report["stages"].items(): + print( + name.ljust(30) + + str(item["count"]).rjust(8) + + f"{item['mean_ms']:.3f}".rjust(12) + + f"{item['p50_ms']:.3f}".rjust(12) + + f"{item['p95_ms']:.3f}".rjust(12) + + f"{item['max_ms']:.3f}".rjust(12) + + f"{item['total_s']:.3f}".rjust(12) + ) + comparison = report.get("norm_comparison") + if comparison: + print( + "\nconversion mean delta (norm - no_norm): " + f"{comparison['mean_delta_ms']:.3f} ms/sample; " + f"ratio={comparison['mean_ratio']:.3f}x" + ) + + +async def _run(args: argparse.Namespace) -> dict[str, Any]: + if args.max_samples <= 0: + raise ValueError("--max-samples must be positive") + if args.max_inflight_requests <= 0 or args.per_endpoint_concurrency <= 0: + raise ValueError("request concurrency values must be positive") + target_layer_ids = [ + int(value.strip()) + for value in args.target_layer_ids.split(",") + if value.strip() + ] + if not target_layer_ids: + raise ValueError("--target-layer-ids must contain at least one layer") + + producer_config = _producer_config(args) + algorithm = args.algorithm.strip().upper() + training_config = { + "speculative_algorithm": algorithm, + "dspark_l1_loss_alpha": args.dspark_l1_loss_alpha, + } + layout = resolve_drafter_hidden_states_layout(algorithm, training_config) + include_final = layout.endswith("_plus_last") + if args.norm_mode in {"both", "norm"} and not include_final: + raise ValueError( + f"algorithm/config resolves to {layout!r}, which does not consume a final " + "hidden block; norm comparison is not part of that production path" + ) + feature_contract = FeatureContract( + algorithm=algorithm, + target_layer_ids=target_layer_ids, + hidden_states_layout=layout, + dtype=_dtype(args.hidden_dtype), + target_model_id=args.target_model_path, + target_model_revision=None, + tokenizer_fingerprint="producer-benchmark", + use_logits=False, + source="producer_benchmark", + require_full_alignment=True, + ) + + timings = Timings() + startup_begin = time.perf_counter() + tokenizer_begin = time.perf_counter() + tokenizer = await asyncio.to_thread(_load_tokenizer, args) + await timings.add(tokenizer_load=time.perf_counter() - tokenizer_begin) + + real_norm: nn.Module | None = None + if args.norm_mode in {"both", "norm"}: + norm_begin = time.perf_counter() + real_norm = await asyncio.to_thread( + load_vllm_final_norm, + args.target_model_path, + dtype=feature_contract.dtype, + trust_remote_code=args.trust_remote_code, + ) + await timings.add(final_norm_load=time.perf_counter() - norm_begin) + identity_norm = nn.Identity() + + pool = VllmFeatureClientPool( + [ + VllmEndpoint(url.rstrip("/"), args.per_endpoint_concurrency) + for url in args.endpoints + ], + model=args.vllm_model or args.target_model_path, + max_inflight_requests=args.max_inflight_requests, + request_timeout=args.request_timeout, + ) + pool_begin = time.perf_counter() + await pool.start() + await timings.add(client_pool_start=time.perf_counter() - pool_begin) + queue: asyncio.Queue[QueuedRequest | object] = asyncio.Queue( + maxsize=args.input_queue_size + ) + + async def read_inputs() -> None: + count = 0 + for record in iter_input_records(args.input_path): + if count >= args.max_samples: + break + begin = time.perf_counter() + request = ( + prepare_generation_request(record, tokenizer, producer_config) + if record.response is None + else tokenize_record(record, tokenizer, producer_config) + ) + preparation = time.perf_counter() - begin + await timings.add(input_prepare=preparation) + await queue.put( + QueuedRequest(request=request, queued_at=time.perf_counter()) + ) + count += 1 + if count == 0: + raise ValueError("input contains no usable samples") + for _ in range(args.max_inflight_requests): + await queue.put(_INPUT_DONE) + + async def request_worker() -> None: + while True: + queued = await queue.get() + if queued is _INPUT_DONE: + return + assert isinstance(queued, QueuedRequest) + request = queued.request + worker_begin = time.perf_counter() + await timings.add(input_queue_wait=worker_begin - queued.queued_at) + raw = None + generated_raw = None + was_generated = isinstance(request, GenerationRequest) + try: + if was_generated: + # Re-assert to narrow the type for mypy: the bool above is + # not tracked as a type guard. + assert isinstance(request, GenerationRequest) + begin = time.perf_counter() + generated_raw = await pool.generate(request) + await timings.add( + vllm_generate_and_load=time.perf_counter() - begin + ) + + begin = time.perf_counter() + request = prepare_generated_prefill_request( + request, + generated_raw.generated_token_ids, + producer_config, + ) + await timings.add( + generated_request_prepare=time.perf_counter() - begin + ) + + begin = time.perf_counter() + await asyncio.to_thread(delete_temporary_result, generated_raw) + await timings.add( + generation_file_cleanup=time.perf_counter() - begin + ) + generated_raw = None + + begin = time.perf_counter() + raw = await pool.prefill(request) + prefill_seconds = time.perf_counter() - begin + await timings.add(vllm_prefill_and_load=prefill_seconds) + assert isinstance(request, TokenizedRequest) + endpoint_url = raw.endpoint_url + + feature_tokens = 0 + conversion_seconds: dict[str, float] = {} + + async def convert(name: str, module: nn.Module | None) -> None: + nonlocal feature_tokens + begin = time.perf_counter() + sample = feature_from_vllm_payload( + raw, + request, + feature_contract, + final_norm=module, + ) + elapsed = time.perf_counter() - begin + conversion_seconds[name] = elapsed + await timings.add(**{name: elapsed}) + feature_tokens = int(sample.input_ids.numel()) + del sample + + # Alternate order in comparison mode so CPU cache warmth does + # not systematically favor either conversion path. + conversion_plan: list[tuple[str, nn.Module | None]] = [] + if args.norm_mode == "both": + conversion_plan = [ + ("convert_no_norm", identity_norm), + ("convert_with_norm", real_norm), + ] + if request.sequence_no % 2: + conversion_plan.reverse() + elif args.norm_mode == "no-norm": + conversion_plan = [ + ( + "convert_no_norm", + identity_norm if include_final else None, + ) + ] + else: + conversion_plan = [("convert_with_norm", real_norm)] + for conversion_name, norm_module in conversion_plan: + await convert(conversion_name, norm_module) + + begin = time.perf_counter() + await asyncio.to_thread(delete_temporary_result, raw) + await timings.add(prefill_file_cleanup=time.perf_counter() - begin) + hidden_bytes = int(raw.byte_size) + raw = None + await timings.add( + worker_total=time.perf_counter() - worker_begin, + queued_to_complete=time.perf_counter() - queued.queued_at, + ) + completed = await timings.complete( + generated=was_generated, + hidden_bytes=hidden_bytes, + feature_tokens=feature_tokens, + ) + if completed <= 3 or completed % args.progress_interval == 0: + logger.info( + "completed=%s/%s sample_id=%s generated=%s endpoint=%s " + "hidden_mib=%.2f", + completed, + args.max_samples, + request.sample_id, + was_generated, + endpoint_url, + hidden_bytes / 1024**2, + ) + logger.info( + "sample timing sample_id=%s prefill_and_load_ms=%.3f " + "convert_no_norm_ms=%s convert_with_norm_ms=%s", + request.sample_id, + prefill_seconds * 1000, + _optional_ms(conversion_seconds.get("convert_no_norm")), + _optional_ms(conversion_seconds.get("convert_with_norm")), + ) + except HiddenStateAlignmentError as exc: + await timings.drop() + logger.warning( + "dropped sample_id=%s because hidden states do not align: %s", + request.sample_id, + exc, + ) + finally: + if generated_raw is not None: + await asyncio.to_thread(delete_temporary_result, generated_raw) + if raw is not None: + await asyncio.to_thread(delete_temporary_result, raw) + + benchmark_begin = time.perf_counter() + try: + await asyncio.gather( + read_inputs(), + *(request_worker() for _ in range(args.max_inflight_requests)), + ) + finally: + await pool.close() + wall_seconds = time.perf_counter() - benchmark_begin + startup_seconds = benchmark_begin - startup_begin + stages = {name: _summary(values) for name, values in timings.values.items()} + report: dict[str, Any] = { + "configuration": { + "input_path": args.input_path, + "endpoints": args.endpoints, + "algorithm": algorithm, + "hidden_states_layout": layout, + "target_layer_ids": target_layer_ids, + "norm_mode": args.norm_mode, + "max_samples": args.max_samples, + "max_inflight_requests": args.max_inflight_requests, + "per_endpoint_concurrency": args.per_endpoint_concurrency, + "max_sequence_length": args.max_sequence_length, + "max_feature_length": args.max_feature_length, + }, + "totals": { + "completed": timings.completed, + "dropped": timings.dropped, + "generated": timings.generated, + "hidden_bytes": timings.hidden_bytes, + "feature_tokens": timings.feature_tokens, + "startup_seconds": startup_seconds, + "wall_seconds": wall_seconds, + "samples_per_second": timings.completed / wall_seconds, + "feature_tokens_per_second": timings.feature_tokens / wall_seconds, + }, + "stages": stages, + } + if "convert_with_norm" in stages and "convert_no_norm" in stages: + norm_mean = float(stages["convert_with_norm"]["mean_ms"]) + no_norm_mean = float(stages["convert_no_norm"]["mean_ms"]) + report["norm_comparison"] = { + "mean_delta_ms": norm_mean - no_norm_mean, + "mean_ratio": norm_mean / no_norm_mean if no_norm_mean else math.inf, + } + return report + + +def main() -> None: + args = _parse_args() + logging.basicConfig( + level=logging.INFO, + format="%(asctime)s %(levelname)s %(name)s: %(message)s", + ) + report = asyncio.run(_run(args)) + _print_report(report) + if args.json_output: + output = Path(args.json_output) + output.parent.mkdir(parents=True, exist_ok=True) + output.write_text( + json.dumps(report, ensure_ascii=False, indent=2), encoding="utf-8" + ) + print(f"JSON report written to {output}") + + +if __name__ == "__main__": + main() diff --git a/verl_speco/standalone_tq_producer.py b/verl_speco/standalone_tq_producer.py index 0a22542f..c241fe9a 100644 --- a/verl_speco/standalone_tq_producer.py +++ b/verl_speco/standalone_tq_producer.py @@ -17,7 +17,9 @@ import asyncio import logging +from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass, replace +from functools import partial from typing import Any, Mapping import torch @@ -63,6 +65,7 @@ logger = logging.getLogger(__name__) _INPUT_DONE = object() _PUBLISH_DONE = object() +_FEATURE_CONVERSION_WORKERS = 8 @dataclass @@ -170,6 +173,7 @@ async def run_producer( stats = ProducerStats() connected = False pool = client_pool + feature_executor: ThreadPoolExecutor | None = None try: logger.info( "Standalone TQ Producer starting run_id=%s input=%s endpoints=%s", @@ -255,6 +259,12 @@ async def run_producer( dtype=feature_contract.dtype, trust_remote_code=bool(producer_cfg.get("trust_remote_code", False)), ) + feature_executor = ThreadPoolExecutor( + max_workers=_FEATURE_CONVERSION_WORKERS, + thread_name_prefix="speco-feature", + ) + feature_slots = asyncio.Semaphore(_FEATURE_CONVERSION_WORKERS) + event_loop = asyncio.get_running_loop() worker_count = int(producer_cfg["max_inflight_requests"]) input_queue: asyncio.Queue[Any] = asyncio.Queue( maxsize=int(producer_cfg["input_queue_size"]) @@ -362,9 +372,17 @@ async def request_worker() -> None: raw = await pool.prefill(request) stats.pending_bytes += int(raw.byte_size) try: - sample = feature_from_vllm_payload( - raw, request, feature_contract, final_norm=final_norm - ) + async with feature_slots: + sample = await event_loop.run_in_executor( + feature_executor, + partial( + feature_from_vllm_payload, + raw, + request, + feature_contract, + final_norm=final_norm, + ), + ) except HiddenStateAlignmentError as exc: stats.dropped_count += 1 stats.pending_bytes = max( @@ -438,8 +456,12 @@ async def publish_results() -> None: return stats finally: try: - if pool is not None: - await pool.close() + try: + if pool is not None: + await pool.close() + finally: + if feature_executor is not None: + feature_executor.shutdown(wait=True) finally: if connected: transport.close_transfer_queue_client() From 18d8a2d6075d2112ccc186d839a41d21d821655e Mon Sep 17 00:00:00 2001 From: xyr <1615681145@qq.com> Date: Fri, 4 Sep 2026 16:46:38 +0800 Subject: [PATCH 50/50] chore: stop tracking personal weight-sync notes Keep the two personal docs (external_vllm_weight_sync / fsdp2 weight sync smoke) as local-only files excluded via .git/info/exclude, consistent with the other personal planning docs. --- docs/external_vllm_weight_sync.md | 91 --------------- docs/fsdp2_vllm_weight_sync_smoke.md | 168 --------------------------- 2 files changed, 259 deletions(-) delete mode 100644 docs/external_vllm_weight_sync.md delete mode 100644 docs/fsdp2_vllm_weight_sync_smoke.md diff --git a/docs/external_vllm_weight_sync.md b/docs/external_vllm_weight_sync.md deleted file mode 100644 index 4e65e2d0..00000000 --- a/docs/external_vllm_weight_sync.md +++ /dev/null @@ -1,91 +0,0 @@ -# co-train 外部 vLLM 权重热更新 - -Last updated: 2026-09-03 - -## 范围 - -依据 `bench_weight_transfer.py` 和本机 py311 中实际安装的 vLLM 0.23.0 源码实现。控制消息通过 HTTP,权重 Tensor 由 actor rank 0 通过 NCCL 直接发送,不经过 driver、Producer Actor 或 TQ。 - -当前支持 CUDA 上的 FSDP/FSDP2/VeOmni actor 导出的完整 HF 权重,包括引擎导出的 merged 权重。不支持未合并 LoRA adapter;Megatron 等其他引擎的单 rank 完整导出语义未验证,暂时明确拒绝。当前 vLLM NCCL API 使用 CUDA,不能直接用于 NPU/HCCL。 - -## 启动配置 - -外部 hidden-state vLLM 保留原来的 hidden-state connector 等启动参数,另外增加: - -```bash -export VLLM_SERVER_DEV_MODE=1 -# 在原 vllm serve 命令上添加: -# --weight-transfer-config '{"backend":"nccl"}' -``` - -不必使用 `--load-format dummy`;如果使用,必须确认首次权重同步成功后才发送推理请求。Dev API 只应暴露在可信网络中。 - -co-train 配置示例(放在原 drafter.training 下): - -```yaml -collect_hidden_states_from_old_logprob: false -collect_hidden_states_from_sgl: false -collect_hidden_states_from_vllm: true -vllm_feature_source: - endpoints: - - http://10.0.0.10:8000/v1 - - http://10.0.0.11:8000/v1 - weight_hot_update: - enabled: true - master_address: null - timeout_seconds: 600 - bucket_size_mb: 256 - packed: true - packed_num_buffers: 2 -``` - -- `endpoints`:复用 Producer 的列表,不再单独填一份。程序去掉末尾 `/v1` 后访问 `/get_world_size` 等管理接口。 -- `master_address`:默认使用 actor rank 0 所在机器的地址;手动填写时必须是 vLLM workers 能访问到的 actor rank 0 地址,不是 driver 地址。端口在该 actor 进程上自动选择。 -- `timeout_seconds`:HTTP 等待超时,不是所有底层 NCCL 故障的统一退出期限。 -- `bucket_size_mb`:发送端分批暂存权重的目标大小,避免自己额外保存一整份完整模型。单个 Tensor 超过这个值时独占一批,packed buffer 同步扩容;不代表总显存上限,也不改变训练引擎自身导出过程的显存需求。 -- `packed` / `packed_num_buffers`:使用 vLLM 原生打包发送及缓冲数量,两端接收相同配置。 - -actor rank 0 和 Producer Actor 都必须能访问所有 endpoint。跨机器不要使用 `127.0.0.1`。外部 vLLM 应由这个训练任务独占更新,不要让其他任务同时写权重或发送请求。每个 endpoint 建立单独 NCCL 组:一个发送 rank,加该 endpoint 的 TP × PP × DP workers。 - -## 执行顺序 - -1. 初始化 actor worker 后,driver 调用所有 actor ranks 的初始化 RPC。仅 rank 0 查询 vLLM worker 数量、在后台发 HTTP 初始化请求,同时在本地调用 `NCCLWeightTransferEngine.trainer_init()` 完成握手。 -2. `fit()` 恢复 actor checkpoint 后,在原来的 rollout 权重同步边界先同步外部 vLLM。因此首次发送的是当前训练模型(含恢复的 checkpoint),不是另加载的模型文件。 -3. 当前轮 Producer 使用这份权重 prefill 当前 rollout 数据,仍按原 scheduler 分桶、提交给 drafter。 -4. actor 更新、drafter 训练后,在下一次 rollout 权重同步边界再次同步外部 vLLM,再执行原 rollout target/drafter 权重发布。没有延后一轮的逻辑。先发送外部权重,是为了尽量在同卡 rollout 恢复显存前释放 actor 导出的暂存 Tensor。 -5. 每次发送的 HTTP 顺序为 `pause(mode=wait, clear_cache=true)`、`start_weight_update`、一批或多批 `update_weights`、`finish_weight_update`、`resume`。后台 HTTP 等待接收时,主线程调用 `trainer_send_weights()`。所有 endpoint 都 finish 成功后才开始 resume。 -6. 所有 actor ranks 消费 `get_per_tensor_param()` 返回的 iterator,完成其内部的分片汇集操作;只有 rank 0 通过 NCCL 发送。需要参数 offload 的引擎在完成后恢复 CPU offload。 -7. 退出时关闭 sender 的 HTTP Session,调用本地 communicator 的 `destroy()`。不关闭用户自己启动的 vLLM 服务。 - -`pause` 清理旧权重的 KV/prefix cache,不等于永久禁用 prefix caching。原 hidden-state 服务为保证完整 hidden rows 所需的禁用 prefix cache 等配置仍须保留。 - -## 改了哪些代码 - -- `verl_speco/integration/external_vllm_weight_sync.py`:替换原占位接口,实现 HTTP/NCCL sender、分批发送、driver 和 actor 侧生命周期。 -- `verl_speco/integration/rollout_publish.py`:在已有 `DraftWeightPublishMixin` 上新增初始化、更新、关闭三个 `ONE_TO_ALL` RPC,复用已有 actor worker,不新加载目标模型。 -- `verl_speco/trainer/speco_ray_trainer.py`:连接初始化、checkpoint 恢复后的首次发送、每轮权重同步边界和最终关闭。未修改上游 verl 源码。 -- `verl_speco/config/speco_base.yaml`:补充上述配置;默认 `enabled: false`,原有模式不发起权重传输。 -- `tests/unit/test_external_vllm_weight_sync.py`:模拟控制面、sender、分批、失败处理和 driver hook 的测试。 -- `tests/special_sanity/check_device_api_usage.py`:声明这一个集成文件使用的是 CUDA 专属 vLLM API。 - -## 日志与失败处理 - -### Producer 的 final norm 同步 - -EAGLE3 和开启 L1 的 DSpark 使用辅助层加最终 hidden;vLLM connector 的最后一块是 final norm 前的 residual,Producer 在公共转换函数内补目标模型 final norm,辅助层不变。DFlash 和关闭 L1 的 DSpark 不需要这一步。 - -Producer 初始化时复用独立训练的 norm loader,只加载 final norm 的 checkpoint 权重。开启 `weight_hot_update.enabled` 后,不能一直使用初始化时的权重:Producer 报告 norm 对应的 HF 参数名;actor rank 0 在同一次完整权重导出中复制这几个小 tensor 到 CPU。发送 vLLM 成功后,现有 Ray RPC 把这些 tensor 返回 driver,再更新 Producer 的 norm。只有两边都成功才记录 `last_synced_step`,之后才进入下一轮取数。不是重新读磁盘,也没有新增一个大模型副本或额外 FSDP 汇集。 - -未开启热更新时仍使用 checkpoint 的 norm,因此 endpoint 必须服务同一份固定目标模型。此模式不能在外部自行更新 endpoint 权重而不更新 Producer。 - -握手成功:`[external vLLM weights] connected endpoint=... workers=...` - -一轮所有 endpoint 更新成功:`[external vLLM weights] updated step=... tensors=... endpoints=...` - -失败会使当前训练同步抛错,不标记成功,也不自动恢复可能只更新了部分权重的服务。后台 HTTP 异常会通过 future.result() 传回;真正 NCCL 故障仍可能等待底层超时。第一版不做失败重试或自动重连,失败后应检查日志并重启服务/任务。 - -## 验证状态 - -已完成 CPU stand-in/mock 测试与相关回归测试(104 passed,8 skipped)。跳过项为本机旧 Transformers 缺少 Qwen3 类的数值测试;包括 Producer norm 更新及失败传播测试,但这不是多卡端到端证明。 - -未启动真实 vLLM 服务。用户将自行在服务器测试实际 HTTP/NCCL 更新、训练与推理结果。 diff --git a/docs/fsdp2_vllm_weight_sync_smoke.md b/docs/fsdp2_vllm_weight_sync_smoke.md deleted file mode 100644 index 58e25c8e..00000000 --- a/docs/fsdp2_vllm_weight_sync_smoke.md +++ /dev/null @@ -1,168 +0,0 @@ -# 四卡独立测试:FSDP2 → vLLM TP=2 权重热更新 - -Last updated: 2026-09-03 - -## 这个测试做什么 - -不启动 co-train、Ray、Producer 或 TQ,只测试实际 FSDP2 模型的权重导出、HTTP/NCCL 传输和 vLLM 加载。 - -测试会真实执行一次或几次 SGD 更新,但不保存模型、optimizer 或 checkpoint,不修改原模型目录。测试结束后内存中的更新权重随进程释放。 - -| 物理 GPU | 进程 | 用途 | -|---|---|---| -| 0、1 | 一个 vLLM 服务的两个 TP workers | 接收并加载权重,执行推理 | -| 2、3 | torchrun 的两个 FSDP2 ranks | 加载分片模型、训练、导出参数 | - -使用有空闲显存的四张 NVIDIA GPU。当前代码不是 NPU/HCCL 测试。先选普通 dense Qwen3-0.6B、Qwen3-4B 或 Llama 类本地模型;不要先用 MoE、量化、LoRA 或多模态模型。该脚本的模型需具有 `model.model.layers`。 - -## 0. 准备代码和环境 - -把当前分支的代码同步到服务器。教程假设仓库在 `/model/xyr/verl-SpeCo-ls`,模型在 `/nas/disk1/Qwen3-4B`;按实际位置修改。 - -激活服务器现有的、能运行 vLLM 0.23.0 的 CUDA Python 环境。两个终端使用同一个环境。 - -```bash -cd /model/xyr/verl-SpeCo-ls -nvidia-smi -python -c 'import torch; from importlib.metadata import version; print("torch:", torch.__version__, "CUDA:", torch.version.cuda, "vllm:", version("vllm")); assert torch.cuda.is_available(), "CUDA unavailable"' -``` - -不要为此随意升级 PyTorch。已核对的 vLLM 0.23.0 包要求 `torch==2.11.0`,实际 CUDA wheel 还需要与服务器驱动兼容。报驱动过旧、CUDA unavailable 或动态库错误时先修环境,尚未进入本测试逻辑。 - -测试脚本会自动优先导入当前仓库,不必重新 `pip install verl-speco`;运行时会打印 `SOURCE=.../verl_speco/integration/external_vllm_weight_sync.py`,确认没有使用旧安装包。 - -## 1. 终端 A:启动 vLLM,使用卡 0、1 - -确认这两个 GPU 没有被其他训练占用,并确认端口 8000 空闲。 - -```bash -cd /model/xyr/verl-SpeCo-ls -MODEL_PATH=/nas/disk1/Qwen3-4B - -CUDA_VISIBLE_DEVICES=0,1 VLLM_SERVER_DEV_MODE=1 \ -vllm serve "$MODEL_PATH" \ - --host 127.0.0.1 \ - --port 8000 \ - --served-model-name weight-sync-smoke \ - --tensor-parallel-size 2 \ - --dtype bfloat16 \ - --max-model-len 1024 \ - --gpu-memory-utilization 0.5 \ - --enforce-eager \ - --generation-config vllm \ - --no-enable-prefix-caching \ - --load-format dummy \ - --weight-transfer-config '{"backend":"nccl"}' -``` - -保持这个终端运行,等服务就绪。 - -- `VLLM_SERVER_DEV_MODE=1`:启用权重管理接口。只在可信环境使用,这里仅监听本机回环地址。 -- `TP=2`:vLLM 内部把目标模型按 Tensor Parallel 分到两张卡。 -- `--load-format dummy`:不先加载 checkpoint 权重,首次有效推理依赖发送端同步成功。不要在首次同步前用输出判断模型质量。 -- `--generation-config vllm`:避免模型目录的采样默认值影响概率对比。 -- 不需要 hidden-state connector,也不需要启动草稿模型,这一步隔离验证目标模型权重传输。 - -本机四卡测试使用 `127.0.0.1`。如果两端不在同一台服务器,需要另外配置监听地址、网络和 actor rank 0 的可达地址,本教程不覆盖跨机部署。 - -## 2. 终端 B:先检查管理接口 - -```bash -curl --noproxy '*' -fsS http://127.0.0.1:8000/get_world_size -``` - -预期: - -```json -{"world_size":2} -``` - -返回 404 通常表示没有启用 dev API,或者访问了错误的 vLLM 版本/端口。连接拒绝说明服务尚未起来。不是 2 则先核对 vLLM 的 TP/PP/DP 配置。 - -这个接口只是读取服务配置,不会执行推理或修改权重。 - -## 3. 终端 B:用卡 2、3 跑 FSDP2 测试 - -```bash -cd /model/xyr/verl-SpeCo-ls -MODEL_PATH=/nas/disk1/Qwen3-4B - -CUDA_VISIBLE_DEVICES=2,3 \ -python -m torch.distributed.run \ - --standalone \ - --nproc_per_node=2 \ - tools/fsdp2_vllm_weight_sync_smoke.py \ - --model "$MODEL_PATH" \ - --endpoint http://127.0.0.1:8000/v1 \ - --served-model-name weight-sync-smoke \ - --bucket-size-mb 256 \ - --train-steps 1 \ - --lr 0.01 \ - --logprob-atol 0.1 -``` - -不需要 `ray start`。 - -参数含义: - -- `--model`:必须与 vLLM 是同一架构、同一份模型;脚本只从本地目录加载,不下载。 -- `--endpoint`:带 `/v1` 的推理地址。生产代码会自动去掉 `/v1` 来访问管理接口。 -- `--served-model-name`:与终端 A 一致,是 HTTP 请求里的模型名,不是文件路径。 -- `--bucket-size-mb`:沿用实际发送代码的分批目标大小;超大单 Tensor 会独占一批并扩展 packed buffer。 -- `--train-steps`:首次对比后真实训练几步,默认 1。 -- `--lr`:本测试的临时 SGD 学习率,不是正式训练的推荐学习率。 -- `--logprob-atol`:不同推理实现下,同一 token 的 logprob 允许的绝对误差。 -- `--no-packed`:可选,用逐 Tensor 广播替代默认 packed。不是绕过 NCCL。 -- `--timeout`:HTTP 超时秒数,默认 600;NCCL 的底层故障可能仍需要等待其自身超时。 - -模型首先在每个进程的 CPU 内存中加载完整 checkpoint,再按 transformer layer 进行 FSDP2 分片并放到设备。因此 CPU 内存也要足够,不适合直接从超大模型开始。 - -## 4. 应看到哪些日志 - -```text -SOURCE=.../verl_speco/integration/external_vllm_weight_sync.py -FSDP_READY rank=0 global_shape=(...) local_shape=(...) -FSDP_READY rank=1 global_shape=(...) local_shape=(...) -CONNECT_OK -INITIAL_SYNC_OK -COMPARE prompt=0 hf_top1=... vllm_top1=... max_logprob_error=... -... -INITIAL_COMPARE_OK -TRAIN_STEP_OK step=1 loss=... -UPDATED_SYNC_OK -COMPARE prompt=0 hf_top1=... vllm_top1=... max_logprob_error=... -... -CHECKED_TOKEN_REFERENCE_CHANGE=... -UPDATED_COMPARE_OK -PASS: initial and trained FSDP2 weights match external vLLM; no checkpoint written -``` - -`FSDP_READY` 会同时打印全局参数形状与本 rank 的分片形状,确认不是两个完整的普通模型。 - -脚本用三个固定 prompt 的 token IDs,分别对比 FSDP2 模型和 vLLM 的下一 token top-5 logprob。不会用“生成的文字是否一样”作为唯一依据,也不会只凭 HTTP 200 判成功。浮点差异可能让非常接近的 top-1 token 换位,脚本同时检查其概率差距。 - -第二轮还检查被比较 token 的参考概率是否相对训练前充分变化。如果变化不足以排除旧权重,打印 `INCONCLUSIVE` 并退出失败,而不声称热更新已验证。遇到这种情况:在终端 A 停止并重新启动测试服务,再把终端 B 改为 `--train-steps 3 --lr 0.05` 重试。不建议直接放宽误差掩盖问题。 - -## 5. 结束与排查 - -发送端正常完成后会自行关闭 NCCL sender、HTTP Session 和 FSDP2 进程组,不保存文件。vLLM 是你单独启动的,因此到终端 A 按 Ctrl+C 关闭。 - -发送失败后不要继续使用这个测试 vLLM:部分权重可能已更新,服务可能保持暂停。先停止发送端,再停止并重新启动终端 A 的测试服务,然后重试。不要使用 `pkill -f ray` 或其他会误杀无关任务的命令。 - -常见定位: - -- 没到 `CONNECT_OK`:检查 dev API、两端网络可达性和 NCCL 握手。传输 rendezvous 默认使用发送 rank 0 的自动检测 IP,可用 `--master-address <该机器可达IP>` 指定。 -- 卡在 `INITIAL_SYNC`:结合发送端和 vLLM 日志,检查 `update_weights` 错误、NCCL 错误及显存。发送端两张卡都必须参与 FSDP2 导出。 -- 首次 `COMPARE` 超阈值:核对模型、dtype、是否量化、参数名、两端包版本;先不要提高阈值。 -- 初次正确、训练后不正确:重点检查后续权重更新和旧缓存。发送代码在 pause 时清缓存,本教程还禁用了 prefix caching。 -- OOM:先换小模型,或减少 bucket 大小。普通模型加载、FSDP2 导出、vLLM TP 权重和 packed 接收缓冲都有显存需求,`bucket_size_mb` 不是总显存上限。 - -## 6. 与正式 co-train 的关系 - -这份脚本复用的是现有 `initialize_worker_weight_sync()`、`update_worker_weights()`、`close_worker_weight_sync()` 和 `ExternalVllmWeightSender`,不是另写一套 NCCL。 - -只有测试引擎适配部分不同:脚本用原生 FSDP2 的 `state_dict()` / `DTensor.full_tensor()`;正式 co-train 使用实际 verl engine 的 `get_per_tensor_param()`。因此测试通过能验证真实分片汇集和传输,但不能证明 MoE 参数转换、LoRA、offload、Ray 调度或整个 co-train 都已经通过。 - -两个 FSDP ranks 都参与汇集;只有 FSDP rank 0 发给两个 vLLM TP workers。FSDP 训练通信组为 2 个进程,外部权重传输组为 3 个进程(sender + 两个 TP workers),不是 4。 - -本机未运行此 GPU/vLLM 测试,脚本和步骤供服务器手动验证。