diff --git a/examples/v1/config/reasoning_rl_qwen3p5vl_mtp_ep.py b/examples/v1/config/reasoning_rl_qwen3p5vl_mtp_ep.py index 92d038acef..b1bc3a650d 100644 --- a/examples/v1/config/reasoning_rl_qwen3p5vl_mtp_ep.py +++ b/examples/v1/config/reasoning_rl_qwen3p5vl_mtp_ep.py @@ -282,10 +282,9 @@ def group_sample_filter_func(group_samples): produce_strategy_config = AsyncProduceStrategyConfig( over_sample_threshold=1, enable_partial_rollout=1, - is_valid_sample_fn=group_sample_filter_func, max_staleness=3, ) -# produce_strategy_config= SyncProduceStrategyConfig(is_valid_sample_fn=group_sample_filter_func) +# produce_strategy_config = SyncProduceStrategyConfig() # 6. agent loop managers agent_loop_config = SingleTurnAgentLoopConfig( @@ -297,6 +296,7 @@ def group_sample_filter_func(group_samples): task_name="train_task", agent_loop_config=agent_loop_config, judger_config=judger_config, + filter_func=group_sample_filter_func, produce_strategy_config=produce_strategy_config, sampler_config=SamplerConfig(dataloader_cfg=dataloader_cfg, prompt_repeat_k=prompt_repeat_k), ), diff --git a/examples/v1/config/rl_dapo_math_async_filter.py b/examples/v1/config/rl_dapo_math_async_filter.py index 6b99b9a1b4..826768c18f 100644 --- a/examples/v1/config/rl_dapo_math_async_filter.py +++ b/examples/v1/config/rl_dapo_math_async_filter.py @@ -157,13 +157,13 @@ def group_samples_filter_func(rollout_states): enable_partial_rollout=True, max_staleness=0, tail_batch_trigger_size=256, - is_valid_sample_fn=group_samples_filter_func ) agent_loop_manager_cfg = AgentLoopManagerConfig( tasks=TaskSpecConfig( task_name="train_task", agent_loop_config=agent_loop_config, judger_config=judger_config, + filter_func=group_samples_filter_func, produce_strategy_config=produce_strategy_config, sampler_config=sampler_config, ), diff --git a/recipe/on_policy_distillation/build_teacher_server_commands.py b/recipe/on_policy_distillation/build_teacher_server_commands.py new file mode 100644 index 0000000000..87ddc68e4e --- /dev/null +++ b/recipe/on_policy_distillation/build_teacher_server_commands.py @@ -0,0 +1,196 @@ +"""Build executable Teacher server commands from an OPD config. + +The NUL-delimited output starts with the Teacher count. Each Teacher record +contains its name, safe name, endpoint, health URL, model-info URL, +command-argument count, and executable command arguments. An externally +managed Teacher has a command-argument count of zero. +""" + +import argparse +import re +import sys +from pathlib import Path +from typing import Literal +from urllib.parse import urlparse + + +REPO_ROOT = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(REPO_ROOT)) + +from xtuner.v1.rl.on_policy_distillation import OPDTeacherConfig # noqa: E402 +from xtuner.v1.utils.config import Config # noqa: E402 + + +def build_teacher_server_command( + teacher: OPDTeacherConfig, + backend: Literal["sglang", "lmdeploy"], +) -> list[str]: + """Build one executable Teacher server command.""" + if backend == "sglang": + return _build_sglang_command(teacher) + return _build_lmdeploy_command(teacher) + + +def build_teacher_server_commands( + config_path: str, + backend: Literal["sglang", "lmdeploy"], +) -> list[list[str]]: + """Build executable Teacher server command records. + + Args: + config_path (str): Path to the XTuner Python config containing + ``opd_config``. + backend (Literal["sglang", "lmdeploy"]): Teacher serving backend. + + Returns: + list[list[str]]: Teacher metadata and executable argv records. + """ + config = Config.fromfile(config_path) + health_path, model_info_path = _get_teacher_server_paths(backend) + records: list[list[str]] = [] + for teacher in config.opd_config.teachers: + command = build_teacher_server_command(teacher, backend) + safe_name = re.sub(r"[^A-Za-z0-9._-]+", "_", teacher.name).strip("_") or "teacher" + endpoint = teacher.endpoint.rstrip("/") + records.append( + [ + teacher.name, + safe_name, + endpoint, + f"{endpoint}/{health_path}", + f"{endpoint}/{model_info_path}", + str(len(command)), + *command, + ] + ) + return records + + +def _get_teacher_server_paths( + backend: Literal["sglang", "lmdeploy"], +) -> tuple[str, str]: + if backend == "sglang": + return "health_generate", "get_model_info" + return "health", "v1/models" + + +def _parse_teacher_endpoint(teacher: OPDTeacherConfig) -> tuple[str, int]: + parsed_endpoint = urlparse(teacher.endpoint) + if parsed_endpoint.scheme not in {"http", "https"}: + raise ValueError(f"Teacher {teacher.name!r} endpoint must use http or https: {teacher.endpoint}") + if parsed_endpoint.hostname is None or parsed_endpoint.port is None: + raise ValueError( + f"Teacher {teacher.name!r} endpoint must contain an explicit host and port: {teacher.endpoint}" + ) + if parsed_endpoint.path not in {"", "/"} or parsed_endpoint.params or parsed_endpoint.query: + raise ValueError( + f"Teacher {teacher.name!r} endpoint must be a server base URL without a path or query: {teacher.endpoint}" + ) + return parsed_endpoint.hostname, parsed_endpoint.port + + +def _build_sglang_command(teacher: OPDTeacherConfig) -> list[str]: + host, port = _parse_teacher_endpoint(teacher) + config = teacher.launch_config + if config is None: + return [] + + tensor_parallel_size = config.tensor_parallel_size + if config.expert_parallel_size > 1: + tensor_parallel_size = config.expert_parallel_size + + command = [ + "env", + f"CUDA_VISIBLE_DEVICES={config.cuda_visible_devices}", + sys.executable, + "-m", + "sglang.launch_server", + "--model-path", + str(config.model_path), + "--host", + host, + "--port", + str(port), + "--dtype", + config.dtype, + "--tp", + str(tensor_parallel_size), + "--ep", + str(config.expert_parallel_size), + "--mem-fraction-static", + str(config.gpu_memory_utilization), + ] + if config.context_length is not None: + command.extend(["--context-length", str(config.context_length)]) + if config.max_batch_size is not None: + command.extend(["--max-running-requests", str(config.max_batch_size)]) + if config.chunked_prefill_size is not None: + command.extend(["--chunked-prefill-size", str(config.chunked_prefill_size)]) + return command + + +def _build_lmdeploy_command(teacher: OPDTeacherConfig) -> list[str]: + host, port = _parse_teacher_endpoint(teacher) + config = teacher.launch_config + if config is None: + return [] + + data_parallel_size = config.expert_parallel_size if config.expert_parallel_size > 1 else 1 + command = [ + "env", + f"CUDA_VISIBLE_DEVICES={config.cuda_visible_devices}", + sys.executable, + "-m", + "lmdeploy", + "serve", + "api_server", + str(config.model_path), + "--backend", + "pytorch", + "--role", + "Hybrid", + "--logprobs-mode", + "raw_logprobs", + "--server-name", + host, + "--server-port", + str(port), + "--dtype", + config.dtype, + "--tp", + str(config.tensor_parallel_size), + "--ep", + str(config.expert_parallel_size), + "--dp", + str(data_parallel_size), + "--cache-max-entry-count", + str(config.gpu_memory_utilization), + ] + if config.context_length is not None: + command.extend(["--session-len", str(config.context_length)]) + if config.max_batch_size is not None: + command.extend(["--max-batch-size", str(config.max_batch_size)]) + if config.max_prefill_token_num is not None: + command.extend(["--max-prefill-token-num", str(config.max_prefill_token_num)]) + return command + + +def _write_teacher_records(records: list[list[str]]) -> None: + fields = [str(len(records))] + for record in records: + fields.extend(record) + + payload = "\0".join(fields) + "\0" + sys.stdout.buffer.write(payload.encode("utf-8")) + + +def _main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("config_path") + parser.add_argument("backend", choices=("sglang", "lmdeploy")) + args = parser.parse_args() + _write_teacher_records(build_teacher_server_commands(args.config_path, args.backend)) + + +if __name__ == "__main__": + _main() diff --git a/recipe/on_policy_distillation/config/rl_dapo_math_mopd.py b/recipe/on_policy_distillation/config/rl_dapo_math_mopd.py new file mode 100644 index 0000000000..a00ed68d80 --- /dev/null +++ b/recipe/on_policy_distillation/config/rl_dapo_math_mopd.py @@ -0,0 +1,335 @@ +import json +import os +from pathlib import Path + +from xtuner.v1.config import AdamWConfig, FSDPConfig, LRConfig +from xtuner.v1.data_proto.rl_data import SampleParams +from xtuner.v1.datasets.config import DataloaderConfig, DatasetConfig +from xtuner.v1.datasets.rl_tokenize_fn import RLQwen3VLTokenizeFnConfig +from xtuner.v1.model import get_model_config_from_hf +from xtuner.v1.rl.advantage import GRPOAdvantageConfig +from xtuner.v1.rl.agent_loop import SingleTurnAgentLoopConfig +from xtuner.v1.rl.agent_loop_manager import ( + AgentLoopManagerConfig, + SamplerConfig, + SyncProduceStrategyConfig, + TaskSpecConfig, +) +from xtuner.v1.rl.evaluator import EvaluatorConfig +from xtuner.v1.rl.judger import ComposedJudgerConfig, GEO3KJudgerConfig, GSM8KJudgerConfig +from xtuner.v1.rl.loss import GRPOLossConfig +from xtuner.v1.rl.on_policy_distillation import ( + OPDConfig, + OPDTeacherConfig, + OPDTeacherLaunchConfig, +) +from xtuner.v1.rl.replay_buffer import SyncReplayBufferConfig +from xtuner.v1.rl.rollout.worker import RolloutConfig +from xtuner.v1.rl.trainer import WorkerConfig +from xtuner.v1.rl.utils import ( + AcceleratorResourcesConfig, + CPUResourcesConfig, +) +from xtuner.v1.train.rl_trainer import RLColocateTrainerConfig + + +work_dir = os.environ["WORK_DIR"] +model_path = os.environ["MODEL_PATH"] +gsm8k_teacher_model_path = os.environ["GSM8K_TEACHER_MODEL_PATH"] +geo3k_teacher_model_path = os.environ["GEO3K_TEACHER_MODEL_PATH"] +meta_data_path = os.environ["DATA_PATH"] +eval_meta_data_path = os.environ.get("EVAL_DATA_PATH", "") +NNODE = int(os.environ.get("WORLD_SIZE", "1")) + + +def _as_list(value): + return value if isinstance(value, list) else [value] + + +# Training shape aligned with verl PR #6051: +# examples/on_policy_distillation_trainer/run_qwen3_mopd_gsm8k_geo3k.sh. +# Teacher roles and model families follow the GSM8K/Geo3K experiment. +experimental_name = "dapo_math_mopd" +total_epochs = 15 +train_batch_size = 128 +prompt_repeat_k = 1 +rollout_tp_size = 1 +rollout_ep_size = 1 +max_prompt_length = 1024 +max_response_length = 2048 +pack_max_length = max_prompt_length + max_response_length +max_num_tokens = pack_max_length +train_optimizer_steps = 1 +enable_evaluate = bool(eval_meta_data_path) +evaluate_step = 5 +eval_prompt_repeat_k = 1 +checkpoint_interval = 200 + +# 1. resources: two colocated Student workers, plus one GPU per Teacher. +resources = AcceleratorResourcesConfig( + accelerator="GPU", + num_workers=2 * NNODE, + num_cpus_per_worker=12, + cpu_memory_per_worker=16 * 1024**3, # 16 GB +) + +# 2. rollout +rollout_config = RolloutConfig( + env=experimental_name, + device=resources.accelerator, + model_path=model_path, + dtype="bfloat16", + tensor_parallel_size=rollout_tp_size, + expert_parallel_size=rollout_ep_size, + gpu_memory_utilization=0.6, + context_length=max_response_length + max_prompt_length, + enable_return_routed_experts=False, + rollout_max_batch_size_per_instance=2048, +) + +# 3. train worker +lr_cfg = LRConfig(lr_type="constant", warmup_ratio=0, lr_min=1e-6) +fsdp_cfg = FSDPConfig( + torch_compile=False, + cpu_offload=False, + ep_size=1, + reduce_dtype="float32", +) +model_cfg = get_model_config_from_hf(Path(model_path)) +if hasattr(model_cfg, "balancing_loss_cfg"): + model_cfg.balancing_loss_cfg = None +if hasattr(model_cfg, "z_loss_cfg"): + model_cfg.z_loss_cfg = None +optim_cfg = AdamWConfig(lr=1e-6, foreach=False, weight_decay=0.1, betas=(0.9, 0.98)) +loss_cfg = GRPOLossConfig( + policy_loss_cfg={ + "cliprange_high": 0.2, + "cliprange_low": 0.2, + "loss_type": "vanilla", + "clip_ratio_c": 10.0, + "log_prob_diff_min": -20.0, + "log_prob_diff_max": 20.0, + }, + ignore_idx=-100, + use_kl_loss=False, + kl_loss_coef=0.0, + kl_loss_type="low_var_kl", + mode="chunk", + chunk_size=512, +) +train_worker_cfg = WorkerConfig( + model_cfg=model_cfg, + load_from=model_path, + optim_cfg=optim_cfg, + loss_cfg=loss_cfg, + lr_cfg=lr_cfg, + fsdp_cfg=fsdp_cfg, + sp_size=1, + optimizer_steps=train_optimizer_steps, + pack_max_length=pack_max_length, +) + +# 4. train agent loop manager +with open(meta_data_path, "r", encoding="utf-8") as f: + ds_collections = json.load(f) + +train_dataset_cfg = [] +for name, data in ds_collections.items(): + annotations = _as_list(data["annotation"]) + for annotation in annotations: + train_dataset_cfg.append( + { + "dataset": DatasetConfig( + name=name, + anno_path=annotation, + media_root=data.get("media_root", ""), + sample_ratio=data.get("sample_ratio", 1.0), + class_name="VLMJsonlDataset", + ), + "tokenize_fn": RLQwen3VLTokenizeFnConfig( + processor_path=model_path, + max_length=max_prompt_length, + system_message=data.get("system_message", None), + chat_template="qwen3-vl", + add_generation_prompt=True, + enable_thinking=True, + ), + } + ) + +dataloader_cfg = DataloaderConfig( + dataset_config_list=train_dataset_cfg, + num_workers=8, + pack_max_length=pack_max_length, + collator="fake_collator", + pack_level="none", +) +sampler_config = SamplerConfig( + dataloader_cfg=dataloader_cfg, + prompt_repeat_k=prompt_repeat_k, +) +training_sample_params = SampleParams( + max_tokens=max_response_length, + top_k=0, + top_p=1.0, + temperature=1.0, + min_tokens=0, + skip_special_tokens=False, + return_routed_experts=False, +) +agent_loop_config = SingleTurnAgentLoopConfig( + hf_checkpoint=model_path, + sample_params=training_sample_params, +) +produce_strategy_config = SyncProduceStrategyConfig() +agent_loop_manager_cfg = AgentLoopManagerConfig( + tasks=TaskSpecConfig( + task_name="train_task", + agent_loop_config=agent_loop_config, + produce_strategy_config=produce_strategy_config, + sampler_config=sampler_config, + ), +) + +# 5. evaluation +eval_agent_loop_manager_cfg = None +evaluator_config = None +if enable_evaluate: + with open(eval_meta_data_path, "r", encoding="utf-8") as f: + eval_ds_collections = json.load(f) + + eval_dataset_cfg = [] + for name, data in eval_ds_collections.items(): + annotations = _as_list(data["annotation"]) + for annotation in annotations: + eval_dataset_cfg.append( + { + "dataset": DatasetConfig( + name=name, + anno_path=annotation, + media_root=data.get("media_root", ""), + sample_ratio=data.get("sample_ratio", 1.0), + class_name="VLMJsonlDataset", + ), + "tokenize_fn": RLQwen3VLTokenizeFnConfig( + processor_path=model_path, + max_length=max_prompt_length, + system_message=data.get("system_message", None), + chat_template="qwen3-vl", + add_generation_prompt=True, + enable_thinking=True, + ignore_multimodal_info=True, + ), + } + ) + + eval_judger_config = ComposedJudgerConfig( + branches={ + "openai/gsm8k": GSM8KJudgerConfig( + judger_name="openai/gsm8k", + cpu_resources=CPUResourcesConfig(num_workers=1, num_cpus_per_worker=1), + ), + "hiyouga/geometry3k": GEO3KJudgerConfig( + cpu_resources=CPUResourcesConfig(num_workers=1, num_cpus_per_worker=1), + ), + } + ) + eval_dataloader_cfg = DataloaderConfig( + dataset_config_list=eval_dataset_cfg, + num_workers=8, + pack_max_length=pack_max_length, + collator="fake_collator", + pack_level="none", + ) + eval_sampler_config = SamplerConfig( + dataloader_cfg=eval_dataloader_cfg, + prompt_repeat_k=eval_prompt_repeat_k, + ) + evaluation_sample_params = SampleParams( + max_tokens=max_response_length, + top_k=0, + top_p=1.0, + temperature=1.0, + min_tokens=0, + skip_special_tokens=False, + return_routed_experts=False, + ) + eval_agent_loop_config = SingleTurnAgentLoopConfig( + hf_checkpoint=model_path, + sample_params=evaluation_sample_params, + ) + eval_agent_loop_manager_cfg = AgentLoopManagerConfig( + tasks=TaskSpecConfig( + task_name="eval_task", + agent_loop_config=eval_agent_loop_config, + judger_config=eval_judger_config, + sampler_config=eval_sampler_config, + ), + ) + evaluator_config = EvaluatorConfig() + +# 6. multi-teacher pure on-policy distillation +# +# This is the only topology block that needs to be edited for an experiment. +# Every training record's data_source must have an entry in +# data_source_teacher_map. Teacher model paths are read from +# GSM8K_TEACHER_MODEL_PATH and GEO3K_TEACHER_MODEL_PATH. +opd_config = OPDConfig( + mode="pg-opd", + task_adv_weight=0.0, + opd_adv_weight=1.0, + teachers=[ + OPDTeacherConfig( + name="gsm8k_teacher", + endpoint="http://127.0.0.1:13141", + launch_config=OPDTeacherLaunchConfig( + model_path=gsm8k_teacher_model_path, + cuda_visible_devices="6", + tensor_parallel_size=1, + expert_parallel_size=1, + context_length=max_num_tokens, + max_batch_size=max_num_tokens, + gpu_memory_utilization=0.8, + ), + ), + OPDTeacherConfig( + name="geo3k_teacher", + endpoint="http://127.0.0.1:13142", + launch_config=OPDTeacherLaunchConfig( + model_path=geo3k_teacher_model_path, + cuda_visible_devices="7", + tensor_parallel_size=1, + expert_parallel_size=1, + context_length=max_num_tokens, + max_batch_size=max_num_tokens, + gpu_memory_utilization=0.8, + ), + ), + ], + data_source_teacher_map={ + "openai/gsm8k": "gsm8k_teacher", + "hiyouga/geometry3k": "geo3k_teacher", + }, +) + +trainer = RLColocateTrainerConfig( + resources=resources, + train_worker_cfg=train_worker_cfg, # TODO: uniform naming of cfg and config + rollout_config=rollout_config, + tokenizer_path=model_path, + replay_buffer_config=SyncReplayBufferConfig(), + agent_loop_manager_cfg=agent_loop_manager_cfg, + eval_agent_loop_manager_cfg=eval_agent_loop_manager_cfg, + evaluator_config=evaluator_config, + load_from=model_path, + train_batch_size=train_batch_size, + advantage_estimator_config=GRPOAdvantageConfig(eps=1e-8), + opd_config=opd_config, + enable_evaluate=enable_evaluate, + enable_initial_evaluate=enable_evaluate, + evaluate_step=evaluate_step, + total_epochs=total_epochs, + work_dir=work_dir, + seed=1234, + debug_rollout=False, +) diff --git a/recipe/on_policy_distillation/config/rl_dapo_math_opd.py b/recipe/on_policy_distillation/config/rl_dapo_math_opd.py new file mode 100644 index 0000000000..a548ebcc62 --- /dev/null +++ b/recipe/on_policy_distillation/config/rl_dapo_math_opd.py @@ -0,0 +1,246 @@ +import os +from pathlib import Path + +from transformers import AutoTokenizer +from xtuner.v1.config import AdamWConfig, FSDPConfig, LRConfig +from xtuner.v1.data_proto.rl_data import SampleParams +from xtuner.v1.datasets.config import DataloaderConfig, DatasetConfig +from xtuner.v1.datasets.rl_tokenize_fn import RLTextTokenizeFnConfig +from xtuner.v1.model import get_model_config_from_hf +from xtuner.v1.rl.advantage import GRPOAdvantageConfig +from xtuner.v1.rl.agent_loop import SingleTurnAgentLoopConfig +from xtuner.v1.rl.agent_loop_manager import ( + AgentLoopManagerConfig, + SamplerConfig, + SyncProduceStrategyConfig, + TaskSpecConfig, +) +from xtuner.v1.rl.evaluator import EvaluatorConfig +from xtuner.v1.rl.judger import DapoMathJudgerConfig +from xtuner.v1.rl.loss import GRPOLossConfig +from xtuner.v1.rl.on_policy_distillation import ( + OPDConfig, + OPDTeacherConfig, + OPDTeacherLaunchConfig, +) +from xtuner.v1.rl.replay_buffer import SyncReplayBufferConfig +from xtuner.v1.rl.rollout.worker import RolloutConfig +from xtuner.v1.rl.trainer import WorkerConfig +from xtuner.v1.rl.utils import ( + AcceleratorResourcesConfig, + CPUResourcesConfig, + get_eos_token, +) +from xtuner.v1.train.rl_trainer import RLColocateTrainerConfig + + +work_dir = os.environ["WORK_DIR"] +model_path = os.environ["MODEL_PATH"] +teacher_model_path = os.environ["TEACHER_MODEL_PATH"] +data_path = os.environ["DATA_PATH"] +eval_data_path = os.environ.get("EVAL_DATA_PATH", "") +NNODE = int(os.environ.get("WORLD_SIZE", "1")) + +# basic settings +experimental_name = "dapo_math_opd" +total_train_steps = 300 +train_batch_size = 16 +prompt_repeat_k = 4 +rollout_tp_size = 1 +rollout_ep_size = 1 +max_prompt_length = 2048 +max_response_length = 16384 +pack_max_length = max_prompt_length + max_response_length +train_optimizer_steps = 1 +enable_evaluate = bool(eval_data_path) +evaluate_step = 20 +eval_prompt_repeat_k = 16 + +# 1. resources +resources = AcceleratorResourcesConfig( + accelerator="GPU", + num_workers=4 * NNODE, + num_cpus_per_worker=12, + cpu_memory_per_worker=16 * 1024**3, # 16 GB +) + +# 2. rollout +rollout_config = RolloutConfig( + env=experimental_name, + device=resources.accelerator, + model_path=model_path, + dtype="bfloat16", + tensor_parallel_size=rollout_tp_size, + expert_parallel_size=rollout_ep_size, + gpu_memory_utilization=0.6, + context_length=max_response_length + max_prompt_length, + enable_return_routed_experts=False, + rollout_max_batch_size_per_instance=2048, +) + +# 3. train worker +lr_cfg = LRConfig(lr_type="constant", warmup_ratio=0, lr_min=1e-6) +fsdp_cfg = FSDPConfig( + torch_compile=False, + cpu_offload=False, + ep_size=1, + reduce_dtype="float32", +) +model_cfg = get_model_config_from_hf(Path(model_path)) +if hasattr(model_cfg, "balancing_loss_cfg"): + model_cfg.balancing_loss_cfg = None +if hasattr(model_cfg, "z_loss_cfg"): + model_cfg.z_loss_cfg = None +optim_cfg = AdamWConfig(lr=1e-6, foreach=False, weight_decay=0.1, betas=(0.9, 0.98)) +loss_cfg = GRPOLossConfig( + policy_loss_cfg={ + "cliprange_high": 0.28, + "cliprange_low": 0.2, + "loss_type": "vanilla", + "clip_ratio_c": 10.0, + "log_prob_diff_min": -20.0, + "log_prob_diff_max": 20.0, + }, + ignore_idx=-100, + use_kl_loss=False, + kl_loss_coef=0.0, + kl_loss_type="low_var_kl", + mode="chunk", + chunk_size=512, +) +train_worker_cfg = WorkerConfig( + model_cfg=model_cfg, + load_from=model_path, + optim_cfg=optim_cfg, + loss_cfg=loss_cfg, + lr_cfg=lr_cfg, + fsdp_cfg=fsdp_cfg, + sp_size=1, + optimizer_steps=train_optimizer_steps, + pack_max_length=pack_max_length, +) + +# 4. train agent loop manager +train_dataset = DatasetConfig(name=experimental_name, anno_path=data_path) +tokenizer_config = RLTextTokenizeFnConfig(max_length=max_prompt_length) +train_dataset_cfg = [{"dataset": train_dataset, "tokenize_fn": tokenizer_config}] +dataloader_cfg = DataloaderConfig( + dataset_config_list=train_dataset_cfg, + pack_max_length=pack_max_length, + collator="fake_collator", + pack_level="none", +) +sampler_config = SamplerConfig( + dataloader_cfg=dataloader_cfg, + prompt_repeat_k=prompt_repeat_k, +) +training_sample_params = SampleParams( + max_tokens=max_response_length, + top_k=0, + top_p=1.0, + temperature=1.0, + min_tokens=0, + skip_special_tokens=False, + return_routed_experts=False, +) +agent_loop_config = SingleTurnAgentLoopConfig( + hf_checkpoint=model_path, + sample_params=training_sample_params, +) +produce_strategy_config = SyncProduceStrategyConfig() +agent_loop_manager_cfg = AgentLoopManagerConfig( + tasks=TaskSpecConfig( + task_name="train_task", + agent_loop_config=agent_loop_config, + produce_strategy_config=produce_strategy_config, + sampler_config=sampler_config, + ), +) + +# 5. evaluation +eval_agent_loop_manager_cfg = None +evaluator_config = None +if enable_evaluate: + eos_token_id = get_eos_token(model_path) + tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True) + eos_token = tokenizer.convert_ids_to_tokens(eos_token_id) + eval_judger_config = DapoMathJudgerConfig( + judger_name="aime_math", + cpu_resources=CPUResourcesConfig(num_workers=1, num_cpus_per_worker=1), + eos_token=eos_token, + enable_overlong_buffer=False, + ) + eval_dataset = DatasetConfig(name="aime", anno_path=eval_data_path, sample_ratio=1.0) + eval_dataset_cfg = [{"dataset": eval_dataset, "tokenize_fn": tokenizer_config}] + eval_dataloader_cfg = DataloaderConfig( + dataset_config_list=eval_dataset_cfg, + pack_max_length=pack_max_length, + collator="fake_collator", + pack_level="none", + ) + eval_sampler_config = SamplerConfig( + dataloader_cfg=eval_dataloader_cfg, + prompt_repeat_k=eval_prompt_repeat_k, + ) + evaluation_sample_params = SampleParams( + max_tokens=max_response_length, + top_k=0, + top_p=1.0, + temperature=1.0, + min_tokens=0, + skip_special_tokens=False, + return_routed_experts=False, + ) + eval_agent_loop_config = SingleTurnAgentLoopConfig( + hf_checkpoint=model_path, + sample_params=evaluation_sample_params, + ) + eval_agent_loop_manager_cfg = AgentLoopManagerConfig( + tasks=TaskSpecConfig( + task_name="eval_task", + agent_loop_config=eval_agent_loop_config, + judger_config=eval_judger_config, + sampler_config=eval_sampler_config, + ), + ) + evaluator_config = EvaluatorConfig() + +# 6. pure on-policy distillation +opd_config = OPDConfig( + mode="pg-opd", + task_adv_weight=0.0, + opd_adv_weight=1.0, + teachers=[ + OPDTeacherConfig( + name="teacher", + endpoint="http://127.0.0.1:13141", + launch_config=OPDTeacherLaunchConfig( + model_path=teacher_model_path, + cuda_visible_devices="7", + ), + ) + ], + data_source_teacher_map={"math_dapo": "teacher"}, +) + +trainer = RLColocateTrainerConfig( + resources=resources, + train_worker_cfg=train_worker_cfg, # TODO: uniform naming of cfg and config + rollout_config=rollout_config, + tokenizer_path=model_path, + replay_buffer_config=SyncReplayBufferConfig(), + agent_loop_manager_cfg=agent_loop_manager_cfg, + eval_agent_loop_manager_cfg=eval_agent_loop_manager_cfg, + evaluator_config=evaluator_config, + load_from=model_path, + train_batch_size=train_batch_size, + advantage_estimator_config=GRPOAdvantageConfig(eps=1e-8), + opd_config=opd_config, + enable_evaluate=enable_evaluate, + enable_initial_evaluate=enable_evaluate, + evaluate_step=evaluate_step, + total_train_steps=total_train_steps, + work_dir=work_dir, + seed=1234, + debug_rollout=False, +) diff --git a/recipe/on_policy_distillation/scripts/launch_teacher_utils.sh b/recipe/on_policy_distillation/scripts/launch_teacher_utils.sh new file mode 100644 index 0000000000..1a6134f833 --- /dev/null +++ b/recipe/on_policy_distillation/scripts/launch_teacher_utils.sh @@ -0,0 +1,247 @@ +#!/usr/bin/env bash + +# Source-only helpers for launching, waiting for, and stopping OPD Teacher servers. + +OPD_RECIPE_DIR=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")/.." && pwd) + +TEACHER_NAMES=() +TEACHER_ENDPOINTS=() +TEACHER_HEALTH_URLS=() +TEACHER_MODEL_INFO_URLS=() +TEACHER_LOG_FILES=() +TEACHER_PIDS=() + +start_single_teacher_server() { + _start_teacher_servers "$1" "$2" "$3" "1" "1" +} + +start_teacher_servers() { + _start_teacher_servers "$1" "$2" "$3" "" "0" +} + +wait_for_teacher_servers() { + local startup_timeout_s=$1 + local deadline=$((SECONDS + startup_timeout_s)) + local teacher_index + local health_check_index + local pid + local log_file + local -a teacher_ready=() + local -a health_check_pids=() + local -a health_check_indices=() + local -a pending_names=() + + if (( ${#TEACHER_NAMES[@]} == 0 )); then + echo "No Teacher servers have been configured." >&2 + return 1 + fi + + for teacher_index in "${!TEACHER_NAMES[@]}"; do + teacher_ready[teacher_index]=0 + done + + while (( SECONDS < deadline )); do + for teacher_index in "${!TEACHER_NAMES[@]}"; do + pid="${TEACHER_PIDS[teacher_index]}" + log_file="${TEACHER_LOG_FILES[teacher_index]}" + if [[ -n "${pid}" ]] && ! kill -0 "${pid}" 2>/dev/null; then + echo "Teacher ${TEACHER_NAMES[teacher_index]} exited before becoming ready." >&2 + if [[ -n "${log_file}" ]]; then + tail -n 50 "${log_file}" >&2 || true + fi + return 1 + fi + done + + health_check_pids=() + health_check_indices=() + for teacher_index in "${!TEACHER_NAMES[@]}"; do + if (( teacher_ready[teacher_index] )); then + continue + fi + curl -sf --max-time 2 \ + "${TEACHER_HEALTH_URLS[teacher_index]}" \ + >/dev/null 2>&1 & + health_check_pids+=("$!") + health_check_indices+=("${teacher_index}") + done + + for health_check_index in "${!health_check_pids[@]}"; do + teacher_index=${health_check_indices[health_check_index]} + if wait "${health_check_pids[health_check_index]}"; then + teacher_ready[teacher_index]=1 + echo "Teacher ${TEACHER_NAMES[teacher_index]} is ready at ${TEACHER_ENDPOINTS[teacher_index]}" + fi + done + + pending_names=() + for teacher_index in "${!TEACHER_NAMES[@]}"; do + if (( ! teacher_ready[teacher_index] )); then + pending_names+=("${TEACHER_NAMES[teacher_index]}") + fi + done + if (( ${#pending_names[@]} == 0 )); then + return 0 + fi + + echo "Waiting for teachers: ${pending_names[*]}" + sleep 5 + done + + for teacher_index in "${!TEACHER_NAMES[@]}"; do + if (( teacher_ready[teacher_index] )); then + continue + fi + echo "Teacher ${TEACHER_NAMES[teacher_index]} did not become ready within ${startup_timeout_s}s." >&2 + log_file="${TEACHER_LOG_FILES[teacher_index]}" + if [[ -n "${log_file}" ]]; then + tail -n 50 "${log_file}" >&2 || true + fi + done + return 1 +} + +stop_teacher_servers() { + local attempt + local any_alive + local pid + + for pid in "${TEACHER_PIDS[@]}"; do + if [[ -n "${pid}" ]]; then + kill -TERM -- "-${pid}" 2>/dev/null || true + fi + done + + for ((attempt = 0; attempt < 30; attempt++)); do + any_alive=0 + for pid in "${TEACHER_PIDS[@]}"; do + if [[ -n "${pid}" ]] && kill -0 -- "-${pid}" 2>/dev/null; then + any_alive=1 + break + fi + done + if (( ! any_alive )); then + break + fi + sleep 1 + done + + for pid in "${TEACHER_PIDS[@]}"; do + if [[ -n "${pid}" ]]; then + if kill -0 -- "-${pid}" 2>/dev/null; then + kill -KILL -- "-${pid}" 2>/dev/null || true + fi + wait "${pid}" 2>/dev/null || true + fi + done + + TEACHER_PIDS=() +} + +_start_teacher_servers() { + local config_file=$1 + local backend=$2 + local work_dir=$3 + local expected_teacher_count=$4 + local require_local_teacher=$5 + local teacher_count + local teacher_index + local teacher_field_offset=1 + local teacher_name + local teacher_safe_name + local teacher_endpoint + local teacher_health_url + local teacher_model_info_url + local teacher_command_arg_count + local teacher_command_offset + local teacher_log_file + local -a teacher_fields=() + local -a teacher_command=() + + _reset_teacher_server_state + mkdir -p "${work_dir}" + + # Command builder output: + # count, then repeated name, safe name, endpoint, health URL, + # model-info URL, command-argument count, and command argv. + mapfile -d "" -t teacher_fields < <( + python "${OPD_RECIPE_DIR}/build_teacher_server_commands.py" \ + "${config_file}" "${backend}" + ) + if (( ${#teacher_fields[@]} == 0 )); then + echo "Teacher command builder did not return any records." >&2 + return 1 + fi + + teacher_count=${teacher_fields[0]} + if [[ -n "${expected_teacher_count}" ]] && (( teacher_count != expected_teacher_count )); then + echo "Expected ${expected_teacher_count} Teacher, got ${teacher_count}." >&2 + return 1 + fi + if (( teacher_count == 0 )); then + echo "OPD config does not contain any Teachers: ${config_file}" >&2 + return 1 + fi + + echo "Teacher backend: ${backend}" + for ((teacher_index = 0; teacher_index < teacher_count; teacher_index++)); do + teacher_name=${teacher_fields[teacher_field_offset]} + teacher_safe_name=${teacher_fields[teacher_field_offset + 1]} + teacher_endpoint=${teacher_fields[teacher_field_offset + 2]} + teacher_health_url=${teacher_fields[teacher_field_offset + 3]} + teacher_model_info_url=${teacher_fields[teacher_field_offset + 4]} + teacher_command_arg_count=${teacher_fields[teacher_field_offset + 5]} + teacher_command_offset=$((teacher_field_offset + 6)) + teacher_command=( + "${teacher_fields[@]:teacher_command_offset:teacher_command_arg_count}" + ) + teacher_field_offset=$((teacher_command_offset + teacher_command_arg_count)) + + TEACHER_NAMES+=("${teacher_name}") + TEACHER_ENDPOINTS+=("${teacher_endpoint}") + TEACHER_HEALTH_URLS+=("${teacher_health_url}") + TEACHER_MODEL_INFO_URLS+=("${teacher_model_info_url}") + + if (( teacher_command_arg_count == 0 )); then + if (( require_local_teacher )); then + echo "Teacher ${teacher_name} must define launch_config for local startup." >&2 + return 1 + fi + echo "Using externally managed Teacher ${teacher_name} at ${teacher_endpoint}" + TEACHER_LOG_FILES+=("") + TEACHER_PIDS+=("") + continue + fi + + if [[ -n "${expected_teacher_count}" ]]; then + teacher_log_file="${work_dir}/teacher.log" + else + teacher_log_file="${work_dir}/teacher_${teacher_index}_${teacher_safe_name}.log" + fi + TEACHER_LOG_FILES+=("${teacher_log_file}") + + echo "Starting Teacher ${teacher_name}" + echo "Teacher endpoint: ${teacher_endpoint}" + echo "Teacher log: ${teacher_log_file}" + + setsid env \ + PYTHONUNBUFFERED=1 \ + "${teacher_command[@]}" \ + >"${teacher_log_file}" 2>&1 & + TEACHER_PIDS+=("$!") + done + + if (( teacher_field_offset != ${#teacher_fields[@]} )); then + echo "Teacher command builder returned malformed records." >&2 + return 1 + fi +} + +_reset_teacher_server_state() { + TEACHER_NAMES=() + TEACHER_ENDPOINTS=() + TEACHER_HEALTH_URLS=() + TEACHER_MODEL_INFO_URLS=() + TEACHER_LOG_FILES=() + TEACHER_PIDS=() +} diff --git a/recipe/on_policy_distillation/scripts/run_sampled_token_mopd.sh b/recipe/on_policy_distillation/scripts/run_sampled_token_mopd.sh new file mode 100755 index 0000000000..4241b4d1a2 --- /dev/null +++ b/recipe/on_policy_distillation/scripts/run_sampled_token_mopd.sh @@ -0,0 +1,75 @@ +#!/usr/bin/env bash + +set -euo pipefail + +SCRIPT_DIR=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd) +REPO_ROOT=$(cd -- "${SCRIPT_DIR}/../../.." && pwd) +source "${SCRIPT_DIR}/launch_teacher_utils.sh" + +export STUDENT_MODEL_PATH=${1:?"Usage: $0 STUDENT_MODEL_PATH DATA_PATH"} +export MODEL_PATH="${STUDENT_MODEL_PATH}" +export DATA_PATH=${2:?"Usage: $0 STUDENT_MODEL_PATH DATA_PATH"} +export GSM8K_TEACHER_MODEL_PATH=${GSM8K_TEACHER_MODEL_PATH:?"GSM8K_TEACHER_MODEL_PATH is required"} +export GEO3K_TEACHER_MODEL_PATH=${GEO3K_TEACHER_MODEL_PATH:?"GEO3K_TEACHER_MODEL_PATH is required"} + +export OPD_CONFIG_PATH="${OPD_CONFIG_PATH:-recipe/on_policy_distillation/config/rl_dapo_math_mopd.py}" +export EVAL_DATA_PATH="${EVAL_DATA_PATH:-}" +export STUDENT_CUDA_VISIBLE_DEVICES="${STUDENT_CUDA_VISIBLE_DEVICES:-0,1,2,3}" +export TEACHER_STARTUP_TIMEOUT_S="${TEACHER_STARTUP_TIMEOUT_S:-1200}" +export PYTHONUNBUFFERED=1 + +USE_SGLANG=${XTUNER_USE_SGLANG:-0} +USE_LMDEPLOY=${XTUNER_USE_LMDEPLOY:-0} +USE_VLLM=${XTUNER_USE_VLLM:-0} + +if [[ "${USE_SGLANG}" == "1" && "${USE_LMDEPLOY}" == "0" && "${USE_VLLM}" == "0" ]]; then + OPD_BACKEND="sglang" +elif [[ "${USE_SGLANG}" == "0" && "${USE_LMDEPLOY}" == "1" && "${USE_VLLM}" == "0" ]]; then + OPD_BACKEND="lmdeploy" +else + echo "Exactly one of XTUNER_USE_SGLANG and XTUNER_USE_LMDEPLOY must be set to 1; XTUNER_USE_VLLM must be 0." >&2 + exit 1 +fi + +export XTUNER_USE_SGLANG="${USE_SGLANG}" +export XTUNER_USE_LMDEPLOY="${USE_LMDEPLOY}" +export XTUNER_USE_VLLM="${USE_VLLM}" + +export WORK_DIR="${WORK_DIR:-${REPO_ROOT}/work_dirs/dapo_math_mopd}" +export OPD_CONFIG_FILE="${REPO_ROOT}/${OPD_CONFIG_PATH}" + +TRAINING_STARTED=0 + +cleanup() { + local exit_code=$? + + trap - EXIT INT TERM + + stop_teacher_servers + + if (( TRAINING_STARTED )); then + ray stop --force >/dev/null 2>&1 || true + fi + + exit "${exit_code}" +} + +trap cleanup EXIT +trap "exit 130" INT +trap "exit 143" TERM + +start_teacher_servers "${OPD_CONFIG_FILE}" "${OPD_BACKEND}" "${WORK_DIR}" +wait_for_teacher_servers "${TEACHER_STARTUP_TIMEOUT_S}" + +echo "All ${#TEACHER_NAMES[@]} teachers are ready." +echo "Starting MOPD training with student GPUs: ${STUDENT_CUDA_VISIBLE_DEVICES}" + +cd "${REPO_ROOT}" +TRAINING_STARTED=1 +CUDA_VISIBLE_DEVICES="${STUDENT_CUDA_VISIBLE_DEVICES}" \ + bash -o pipefail examples/v1/scripts/run_rl.sh \ + "${OPD_CONFIG_FILE}" \ + "${OPD_BACKEND}" \ + "${STUDENT_MODEL_PATH}" \ + "${DATA_PATH}" \ + "${EVAL_DATA_PATH}" diff --git a/recipe/on_policy_distillation/scripts/run_sampled_token_opd.sh b/recipe/on_policy_distillation/scripts/run_sampled_token_opd.sh new file mode 100755 index 0000000000..cff10eceea --- /dev/null +++ b/recipe/on_policy_distillation/scripts/run_sampled_token_opd.sh @@ -0,0 +1,72 @@ +#!/usr/bin/env bash + +set -euo pipefail + +SCRIPT_DIR=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd) +REPO_ROOT=$(cd -- "${SCRIPT_DIR}/../../.." && pwd) +source "${SCRIPT_DIR}/launch_teacher_utils.sh" + +export STUDENT_MODEL_PATH=${1:?"Usage: $0 STUDENT_MODEL_PATH TEACHER_MODEL_PATH DATA_PATH"} +export TEACHER_MODEL_PATH=${2:?"Usage: $0 STUDENT_MODEL_PATH TEACHER_MODEL_PATH DATA_PATH"} +export DATA_PATH=${3:?"Usage: $0 STUDENT_MODEL_PATH TEACHER_MODEL_PATH DATA_PATH"} +export OPD_CONFIG_PATH="${OPD_CONFIG_PATH:-recipe/on_policy_distillation/config/rl_dapo_math_opd.py}" +export EVAL_DATA_PATH="${EVAL_DATA_PATH:-}" +export STUDENT_CUDA_VISIBLE_DEVICES="${STUDENT_CUDA_VISIBLE_DEVICES:-0,1,2,3}" +export TEACHER_STARTUP_TIMEOUT_S="${TEACHER_STARTUP_TIMEOUT_S:-1200}" +export PYTHONUNBUFFERED=1 + +USE_SGLANG=${XTUNER_USE_SGLANG:-0} +USE_LMDEPLOY=${XTUNER_USE_LMDEPLOY:-0} +USE_VLLM=${XTUNER_USE_VLLM:-0} + +if [[ "${USE_SGLANG}" == "1" && "${USE_LMDEPLOY}" == "0" && "${USE_VLLM}" == "0" ]]; then + OPD_BACKEND="sglang" +elif [[ "${USE_SGLANG}" == "0" && "${USE_LMDEPLOY}" == "1" && "${USE_VLLM}" == "0" ]]; then + OPD_BACKEND="lmdeploy" +else + echo "Exactly one of XTUNER_USE_SGLANG and XTUNER_USE_LMDEPLOY must be set to 1; XTUNER_USE_VLLM must be 0." >&2 + exit 1 +fi + +export XTUNER_USE_SGLANG="${USE_SGLANG}" +export XTUNER_USE_LMDEPLOY="${USE_LMDEPLOY}" +export XTUNER_USE_VLLM="${USE_VLLM}" + +export WORK_DIR="${WORK_DIR:-${REPO_ROOT}/work_dirs/dapo_math_opd}" +export OPD_CONFIG_FILE="${REPO_ROOT}/${OPD_CONFIG_PATH}" + +TRAINING_STARTED=0 + +cleanup() { + local exit_code=$? + + trap - EXIT INT TERM + + stop_teacher_servers + + if (( TRAINING_STARTED )); then + ray stop --force >/dev/null 2>&1 || true + fi + + exit "${exit_code}" +} + +trap cleanup EXIT +trap "exit 130" INT +trap "exit 143" TERM + +start_single_teacher_server "${OPD_CONFIG_FILE}" "${OPD_BACKEND}" "${WORK_DIR}" +wait_for_teacher_servers "${TEACHER_STARTUP_TIMEOUT_S}" +curl -sS --max-time 10 "${TEACHER_MODEL_INFO_URLS[0]}" +echo +echo "Starting Pure PG-OPD training with student GPUs: ${STUDENT_CUDA_VISIBLE_DEVICES}" + +cd "${REPO_ROOT}" +TRAINING_STARTED=1 +CUDA_VISIBLE_DEVICES="${STUDENT_CUDA_VISIBLE_DEVICES}" \ + bash -o pipefail examples/v1/scripts/run_rl.sh \ + "${OPD_CONFIG_FILE}" \ + "${OPD_BACKEND}" \ + "${STUDENT_MODEL_PATH}" \ + "${DATA_PATH}" \ + "${EVAL_DATA_PATH}" diff --git a/tests/rl/test_on_policy_distillation.py b/tests/rl/test_on_policy_distillation.py new file mode 100644 index 0000000000..fa627bbee4 --- /dev/null +++ b/tests/rl/test_on_policy_distillation.py @@ -0,0 +1,580 @@ +import math +import os +import signal +import subprocess +import time +import unittest +from pathlib import Path +from unittest.mock import patch + +import httpx +import torch + +from recipe.on_policy_distillation.build_teacher_server_commands import ( + build_teacher_server_command, +) +from xtuner.v1.data_proto.rl_data import RolloutState, Status +from xtuner.v1.rl.loss import GRPOLossConfig +from xtuner.v1.rl.on_policy_distillation import ( + OPDConfig, + OPDTeacherConfig, + OPDTeacherLaunchConfig, + TeacherLogprobClient, + apply_opd_kl_to_advantages, +) +from xtuner.v1.rl.utils import find_free_ports + + +REPO_ROOT = Path(__file__).resolve().parents[2] +BASELINE_PATH = os.getenv("XTUNER_OPD_BASELINE") +TEACHER_MODEL_PATH = os.getenv("XTUNER_OPD_TEACHER_MODEL") +STUDENT_MODEL_PATH = os.getenv("XTUNER_OPD_STUDENT_MODEL") +TEACHER_STARTUP_TIMEOUT_S = float(os.getenv("XTUNER_OPD_TEACHER_STARTUP_TIMEOUT_S", "1200")) +TRAINER_CONFIG_PATH = ( + REPO_ROOT / "recipe/on_policy_distillation/config/rl_dapo_math_opd.py" +) + + +def _wait_for_teacher(process: subprocess.Popen, endpoint: str, backend: str) -> None: + deadline = time.monotonic() + TEACHER_STARTUP_TIMEOUT_S + health_path = "health_generate" if backend == "sglang" else "health" + with httpx.Client(timeout=1.0, trust_env=False) as client: + while time.monotonic() < deadline: + return_code = process.poll() + if return_code is not None: + raise RuntimeError(f"Teacher process exited during startup with code {return_code}") + try: + if client.get(f"{endpoint}/{health_path}").status_code == 200: + return + except httpx.RequestError: + pass + time.sleep(1.0) + raise TimeoutError(f"Teacher did not become ready within {TEACHER_STARTUP_TIMEOUT_S} seconds") + + +def _build_debug_rollout_batch(samples: list[dict]) -> list[list[RolloutState]]: + train_batch = [] + for sample in samples: + sample_index = int(sample["sample_index"]) + group_index = sample.get("group_index") + prompt_ids = torch.as_tensor(sample["prompt_token_ids"], dtype=torch.long).tolist() + response_ids = torch.as_tensor(sample["response_token_ids"], dtype=torch.long).tolist() + rollout_logprobs = torch.as_tensor(sample["rollout_log_probs"], dtype=torch.float32).tolist() + teacher_logprobs = torch.as_tensor(sample["teacher_log_probs"], dtype=torch.float32).tolist() + response_mask = torch.as_tensor(sample["loss_mask"], dtype=torch.bool).int().tolist() + response = str(sample["response"]) + + train_batch.append( + [ + RolloutState( + rollout_id=sample_index, + group_id=sample_index if group_index is None else int(group_index), + message=[], + prompt_ids=prompt_ids, + tokens=prompt_ids, + response=response, + response_ids=response_ids, + logprobs=rollout_logprobs, + teacher_tokens=response_ids, + teacher_logprobs=teacher_logprobs, + response_mask=response_mask, + reward={"score": 0.0}, + finish_reason="stop", + status=Status.COMPLETED, + extra_fields={"origin_data_source": "baseline"}, + ) + ] + ) + return train_batch + + +@unittest.skipUnless(BASELINE_PATH, "XTUNER_OPD_BASELINE is required") +class TestPGOPDAdvantage(unittest.TestCase): + @classmethod + def setUpClass(cls) -> None: + super().setUpClass() + baseline = torch.load(BASELINE_PATH, map_location="cpu", weights_only=False) + cls.samples = baseline["samples"] + + def test_compute_advantages_matches_baseline(self) -> None: + config = OPDConfig( + teachers=[OPDTeacherConfig(name="teacher", endpoint="http://unused")], + data_source_teacher_map={"baseline": "teacher"}, + ) + loss_cfg = GRPOLossConfig(policy_loss_cfg={"loss_type": "vanilla"}) + + for sample in self.samples: + with self.subTest(sample_index=sample["sample_index"]): + old_logprobs = torch.as_tensor(sample["old_log_probs"], dtype=torch.float32) + teacher_logprobs = torch.as_tensor(sample["teacher_log_probs"], dtype=torch.float32) + loss_mask = torch.as_tensor(sample["loss_mask"], dtype=torch.float32) + shifted_labels = torch.where( + loss_mask.bool(), + torch.zeros_like(loss_mask, dtype=torch.long), + torch.full_like(loss_mask, -100, dtype=torch.long), + ) + loss_ctx = loss_cfg.build( + { + "shifted_labels": shifted_labels, + "advantages": torch.zeros_like(old_logprobs), + "old_logprobs": old_logprobs, + "teacher_logprobs": teacher_logprobs, + } + ) + assert loss_ctx is not None + apply_opd_kl_to_advantages(loss_ctx, config=config) + actual_advantages = loss_ctx.loss_kwargs.advantages.cpu() + expected_advantages = torch.as_tensor(sample["advantages"], dtype=torch.float32) * loss_mask + try: + torch.testing.assert_close( + actual_advantages, + expected_advantages, + rtol=1e-5, + atol=1e-5, + ) + except AssertionError as error: + mismatch_indices = torch.nonzero( + ~torch.isclose(actual_advantages, expected_advantages, rtol=1e-5, atol=1e-5) + ).flatten() + mismatch_values = "\n".join( + ( + f"index={index}: " + f"actual={actual_advantages[index].item()!r}, " + f"expected={expected_advantages[index].item()!r}, " + f"abs_diff={abs(actual_advantages[index] - expected_advantages[index]).item()!r}" + ) + for index in mismatch_indices.tolist() + ) + raise AssertionError(f"{error}\n\nMismatched advantages:\n{mismatch_values}") from None + + +def _load_trainer_samples(capture_dir: Path) -> dict[int, dict]: + trainer_samples = {} + for capture_file in sorted(capture_dir.glob("rank_*.pt")): + for batch in torch.load(capture_file, map_location="cpu", weights_only=True): + shifted_labels = torch.as_tensor(batch["shifted_labels"], dtype=torch.long).reshape(-1) + old_log_probs = torch.as_tensor(batch["old_log_probs"], dtype=torch.float32).reshape(-1) + advantages = torch.as_tensor(batch["advantages"], dtype=torch.float32).reshape(-1) + boundaries = torch.as_tensor(batch["cu_seq_lens_q"], dtype=torch.long).tolist() + num_padding = int(batch["num_padding"]) + padding_start = shifted_labels.numel() - num_padding + + assert boundaries[len(batch["rollout_ids"])] == padding_start + torch.testing.assert_close( + shifted_labels[padding_start:], + torch.full_like(shifted_labels[padding_start:], -100), + rtol=0, + atol=0, + ) + + for rollout_id, start, end in zip(batch["rollout_ids"], boundaries[:-1], boundaries[1:]): + response_mask = shifted_labels[start:end] != -100 + trainer_samples[int(rollout_id)] = { + "response_token_ids": shifted_labels[start:end][response_mask], + "old_log_probs": old_log_probs[start:end][response_mask], + "advantages": advantages[start:end][response_mask], + } + + return trainer_samples + + +def _run_trainer_once( + baseline_samples: list[dict], + *, + baseline_path: Path, + student_model_path: Path, + debug_rollout_dir: Path, + capture_dir: Path, + run_dir: Path, +) -> dict[int, dict]: + import ray + + from xtuner.v1.rl.trainer.controller import TrainingController + from xtuner.v1.rl.trainer.worker import TrainingWorker, WorkerConfig + from xtuner.v1.train.rl_trainer import RLColocateTrainer + from xtuner.v1.utils import Config + + max_prompt_length = max(len(sample["prompt_token_ids"]) for sample in baseline_samples) + max_response_length = max(len(sample["response_token_ids"]) for sample in baseline_samples) + max_input_length = max(len(sample["token_ids"]) - 1 for sample in baseline_samples) + pack_max_length = math.ceil(max_input_length / 512) * 512 + num_workers = 8 + + class CapturingRLColocateTrainer(RLColocateTrainer): + def _prepare_train_data( + self, + data_groups, + pack_max_length, + raw_rewards_sum=0.0, + raw_rewards_count=0, + ): + data_batches, data_info = super()._prepare_train_data( + data_groups, + pack_max_length, + raw_rewards_sum=raw_rewards_sum, + raw_rewards_count=raw_rewards_count, + ) + rollout_states = (state for group in data_groups for state in group) + for data_batch, rollout_state in zip(data_batches, rollout_states): + data_batch["rollout_id"] = rollout_state.rollout_id + return data_batches, data_info + + class CapturingTrainingController(TrainingController): + def _packing(self, data_batches, pack_max_length, language_cfg): + pack_infos = self._get_pack_infos( + data_batches, + [data["seq_ctx"].input_ids.numel() for data in data_batches], + pack_max_length, + ) + packed_data_batches = super()._packing(data_batches, pack_max_length, language_cfg) + for packed_data, pack_info in zip(packed_data_batches, pack_infos): + packed_data["rollout_ids"] = [data_batches[index]["rollout_id"] for index in pack_info["indices"]] + return packed_data_batches + + class CapturingTrainingWorker(TrainingWorker): + trainer_capture_dir = capture_dir + + def fit(self, data_batches, rollout_idx): + from unittest.mock import patch as mock_patch + + from xtuner.v1.rl.trainer import worker as worker_module + + captured_batches = [ + { + "rollout_ids": data.get("rollout_ids", []), + "cu_seq_lens_q": data["seq_ctx"].cu_seq_lens_q.detach().cpu(), + "num_padding": data["seq_ctx"].num_padding, + "shifted_labels": data["shifted_labels"].detach().cpu().reshape(-1), + } + for data in data_batches + ] + captured_batch_index = 0 + apply_opd = worker_module.apply_opd_kl_to_advantages + + def capture_opd_result(loss_ctx, *, config): + nonlocal captured_batch_index + reverse_kl_sum = apply_opd(loss_ctx, config=config) + captured_batches[captured_batch_index]["old_log_probs"] = ( + loss_ctx.loss_kwargs.old_logprobs.detach().cpu().reshape(-1) + ) + captured_batches[captured_batch_index]["advantages"] = ( + loss_ctx.loss_kwargs.advantages.detach().cpu().reshape(-1) + ) + captured_batch_index += 1 + return reverse_kl_sum + + with mock_patch.object(worker_module, "apply_opd_kl_to_advantages", capture_opd_result): + worker_log_item = TrainingWorker.fit(self, data_batches, rollout_idx) + + torch.save(captured_batches, self.trainer_capture_dir / f"rank_{self.rank}.pt") + return worker_log_item + + def build_capturing_training_workers(self, placement_group): + from xtuner.v1.rl.utils import AutoAcceleratorWorkers + + capturing_worker_cls = ray.remote( + runtime_env={ + "env_vars": { + "RAY_EXPERIMENTAL_NOSET_CUDA_VISIBLE_DEVICES": "1", + "RAY_EXPERIMENTAL_NOSET_ASCEND_RT_VISIBLE_DEVICES": "1", + "HCCL_NPU_SOCKET_PORT_RANGE": "auto", + } + } + )(CapturingTrainingWorker) + train_workers, _ = AutoAcceleratorWorkers.from_placement_group( + capturing_worker_cls, + self, + placement_group, + ) + ray.wait([worker.ready.remote() for worker in train_workers]) + return CapturingTrainingController(workers=train_workers) + + trainer_environment = { + "WORK_DIR": str(run_dir / "trainer"), + "MODEL_PATH": str(student_model_path), + "TEACHER_MODEL_PATH": os.getenv("XTUNER_OPD_TEACHER_MODEL", str(student_model_path)), + "DATA_PATH": os.getenv("XTUNER_OPD_DATA_PATH", str(baseline_path)), + "WORLD_SIZE": "1", + "ONLY_CALC_MISMATCH_RATIO": "1", + "XTUNER_DETERMINISTIC": "true", + "XTUNER_USE_FA3": os.getenv("XTUNER_OPD_USE_FA3", "1"), + "TOKENIZERS_PARALLELISM": "false", + "PYTHONUNBUFFERED": "1", + } + + with patch.dict(os.environ, trainer_environment, clear=False): + try: + ray.init( + num_cpus=12 * num_workers + 10, + num_gpus=num_workers, + include_dashboard=False, + _temp_dir="/dev/shm/xtuner-opd-old-logprobs", + ) + cfg = Config.fromfile(TRAINER_CONFIG_PATH) + cfg.trainer.resources.num_workers = num_workers + cfg.trainer.total_train_steps = 1 + cfg.trainer.train_batch_size = len(baseline_samples) + cfg.trainer.train_worker_cfg.optimizer_steps = 1 + cfg.trainer.train_worker_cfg.pack_max_length = pack_max_length + cfg.trainer.debug_train = True + cfg.trainer.debug_rollout_dir = debug_rollout_dir + with ( + patch("xtuner.v1.train.rl_trainer.XTUNER_DETERMINISTIC", True), + patch( + "xtuner.v1.train.rl_trainer.RLColocateTrainer", + CapturingRLColocateTrainer, + ), + patch.object(WorkerConfig, "build", build_capturing_training_workers), + ): + trainer = cfg.trainer.build() + trainer.fit() + finally: + if ray.is_initialized(): + ray.shutdown() + + return _load_trainer_samples(capture_dir) + + +@unittest.skipUnless( + BASELINE_PATH and STUDENT_MODEL_PATH, + "XTUNER_OPD_BASELINE and XTUNER_OPD_STUDENT_MODEL are required", +) +class TestPGOPDOldLogprobs(unittest.TestCase): + def test_trainer_old_logprobs_and_advantage_error_propagation(self) -> None: + baseline_path = Path(str(BASELINE_PATH)).expanduser().resolve() + student_model_path = Path(str(STUDENT_MODEL_PATH)).expanduser().resolve() + baseline = torch.load(baseline_path, map_location="cpu", weights_only=True) + baseline_samples = baseline["samples"] + + work_root = Path( + os.getenv( + "XTUNER_OPD_OLD_LOGPROB_WORK_DIR", + str(REPO_ROOT / "work_dirs/test_pg_opd_old_logprobs"), + ) + ).expanduser() + run_dir = work_root / f"run_{time.strftime('%Y%m%d_%H%M%S')}_{os.getpid()}" + debug_rollout_dir = run_dir / "debug_rollout" + capture_dir = run_dir / "trainer_capture" + debug_rollout_dir.mkdir(parents=True) + capture_dir.mkdir() + + train_batch = _build_debug_rollout_batch(baseline_samples) + torch.save(train_batch, debug_rollout_dir / "debug_rollout_1.pt") + trainer_samples_by_rollout_id = _run_trainer_once( + baseline_samples, + baseline_path=baseline_path, + student_model_path=student_model_path, + debug_rollout_dir=debug_rollout_dir, + capture_dir=capture_dir, + run_dir=run_dir, + ) + + result_path = run_dir / "result.pt" + result_samples = [] + + for baseline_sample in baseline_samples: + rollout_id = int(baseline_sample["sample_index"]) + trainer_sample = trainer_samples_by_rollout_id[rollout_id] + loss_mask = torch.as_tensor(baseline_sample["loss_mask"], dtype=torch.bool) + response_token_ids = torch.as_tensor(baseline_sample["response_token_ids"], dtype=torch.long)[loss_mask] + torch.testing.assert_close( + trainer_sample["response_token_ids"], + response_token_ids, + rtol=0, + atol=0, + ) + baseline_old_log_probs = torch.as_tensor( + baseline_sample["old_log_probs"], + dtype=torch.float32, + )[loss_mask] + trainer_old_log_probs = trainer_sample["old_log_probs"] + baseline_advantages = torch.as_tensor( + baseline_sample["advantages"], + dtype=torch.float32, + )[loss_mask] + trainer_advantages = trainer_sample["advantages"] + + old_logprob_error = trainer_old_log_probs - baseline_old_log_probs + advantage_error = trainer_advantages - baseline_advantages + propagation_residual = advantage_error + old_logprob_error + sample_num_tokens = old_logprob_error.numel() + sample_summary = { + "num_tokens": sample_num_tokens, + "old_logprobs_mean_abs_error": old_logprob_error.abs().mean().item(), + "old_logprobs_max_abs_error": old_logprob_error.abs().max().item(), + "advantages_mean_abs_error": advantage_error.abs().mean().item(), + "advantages_max_abs_error": advantage_error.abs().max().item(), + "propagation_mean_abs_error": propagation_residual.abs().mean().item(), + "propagation_max_abs_error": propagation_residual.abs().max().item(), + } + result_samples.append( + { + "sample_index": rollout_id, + "response_token_ids": response_token_ids, + "baseline_old_log_probs": baseline_old_log_probs, + "trainer_old_log_probs": trainer_old_log_probs, + "old_logprob_error": old_logprob_error, + "baseline_advantages": baseline_advantages, + "trainer_advantages": trainer_advantages, + "advantage_error": advantage_error, + "propagation_residual": propagation_residual, + "summary": sample_summary, + } + ) + + with self.subTest(sample_index=rollout_id): + torch.testing.assert_close( + advantage_error, + -old_logprob_error, + rtol=1e-5, + atol=1e-5, + msg=f"Result: {result_path}", + ) + + global_result = { + "response_token_ids": torch.cat([sample["response_token_ids"] for sample in result_samples]), + "baseline_old_log_probs": torch.cat([sample["baseline_old_log_probs"] for sample in result_samples]), + "trainer_old_log_probs": torch.cat([sample["trainer_old_log_probs"] for sample in result_samples]), + "old_logprob_error": torch.cat([sample["old_logprob_error"] for sample in result_samples]), + "baseline_advantages": torch.cat([sample["baseline_advantages"] for sample in result_samples]), + "trainer_advantages": torch.cat([sample["trainer_advantages"] for sample in result_samples]), + "advantage_error": torch.cat([sample["advantage_error"] for sample in result_samples]), + "propagation_residual": torch.cat([sample["propagation_residual"] for sample in result_samples]), + } + summary = { + "num_samples": len(result_samples), + "num_tokens": global_result["old_logprob_error"].numel(), + "old_logprobs_mean_abs_error": global_result["old_logprob_error"].abs().mean().item(), + "old_logprobs_max_abs_error": global_result["old_logprob_error"].abs().max().item(), + "advantages_mean_abs_error": global_result["advantage_error"].abs().mean().item(), + "advantages_max_abs_error": global_result["advantage_error"].abs().max().item(), + "propagation_mean_abs_error": global_result["propagation_residual"].abs().mean().item(), + "propagation_max_abs_error": global_result["propagation_residual"].abs().max().item(), + } + torch.save( + { + "summary": summary, + "global": global_result, + "samples": result_samples, + }, + result_path, + ) + + with self.subTest(scope="global"): + torch.testing.assert_close( + global_result["advantage_error"], + -global_result["old_logprob_error"], + rtol=1e-5, + atol=1e-5, + msg=f"Result: {result_path}", + ) + + print( + f"old_logprobs: mean_abs={summary['old_logprobs_mean_abs_error']}, " + f"max_abs={summary['old_logprobs_max_abs_error']}\n" + f"advantages: mean_abs={summary['advantages_mean_abs_error']}, " + f"max_abs={summary['advantages_max_abs_error']}\n" + f"propagation: mean_abs={summary['propagation_mean_abs_error']}, " + f"max_abs={summary['propagation_max_abs_error']}\n" + f"result: {result_path}" + ) + + +@unittest.skipUnless( + BASELINE_PATH and TEACHER_MODEL_PATH, + "XTUNER_OPD_BASELINE and XTUNER_OPD_TEACHER_MODEL are required", +) +class TestTeacherLogprobClient(unittest.IsolatedAsyncioTestCase): + @classmethod + def setUpClass(cls) -> None: + super().setUpClass() + baseline = torch.load(BASELINE_PATH, map_location="cpu", weights_only=False) + cls.samples = baseline["samples"] + port = find_free_ports()[0] + cls.teacher_endpoint = f"http://127.0.0.1:{port}" + cls.teacher_backend = TeacherLogprobClient._resolve_backend_from_env() + teacher_env = os.environ.copy() + teacher_command = build_teacher_server_command( + OPDTeacherConfig( + name="teacher", + endpoint=cls.teacher_endpoint, + launch_config=OPDTeacherLaunchConfig( + model_path=str(TEACHER_MODEL_PATH), + cuda_visible_devices=teacher_env.get("CUDA_VISIBLE_DEVICES") + or "7", + ), + ), + cls.teacher_backend, + ) + if not teacher_command: + raise ValueError("Teacher must define launch_config for local startup") + cls.teacher_process = subprocess.Popen( + teacher_command, + cwd=REPO_ROOT, + env=teacher_env, + start_new_session=True, + ) + cls.addClassCleanup(cls._stop_teacher) + _wait_for_teacher(cls.teacher_process, cls.teacher_endpoint, cls.teacher_backend) + + @classmethod + def _stop_teacher(cls) -> None: + if cls.teacher_process.poll() is not None: + return + os.killpg(cls.teacher_process.pid, signal.SIGTERM) + try: + cls.teacher_process.wait(timeout=30) + except subprocess.TimeoutExpired: + os.killpg(cls.teacher_process.pid, signal.SIGKILL) + cls.teacher_process.wait() + + @unittest.skipUnless(os.getenv("XTUNER_USE_SGLANG", "0") == "1", "XTUNER_USE_SGLANG=1 is required") + async def test_compute_logprobs_with_sglang_matches_baseline(self) -> None: + self.assertEqual(self.teacher_backend, "sglang") + await self._assert_compute_logprobs_matches_baseline() + + @unittest.skipUnless(os.getenv("XTUNER_USE_LMDEPLOY", "0") == "1", "XTUNER_USE_LMDEPLOY=1 is required") + async def test_compute_logprobs_with_lmdeploy_matches_baseline(self) -> None: + self.assertEqual(self.teacher_backend, "lmdeploy") + await self._assert_compute_logprobs_matches_baseline() + + async def _assert_compute_logprobs_matches_baseline(self) -> None: + client = TeacherLogprobClient(OPDTeacherConfig(name="teacher", endpoint=self.teacher_endpoint)) + self.addAsyncCleanup(client._client.aclose) + + for sample in self.samples: + with self.subTest(sample_index=sample["sample_index"]): + prompt_ids = torch.as_tensor(sample["prompt_token_ids"], dtype=torch.long).tolist() + response_ids = torch.as_tensor(sample["response_token_ids"], dtype=torch.long).tolist() + expected_logprobs = torch.as_tensor(sample["teacher_log_probs"], dtype=torch.float32) + state = RolloutState( + rollout_id=int(sample["sample_index"]), + group_id=int(sample["group_index"]), + message=[], + prompt_ids=prompt_ids, + tokens=prompt_ids, + response="", + response_ids=response_ids, + status=Status.COMPLETED, + ) + + result = await client.compute_logprobs(state) + + self.assertEqual(result.status, Status.COMPLETED, result.error_msg) + self.assertEqual(result.teacher_tokens, response_ids) + actual_logprobs = torch.tensor(result.teacher_logprobs, dtype=torch.float32) + self.assertEqual(actual_logprobs.shape, expected_logprobs.shape) + absolute_errors = torch.abs(actual_logprobs - expected_logprobs) + mae = absolute_errors.mean().item() + self.assertLessEqual( + mae, + 0.05, + ( + f"Teacher logprob MAE {mae:.8f} exceeds threshold 0.05; " + f"max_abs_error={absolute_errors.max().item():.8f}" + ), + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/rl/test_prepare_train_data.py b/tests/rl/test_prepare_train_data.py index 029a3f0306..efb32d61c5 100644 --- a/tests/rl/test_prepare_train_data.py +++ b/tests/rl/test_prepare_train_data.py @@ -7,9 +7,6 @@ - VLM 样本使用 train_prompt_ids,并保留 multimodal 训练字段。 - 无效 rollout group 会被跳过。 - 缺失 reward、logprob/mask 长度不一致、pack_max_length 过小时 fail fast。 - -注意:当前训练 contract 中 data_dict["advantage"] 比 shifted_labels 多 1 个元素; -metric 统计使用 actual_advantages[:-1],测试会显式固定这个行为。 """ import unittest @@ -36,6 +33,7 @@ class TestPrepareTrainData(unittest.TestCase): def _build_trainer(self, advantages: list[float]): trainer = BaseRLTrainer.__new__(BaseRLTrainer) trainer._advantage_estimator = _FakeAdvantageEstimator(advantages) + trainer._opd_config = None trainer.tokenizer = MagicMock(return_value={"input_ids": torch.tensor([[999]])}) trainer.logger = MagicMock() return trainer @@ -102,8 +100,8 @@ def test_text_path_builds_shifted_training_tensors(self): batch["rollout_logprobs"], torch.tensor([[0.0, 0.0, 0.1, 0.2, 0.3]], dtype=torch.float32), ) - self.assertEqual(batch["advantage"], [1.5, 1.5, 1.5, 1.5, 0.0, 1.5]) - self.assertEqual(len(batch["advantage"]), batch["shifted_labels"].numel() + 1) + self.assertEqual(batch["advantage"], [0.0, 0.0, 1.5, 0.0, 1.5]) + self.assertEqual(len(batch["advantage"]), batch["shifted_labels"].numel()) self.assertIs(batch["seq_ctx"].rollout_routed_experts, routed_experts) self.assertEqual(info["training_samples"], 1) self.assertEqual(info["training_tokens"], 5) @@ -120,8 +118,8 @@ def test_multi_sample_group_uses_each_sample_reward_and_advantage(self): data_batches, info = self._prepare(trainer, [[first, second]]) self.assertEqual(len(data_batches), 2) - self.assertEqual(data_batches[0]["advantage"], [1.5, 1.5, 1.5, 1.5, 1.5]) - self.assertEqual(data_batches[1]["advantage"], [-2.0, -2.0, -2.0, -2.0, -2.0]) + self.assertEqual(data_batches[0]["advantage"], [0.0, 0.0, 1.5, 1.5]) + self.assertEqual(data_batches[1]["advantage"], [0.0, 0.0, -2.0, -2.0]) self.assertEqual(info["batch_size"], 2) self.assertEqual(info["rewards/min"], -1.0) self.assertEqual(info["rewards/max"], 3.0) diff --git a/tests/rl/test_producer.py b/tests/rl/test_producer.py index 3514542f89..c1c187779b 100644 --- a/tests/rl/test_producer.py +++ b/tests/rl/test_producer.py @@ -23,6 +23,7 @@ from unittest.mock import AsyncMock, MagicMock from xtuner.v1.data_proto.rl_data import RolloutState, Status, discard_rollout_state +from xtuner.v1.rl.agent_loop import AgentLoop from xtuner.v1.rl.agent_loop_manager import ( AsyncProduceStrategyConfig, DisaggAsyncProduceStrategyConfig, @@ -117,6 +118,8 @@ async def mock_gen(rs, **kwargs): return rs mock_agent_loop.generate_group = mock_gen + mock_agent_loop.teacher_clients = {} + mock_agent_loop.collect_rollout_group = AgentLoop.collect_rollout_group.__get__(mock_agent_loop) return mock_agent_loop def _build_context( @@ -130,6 +133,7 @@ def _build_context( train_step: int = 0, model_step: int = 0, progress: ProduceProgress | None = None, + is_valid_sample_fn=None, ) -> ProduceContext: # 测试只走新的 ProduceContext 入口,不再覆盖旧散装参数兼容逻辑。 if progress is None: @@ -143,7 +147,7 @@ def _build_context( train_step=train_step, model_step=model_step, progress=progress, - is_valid_sample_fn=strategy.is_valid_sample_fn, + is_valid_sample_fn=is_valid_sample_fn, stale_threshold=getattr(strategy, "stale_threshold", None), ) @@ -193,7 +197,6 @@ def _build_disagg_context( update_event=update_event, model_step=model_step, progress=progress, - is_valid_sample_fn=strategy.is_valid_sample_fn, stale_threshold=getattr(strategy, "stale_threshold", None), ) @@ -280,8 +283,8 @@ async def test_discard_rollout_state_keeps_required_fields_valid(self): self.assertIsNone(discarded.routed_experts) self.assertEqual(discarded.extra_fields, {}) - async def test_put_generated_group_only_validates_completed_group(self): - # 验证 ProduceContext 只对 completed group 执行业务过滤,aborted group 保持可重试状态。 + async def test_collection_filters_only_completed_group(self): + # 过滤由 AgentLoop collection 执行;Producer 只处理返回状态和数据所有权。 task_name = "test_valid_completed_only" valid_checked_statuses = [] @@ -289,16 +292,18 @@ def is_valid_sample_fn(samples): valid_checked_statuses.append([sample.status for sample in samples]) return False - strategy = SyncProduceStrategyConfig(is_valid_sample_fn=is_valid_sample_fn).build() + strategy = SyncProduceStrategyConfig().build() ctx = self._build_context( strategy, task_name, self._build_agent_loop(), self._build_sampler(), batch_size=1, + is_valid_sample_fn=is_valid_sample_fn, ) completed_group = [make_rollout_state(1, status=Status.COMPLETED)] + completed_group = await ctx.collect_rollout_group(completed_group) self.assertFalse(await ctx.put_generated_group(completed_group)) self.assertIsNone(completed_group[0].uid) @@ -339,19 +344,21 @@ async def test_put_generated_group_records_raw_rewards_before_filtering(self): def is_valid_sample_fn(samples): return False - strategy = SyncProduceStrategyConfig(is_valid_sample_fn=is_valid_sample_fn).build() + strategy = SyncProduceStrategyConfig().build() ctx = self._build_context( strategy, task_name, self._build_agent_loop(), self._build_sampler(), batch_size=1, + is_valid_sample_fn=is_valid_sample_fn, ) completed_group = [ make_rollout_state(1, status=Status.COMPLETED, reward_score=0.25), make_rollout_state(2, status=Status.COMPLETED, reward_score=0.75), ] + completed_group = await ctx.collect_rollout_group(completed_group) self.assertFalse(await ctx.put_generated_group(completed_group)) self.assertTrue(all(item.uid is None for item in completed_group)) @@ -425,7 +432,7 @@ async def mock_gen(rs, **kwargs): mock_agent_loop = self._build_agent_loop() mock_agent_loop.generate_group = mock_gen - strategy = SyncProduceStrategyConfig(is_valid_sample_fn=is_valid_sample_fn).build() + strategy = SyncProduceStrategyConfig().build() sampler = self._build_sampler() ctx = self._build_context( strategy, @@ -436,6 +443,7 @@ async def mock_gen(rs, **kwargs): train_step=4, model_step=3, progress=self._build_progress(task_name, target=2), + is_valid_sample_fn=is_valid_sample_fn, ) await strategy.produce_batch(ctx) @@ -465,7 +473,7 @@ async def mock_gen(rs, **kwargs): r.status = Status.COMPLETED return rs - mock_agent_loop.generate_group = mock_gen + mock_agent_loop.collect_rollout_group = mock_gen sampler_cfg = SamplerConfig.model_construct(dataloader_cfg=self.mock_dataloader_cfg) produce_strategy_cfg = AsyncProduceStrategyConfig(over_sample_threshold=1) diff --git a/tests/rl/test_rl_colocate_trainer.py b/tests/rl/test_rl_colocate_trainer.py index 69f902d6d0..3328ec90c4 100644 --- a/tests/rl/test_rl_colocate_trainer.py +++ b/tests/rl/test_rl_colocate_trainer.py @@ -98,7 +98,7 @@ async def generate_group(rollout_states, **kwargs): state.response_model_steps = [model_step] return rollout_states - agent_loop.generate_group = generate_group + agent_loop.collect_rollout_group = generate_group return agent_loop diff --git a/xtuner/v1/data_proto/rl_data.py b/xtuner/v1/data_proto/rl_data.py index 7db20f2940..fb22c635be 100644 --- a/xtuner/v1/data_proto/rl_data.py +++ b/xtuner/v1/data_proto/rl_data.py @@ -109,6 +109,8 @@ class RolloutState(BaseModel): tool_calls: list[RolloutToolCall] | None = None response_ids: list[int] | None = None logprobs: list[float] | None = None + teacher_tokens: list[int] | None = None + teacher_logprobs: list[float] | None = None routed_experts: np.ndarray | RayObjectRef | list[RayObjectRef] | None = None finish_reason: str | None = None # response_mask: 记录response_ids中哪个token算loss, 与response_ids长度相同,每轮rollout在 agent_loop.generate 中覆盖写 @@ -244,6 +246,8 @@ def reset_rollout_response(rollout_state: RolloutState) -> RolloutState: rollout_state.response = "" rollout_state.response_ids = [] rollout_state.logprobs = [] + rollout_state.teacher_tokens = None + rollout_state.teacher_logprobs = None rollout_state.routed_experts = None rollout_state.finish_reason = None rollout_state.response_mask = [] diff --git a/xtuner/v1/rl/agent_loop/agent_loop.py b/xtuner/v1/rl/agent_loop/agent_loop.py index 6a69d0f80c..4fa21fdf0f 100644 --- a/xtuner/v1/rl/agent_loop/agent_loop.py +++ b/xtuner/v1/rl/agent_loop/agent_loop.py @@ -2,6 +2,7 @@ import asyncio from abc import ABC, abstractmethod +from collections.abc import Callable from typing import Any, TypeAlias, cast, overload import ray @@ -9,8 +10,14 @@ from ray.actor import ActorClass, ActorProxy from ray.util.placement_group import PlacementGroup -from xtuner.v1.data_proto.rl_data import RolloutState, SampleParams, Status +from xtuner.v1.data_proto.rl_data import RolloutState, SampleParams, Status, get_group_status from xtuner.v1.rl.judger import Judger +from xtuner.v1.rl.on_policy_distillation import ( + OPDConfig, + TeacherLogprobClient, + route_teacher_client, + validate_opd_sample_params, +) from xtuner.v1.rl.rollout import RolloutController from xtuner.v1.rl.rollout.constants import AGENT_LOOP_RAY_GENERATE_MAX_CONCURRENCY from xtuner.v1.rl.trace.rollout_api import ( @@ -39,13 +46,22 @@ class AgentLoopConfig(ABC, BaseModel): enable_batch_judge: bool = False requires_rollout_proxy: bool = False - def build(self, rollout_controller, judger: Judger | None = None, logger=None) -> AgentLoopSpec: + def build( + self, + rollout_controller, + judger: Judger | None = None, + logger=None, + *, + opd_config: OPDConfig | None = None, + ) -> AgentLoopSpec: if self.cpu_resources is None: - return self.build_local( + agent_loop = self.build_local( rollout_controller=rollout_controller, judger=judger, logger=logger, ) + agent_loop.configure_opd(opd_config) + return agent_loop concurrency = AGENT_LOOP_RAY_GENERATE_MAX_CONCURRENCY @@ -61,6 +77,7 @@ def build(self, rollout_controller, judger: Judger | None = None, logger=None) - concurrency=concurrency, judger=judger, logger=logger, + opd_config=opd_config, ) return self._build_ray_actor( rollout_controller=rollout_controller, @@ -68,6 +85,7 @@ def build(self, rollout_controller, judger: Judger | None = None, logger=None) - concurrency=concurrency, judger=judger, logger=logger, + opd_config=opd_config, ) @abstractmethod @@ -86,6 +104,7 @@ def _build_ray_actor( pg: PlacementGroup | None = None, judger: Judger | None = None, logger=None, + opd_config: OPDConfig | None = None, ) -> RayAgentLoopProxy: ray_agent_loop = ray.remote( concurrency_groups={ @@ -104,6 +123,7 @@ def _build_ray_actor( actor_num_cpus=cpu_resources.num_cpus_per_worker, actor_memory=cpu_resources.cpu_memory_per_worker, capture_child_tasks=True, + opd_config=opd_config, ), ) @@ -116,6 +136,7 @@ def _build_ray_actors( judger: Judger | None = None, logger=None, start_bundle_idx: int = 0, + opd_config: OPDConfig | None = None, ) -> list[RayAgentLoopProxy]: ray_agent_loop = ray.remote( concurrency_groups={ @@ -135,6 +156,7 @@ def _build_ray_actors( actor_num_cpus_per_worker=cpu_resources.num_cpus_per_worker, actor_memory_per_worker=cpu_resources.cpu_memory_per_worker, capture_child_tasks=True, + opd_config=opd_config, ), ) @@ -147,6 +169,7 @@ def _build_router( judger: Judger | None = None, logger=None, start_bundle_idx: int = 0, + opd_config: OPDConfig | None = None, ) -> RouterAgentLoop: return RouterAgentLoop( workers=self._build_ray_actors( @@ -157,6 +180,7 @@ def _build_router( judger=judger, logger=logger, start_bundle_idx=start_bundle_idx, + opd_config=opd_config, ), rollout_ctl=rollout_controller, ) @@ -184,6 +208,16 @@ def __init__( else: self.logger = logger self._judger_pause_event = asyncio.Event() + self.teacher_clients: dict[str, TeacherLogprobClient] = {} + self.data_source_teacher_map: dict[str, str] = {} + + def configure_opd(self, opd_config: OPDConfig | None) -> None: + if opd_config is None: + return + + validate_opd_sample_params(self.sample_params) + self.teacher_clients = {teacher.name: TeacherLogprobClient(teacher) for teacher in opd_config.teachers} + self.data_source_teacher_map = dict(opd_config.data_source_teacher_map) @abstractmethod async def generate_sample(self, rollout_state: RolloutState, **kwargs) -> RolloutState: ... @@ -201,6 +235,48 @@ async def generate_group(self, rollout_state: list[RolloutState], **kwargs) -> l group_samples = await self.run_judger(group_samples) return group_samples + async def collect_rollout_group( + self, + rollout_state: list[RolloutState], + *, + is_valid_sample_func: Callable[[list[RolloutState]], bool] | None = None, + **kwargs, + ) -> list[RolloutState]: + if is_valid_sample_func is None and self.teacher_clients: + teacher = route_teacher_client( + rollout_state[0], + data_source_teacher_map=self.data_source_teacher_map, + teacher_clients=self.teacher_clients, + ) + + async def generate_and_score(state: RolloutState) -> RolloutState: + state.sample_params = self.sample_params + state = await self.generate_sample(state, **kwargs) + if state.status == Status.COMPLETED: + state = await teacher.compute_logprobs(state) + return state + + group = list(await asyncio.gather(*(create_task(generate_and_score(state)) for state in rollout_state))) + if self.judger is not None and self.enable_batch_judge and get_group_status(group) == Status.COMPLETED: + group = await self.run_judger(group) + return group + + group = await self.generate_group(rollout_state, **kwargs) + if get_group_status(group) != Status.COMPLETED: + return group + if is_valid_sample_func is not None and not is_valid_sample_func(group): + for state in group: + state.status = Status.FILTERED + return group + if self.teacher_clients: + teacher = route_teacher_client( + group[0], + data_source_teacher_map=self.data_source_teacher_map, + teacher_clients=self.teacher_clients, + ) + group = list(await asyncio.gather(*(create_task(teacher.compute_logprobs(state)) for state in group))) + return group + @overload async def run_judger(self, rollout_state: RolloutState) -> RolloutState: ... @@ -287,6 +363,13 @@ async def generate_group(self, rollout_state: list[RolloutState], **kwargs) -> l finally: await self._release_worker(worker) + async def collect_rollout_group(self, rollout_state: list[RolloutState], **kwargs) -> list[RolloutState]: + worker = await self._pick_worker() + try: + return await worker.collect_rollout_group.remote(rollout_state, **kwargs) + finally: + await self._release_worker(worker) + def get_worker_status(self) -> dict[str, int]: return {str(worker): load for worker, load in self._worker_loads.items()} @@ -314,12 +397,15 @@ def __init__( rollout_controller: RolloutController, judger: Judger | None = None, logger=None, + *, + opd_config: OPDConfig | None = None, ): self.agent_loop = agent_loop_config.build_local( rollout_controller=rollout_controller, judger=judger, logger=logger, ) + self.agent_loop.configure_opd(opd_config) @ray_method(concurrency_group=AGENT_LOOP_CONCURRENCY_GROUP_GENERATE) async def generate_sample(self, rollout_state: RolloutState, **kwargs) -> RolloutState: @@ -329,6 +415,10 @@ async def generate_sample(self, rollout_state: RolloutState, **kwargs) -> Rollou async def generate_group(self, rollout_state: list[RolloutState], **kwargs) -> list[RolloutState]: return await self.agent_loop.generate_group(rollout_state, **kwargs) + @ray_method(concurrency_group=AGENT_LOOP_CONCURRENCY_GROUP_GENERATE) + async def collect_rollout_group(self, rollout_state: list[RolloutState], **kwargs) -> list[RolloutState]: + return await self.agent_loop.collect_rollout_group(rollout_state, **kwargs) + @ray_method async def get_rollout_ctl(self): return self.agent_loop.rollout_ctl diff --git a/xtuner/v1/rl/agent_loop_manager/agent_loop_manager.py b/xtuner/v1/rl/agent_loop_manager/agent_loop_manager.py index 2b45a6ca15..57f0e92f4d 100644 --- a/xtuner/v1/rl/agent_loop_manager/agent_loop_manager.py +++ b/xtuner/v1/rl/agent_loop_manager/agent_loop_manager.py @@ -10,6 +10,7 @@ from xtuner.v1.data_proto.rl_data import Status from xtuner.v1.rl.agent_loop import AgentLoopConfig from xtuner.v1.rl.judger import ComposedJudgerConfig, JudgerConfig, build_judger +from xtuner.v1.rl.on_policy_distillation import OPDConfig from xtuner.v1.rl.replay_buffer import ReplayBuffer from xtuner.v1.rl.rollout import RolloutController from xtuner.v1.utils import get_logger @@ -18,6 +19,7 @@ _MANAGER_STATE_PATH, _STATUS_POLL_INTERVAL_S, _TASK_CHECKPOINT_DIR, + IsValidSampleFn, ProduceBatchResult, _TaskRunner, _TaskSamplerView, @@ -55,6 +57,8 @@ class TaskSpecConfig(BaseModel): judger_config (JudgerConfig | ComposedJudgerConfig | None): Optional judger configuration used to score generated samples. Defaults to None. + filter_func (IsValidSampleFn | None): Optional group filter applied by the + agent loop after generation. Defaults to None. produce_strategy_config (ProduceStrategyConfig): Strategy used to produce rollout samples. Defaults to ``SyncProduceStrategyConfig``. sampler_config (SamplerConfig): Dataset sampler configuration for this @@ -82,6 +86,7 @@ class TaskSpecConfig(BaseModel): weight: float = Field(default=1.0, ge=0.0) agent_loop_config: AgentLoopConfig judger_config: JudgerConfig | ComposedJudgerConfig | None = None + filter_func: IsValidSampleFn | None = None produce_strategy_config: ProduceStrategyConfig = SyncProduceStrategyConfig() sampler_config: SamplerConfig @@ -126,6 +131,7 @@ def build( replay_buffer: ReplayBuffer, logger=None, sync_weights_interval: int = 1, + opd_config: OPDConfig | None = None, ) -> "AgentLoopManager": tasks = self.tasks if isinstance(self.tasks, list) else [self.tasks] if not tasks: @@ -142,6 +148,7 @@ def build( rollout_controller=rollout_controller, judger=build_judger(task_cfg.judger_config) if task_cfg.judger_config is not None else None, logger=logger, + opd_config=opd_config, ) produce_strategy = task_cfg.produce_strategy_config.build( sync_weights_interval=sync_weights_interval, @@ -154,6 +161,7 @@ def build( agent_loop=agent_loop, produce_strategy=produce_strategy, sampler=sampler, + is_valid_sample_fn=task_cfg.filter_func, weight=task_cfg.weight, order=order, ) diff --git a/xtuner/v1/rl/agent_loop_manager/disagg_agent_loop_manager.py b/xtuner/v1/rl/agent_loop_manager/disagg_agent_loop_manager.py index e73dc63da3..e070b0026c 100644 --- a/xtuner/v1/rl/agent_loop_manager/disagg_agent_loop_manager.py +++ b/xtuner/v1/rl/agent_loop_manager/disagg_agent_loop_manager.py @@ -11,6 +11,7 @@ from xtuner.v1.data_proto.rl_data import Status from xtuner.v1.rl.agent_loop import AgentLoopConfig from xtuner.v1.rl.judger import ComposedJudgerConfig, JudgerConfig, build_judger +from xtuner.v1.rl.on_policy_distillation import OPDConfig from xtuner.v1.rl.replay_buffer import ReplayBuffer from xtuner.v1.rl.rollout import RolloutController from xtuner.v1.utils import get_logger @@ -27,6 +28,7 @@ _MANAGER_STATE_PATH, _STATUS_POLL_INTERVAL_S, _TASK_CHECKPOINT_DIR, + IsValidSampleFn, ProduceBatchResult, ProduceBatchStatus, _TaskRunner, @@ -50,6 +52,7 @@ class DisaggTaskSpecConfig(BaseModel): weight: float = Field(default=1.0, ge=0.0) agent_loop_config: AgentLoopConfig judger_config: JudgerConfig | ComposedJudgerConfig | None = None + filter_func: IsValidSampleFn | None = None produce_strategy_config: DisaggProduceStrategyConfig = DisaggAsyncProduceStrategyConfig() sampler_config: SamplerConfig @@ -68,6 +71,7 @@ def build( replay_buffer: ReplayBuffer, logger=None, sync_weights_interval: int = 1, + opd_config: OPDConfig | None = None, ) -> "DisaggAgentLoopManager": tasks = self.tasks if isinstance(self.tasks, list) else [self.tasks] if not tasks: @@ -84,6 +88,7 @@ def build( rollout_controller=rollout_controller, judger=build_judger(task_cfg.judger_config) if task_cfg.judger_config is not None else None, logger=logger, + opd_config=opd_config, ) produce_strategy = task_cfg.produce_strategy_config.build( sync_weights_interval=sync_weights_interval, @@ -96,6 +101,7 @@ def build( agent_loop=agent_loop, produce_strategy=produce_strategy, sampler=sampler, + is_valid_sample_fn=task_cfg.filter_func, weight=task_cfg.weight, order=order, ) diff --git a/xtuner/v1/rl/agent_loop_manager/disagg_producer.py b/xtuner/v1/rl/agent_loop_manager/disagg_producer.py index d75c34b7c1..2bedd94e9f 100644 --- a/xtuner/v1/rl/agent_loop_manager/disagg_producer.py +++ b/xtuner/v1/rl/agent_loop_manager/disagg_producer.py @@ -13,14 +13,12 @@ from .produce_utils import ( PERIODIC_ABORT_INTERVAL_S, BaseProduceContext, - IsValidSampleFn, ProduceBatchStatus, ShouldContinueFn, _PendingTasks, _ProgressDisplayer, _put_claimed_tasks, calculate_stale_threshold, - default_is_valid_sample_fn, default_should_continue_fn, pause_pending_tasks, ) @@ -230,7 +228,6 @@ class DisaggProduceStrategyConfig(ABC, BaseModel): """非共卡后台 producer strategy 配置。""" model_config = ConfigDict(extra="forbid", arbitrary_types_allowed=True) - is_valid_sample_fn: IsValidSampleFn = default_is_valid_sample_fn should_continue_fn: ShouldContinueFn = default_should_continue_fn @abstractmethod @@ -266,7 +263,6 @@ def build( max_staleness=self.max_staleness, sync_weights_interval=sync_weights_interval, tail_batch_trigger_size=self.tail_batch_trigger_size, - is_valid_sample_fn=self.is_valid_sample_fn, should_continue_fn=self.should_continue_fn, ) @@ -274,10 +270,8 @@ def build( class DisaggProduceStrategy(ABC): def __init__( self, - is_valid_sample_fn: IsValidSampleFn, should_continue_fn: ShouldContinueFn, ): - self.is_valid_sample_fn = is_valid_sample_fn self.should_continue_fn = should_continue_fn @abstractmethod @@ -305,10 +299,9 @@ def __init__( tail_batch_trigger_size: int, max_staleness: int, sync_weights_interval: int, - is_valid_sample_fn: IsValidSampleFn, should_continue_fn: ShouldContinueFn, ): - super().__init__(is_valid_sample_fn, should_continue_fn) + super().__init__(should_continue_fn) if not enable_partial_rollout and max_staleness > 0: logger.warning( @@ -382,7 +375,7 @@ async def produce_batch(self, ctx: DisaggProduceContext) -> ProduceBatchStatus: async def spawn_one() -> asyncio.Task: rollout_state = await ctx.sample_group(from_expired_pool=sample_from_expired) return create_task( - ctx.generate_group( + ctx.collect_rollout_group( rollout_state, enable_partial_rollout=self.enable_partial_rollout, ) diff --git a/xtuner/v1/rl/agent_loop_manager/produce_utils.py b/xtuner/v1/rl/agent_loop_manager/produce_utils.py index b80a2f9e99..a669fe2005 100644 --- a/xtuner/v1/rl/agent_loop_manager/produce_utils.py +++ b/xtuner/v1/rl/agent_loop_manager/produce_utils.py @@ -83,10 +83,6 @@ class ProduceBatchStatus(Enum): EXPIRED_BATCH = auto() -def default_is_valid_sample_fn(samples: list[RolloutState]) -> bool: - return True - - def default_should_continue_fn(completed_count: int, batch_size: int, **kwargs) -> bool: return completed_count < batch_size @@ -123,7 +119,7 @@ class BaseProduceContext: train_step: int model_step: int progress: "ProduceProgress | DisaggProduceProgress" - is_valid_sample_fn: IsValidSampleFn = default_is_valid_sample_fn + is_valid_sample_fn: IsValidSampleFn | None = None stale_threshold: int | None = None @property @@ -137,7 +133,7 @@ async def sample_group(self, *, from_expired_pool: bool) -> list[RolloutState]: group_status = [Status.EXPIRED, Status.ABORTED] if from_expired_pool else [Status.ABORTED] return await self.sampler.sample(task_name=self.task_name, group_status=group_status) - async def generate_group( + async def collect_rollout_group( self, rollout_state: list[RolloutState], *, @@ -151,13 +147,15 @@ async def generate_group( start = time.perf_counter() if isinstance(self.agent_loop, ray.actor.ActorHandle): - result = await self.agent_loop.generate_group.remote( + result = await self.agent_loop.collect_rollout_group.remote( rollout_state, + is_valid_sample_func=self.is_valid_sample_fn, enable_partial_rollout=enable_partial_rollout, ) else: - result = await self.agent_loop.generate_group( + result = await self.agent_loop.collect_rollout_group( rollout_state, + is_valid_sample_func=self.is_valid_sample_fn, enable_partial_rollout=enable_partial_rollout, ) elapsed = time.perf_counter() - start @@ -172,9 +170,8 @@ async def generate_group( async def put_generated_group(self, group: list[RolloutState]) -> bool: produced_tokens = sum(len(item.response_ids) for item in group if item.response_ids is not None) initial_status = get_group_status(group) - discard_status: Status | None = None - if initial_status == Status.COMPLETED: + if initial_status in (Status.COMPLETED, Status.FILTERED): rewards_sum = 0.0 rewards_count = 0 for item in group: @@ -188,15 +185,10 @@ async def put_generated_group(self, group: list[RolloutState]) -> bool: rewards_count += 1 self.progress.add_raw_rewards(self.task_name, rewards_sum, rewards_count) - if not self.is_valid_sample_fn(group): - discard_status = Status.FILTERED - elif initial_status == Status.FAILED: - discard_status = Status.FAILED - - if discard_status is not None: + if initial_status in (Status.FAILED, Status.FILTERED): # 失败样本和业务过滤样本都不进入 replay buffer。 self.progress.add_produced(self.task_name, samples=len(group), tokens=produced_tokens) - self.progress.add_discarded(self.task_name, discard_status, samples=len(group)) + self.progress.add_discarded(self.task_name, initial_status, samples=len(group)) for item in group: discard_rollout_state(item) return False @@ -272,13 +264,10 @@ class _TaskRunner: agent_loop: AgentLoopSpec produce_strategy: Any sampler: Sampler + is_valid_sample_fn: IsValidSampleFn | None = None weight: float = 1.0 order: int = 0 - @property - def is_valid_sample_fn(self) -> IsValidSampleFn: - return getattr(self.produce_strategy, "is_valid_sample_fn", default_is_valid_sample_fn) - @property def stale_threshold(self) -> int | None: return getattr(self.produce_strategy, "stale_threshold", None) diff --git a/xtuner/v1/rl/agent_loop_manager/producer.py b/xtuner/v1/rl/agent_loop_manager/producer.py index 5620be7772..f309e89ba5 100644 --- a/xtuner/v1/rl/agent_loop_manager/producer.py +++ b/xtuner/v1/rl/agent_loop_manager/producer.py @@ -13,12 +13,10 @@ from .produce_utils import ( PERIODIC_ABORT_INTERVAL_S, BaseProduceContext, - IsValidSampleFn, ShouldContinueFn, _ProgressDisplayer, _put_claimed_tasks, calculate_stale_threshold, - default_is_valid_sample_fn, default_should_continue_fn, pause_pending_tasks, ) @@ -127,16 +125,12 @@ class ProduceStrategyConfig(ABC, BaseModel): when it should stop producing samples for the current training step. Args: - is_valid_sample_fn (IsValidSampleFn): Function used to decide whether a - generated rollout group is trainable. Defaults to - ``default_is_valid_sample_fn``. should_continue_fn (ShouldContinueFn): Function used to decide whether production should continue after a group is processed. Defaults to ``default_should_continue_fn``. """ model_config = ConfigDict(extra="forbid", arbitrary_types_allowed=True) - is_valid_sample_fn: IsValidSampleFn = default_is_valid_sample_fn should_continue_fn: ShouldContinueFn = default_should_continue_fn @abstractmethod @@ -156,9 +150,6 @@ class SyncProduceStrategyConfig(ProduceStrategyConfig): in a colocated or tightly synchronized workflow. Args: - is_valid_sample_fn (IsValidSampleFn): Function used to decide whether a - generated rollout group is trainable. Defaults to - ``default_is_valid_sample_fn``. should_continue_fn (ShouldContinueFn): Function used to decide whether production should continue after a group is processed. Defaults to ``default_should_continue_fn``. @@ -176,9 +167,7 @@ def build( sync_weights_interval: int = 1, rollout_controller: "Optional[RolloutControllerProxy]" = None, ) -> "SyncProduceStrategy": - return SyncProduceStrategy( - is_valid_sample_fn=self.is_valid_sample_fn, should_continue_fn=self.should_continue_fn - ) + return SyncProduceStrategy(should_continue_fn=self.should_continue_fn) class AsyncProduceStrategyConfig(ProduceStrategyConfig): @@ -190,9 +179,6 @@ class AsyncProduceStrategyConfig(ProduceStrategyConfig): discard samples that are too stale relative to the current training step. Args: - is_valid_sample_fn (IsValidSampleFn): Function used to decide whether a - generated rollout group is trainable. Defaults to - ``default_is_valid_sample_fn``. should_continue_fn (ShouldContinueFn): Function used to decide whether production should continue after a group is processed. Defaults to ``default_should_continue_fn``. @@ -237,7 +223,6 @@ def build( max_staleness=self.max_staleness, sync_weights_interval=sync_weights_interval, tail_batch_trigger_size=self.tail_batch_trigger_size, - is_valid_sample_fn=self.is_valid_sample_fn, should_continue_fn=self.should_continue_fn, ) @@ -245,10 +230,8 @@ def build( class ProduceStrategy(ABC): def __init__( self, - is_valid_sample_fn: IsValidSampleFn, should_continue_fn: ShouldContinueFn, ): - self.is_valid_sample_fn = is_valid_sample_fn self.should_continue_fn = should_continue_fn @abstractmethod @@ -268,7 +251,7 @@ async def produce_batch(self, ctx: ProduceContext) -> None: for _ in range(ctx.task_batch_size): rollout_state = await ctx.sampler.sample(task_name=ctx.task_name) - task = create_task(ctx.generate_group(rollout_state)) + task = create_task(ctx.collect_rollout_group(rollout_state)) pending_tasks.add(task) logger.info(f"[SyncProduceStrategy] Started {len(pending_tasks)} initial tasks.") @@ -286,7 +269,7 @@ async def produce_batch(self, ctx: ProduceContext) -> None: done_tasks, pending_tasks = await asyncio.wait( pending_tasks, timeout=1, return_when=asyncio.FIRST_COMPLETED ) - # put_generated_group 负责过滤和入库。 + # AgentLoop 已完成过滤;put_generated_group 只处理状态、数据入库和释放。 for task in done_tasks: items = task.result() @@ -301,7 +284,7 @@ async def produce_batch(self, ctx: ProduceContext) -> None: completed_sample_count, ctx.task_batch_size ): rollout_state = await ctx.sampler.sample(task_name=ctx.task_name) - task = create_task(ctx.generate_group(rollout_state)) + task = create_task(ctx.collect_rollout_group(rollout_state)) pending_tasks.add(task) progress_displayer.close() @@ -316,10 +299,9 @@ def __init__( tail_batch_trigger_size: int, max_staleness: int, sync_weights_interval: int, - is_valid_sample_fn: IsValidSampleFn, should_continue_fn: ShouldContinueFn, ): - super().__init__(is_valid_sample_fn, should_continue_fn) + super().__init__(should_continue_fn) # TODO: 需要添加 tail_batch_max_tries # 作用是:如果一个样本多次重试,则将它置为特殊状态 MAX_TRIES,这类样本和过期样本一起触发tail batch逻辑 @@ -381,7 +363,7 @@ async def produce_batch(self, ctx: ProduceContext) -> None: async def spawn_one() -> asyncio.Task: rollout_state = await ctx.sample_group(from_expired_pool=sample_from_expired) return create_task( - ctx.generate_group( + ctx.collect_rollout_group( rollout_state, enable_partial_rollout=self.enable_partial_rollout, ) diff --git a/xtuner/v1/rl/loss/base_loss.py b/xtuner/v1/rl/loss/base_loss.py index 108af41ed8..18fc2d7175 100644 --- a/xtuner/v1/rl/loss/base_loss.py +++ b/xtuner/v1/rl/loss/base_loss.py @@ -106,6 +106,7 @@ def build( - shifted_labels (torch.Tensor): The shifted labels - advantages (torch.Tensor): Advantage estimates - rollout_logprobs (torch.Tensor | None): Rollout log probabilities + - teacher_logprobs (torch.Tensor | None): Teacher log probabilities for OPD - old_logprobs (torch.Tensor | None): Old policy log probabilities (optional, can be set later) - rollout_is_weights (torch.Tensor | None): Importance sampling weights - ref_logprobs (torch.Tensor | None): Reference model log probabilities @@ -122,6 +123,7 @@ def build( shifted_labels = data["shifted_labels"] advantages = data["advantages"] rollout_logprobs = data.get("rollout_logprobs", None) + teacher_logprobs = data.get("teacher_logprobs", None) old_logprobs = data.get("old_logprobs", None) rollout_is_weights = data.get("rollout_is_weights", None) ref_logprobs = data.get("ref_logprobs", None) @@ -132,6 +134,7 @@ def build( old_logprobs=old_logprobs, advantages=advantages, rollout_logprobs=rollout_logprobs, + teacher_logprobs=teacher_logprobs, is_weights=rollout_is_weights, ref_logprobs=ref_logprobs, ).to(DEVICE) @@ -153,10 +156,12 @@ class BaseRLLossKwargs(CELossKwargs): ref_logprobs (torch.Tensor | None): Reference log probabilities for KL penalty, if used. kl_loss_weight (torch.Tensor | None): Weights for each token in the KL loss computation, if used. rollout_logprobs (torch.Tensor | None): Rollout log probabilities from inference engine, used for importance sampling. + teacher_logprobs (torch.Tensor | None): Teacher log probabilities used to apply the OPD KL penalty. is_weights (torch.Tensor | None): Importance sampling weights. If None, importance sampling is not used. """ rollout_logprobs: torch.Tensor | None = None + teacher_logprobs: torch.Tensor | None = None advantages: torch.Tensor old_logprobs: torch.Tensor | None = None policy_loss_weight: torch.Tensor | None = None @@ -172,6 +177,8 @@ def sp_split(self, sp_mesh: DeviceMesh) -> Self: self.advantages = sp_split(self.advantages, sp_mesh=sp_mesh, split_dim=1, padding_value=0.0) if self.rollout_logprobs is not None: self.rollout_logprobs = sp_split(self.rollout_logprobs, sp_mesh=sp_mesh, split_dim=1, padding_value=0.0) + if self.teacher_logprobs is not None: + self.teacher_logprobs = sp_split(self.teacher_logprobs, sp_mesh=sp_mesh, split_dim=1, padding_value=0.0) if self.is_weights is not None: self.is_weights = sp_split(self.is_weights, sp_mesh=sp_mesh, split_dim=1, padding_value=1.0) # 1. 这里不用对old_logprobs和ref_logprobs进行sp_split,因为他是模型 fwd 生成的, @@ -192,6 +199,8 @@ def to(self, device: torch.device | str) -> Self: self.ref_logprobs = self.ref_logprobs.to(device) if self.rollout_logprobs is not None: self.rollout_logprobs = self.rollout_logprobs.to(device) + if self.teacher_logprobs is not None: + self.teacher_logprobs = self.teacher_logprobs.to(device) if self.is_weights is not None: self.is_weights = self.is_weights.to(device) if self.global_grad_tokens is not None: diff --git a/xtuner/v1/rl/on_policy_distillation.py b/xtuner/v1/rl/on_policy_distillation.py new file mode 100644 index 0000000000..66a9dece16 --- /dev/null +++ b/xtuner/v1/rl/on_policy_distillation.py @@ -0,0 +1,283 @@ +from __future__ import annotations + +import asyncio +import math +import os +import time +from pathlib import Path +from typing import Any, Literal, cast + +import httpx +import torch +from pydantic import BaseModel, ConfigDict, Field, model_validator + +from xtuner.v1.data_proto.rl_data import RolloutState, SampleParams, Status +from xtuner.v1.rl.loss.base_loss import BaseRLLossContext + + +class OPDTeacherLaunchConfig(BaseModel): + model_config = ConfigDict(extra="forbid") + + model_path: str | Path + cuda_visible_devices: str + dtype: Literal["auto", "float16", "bfloat16"] = "bfloat16" + tensor_parallel_size: int = Field(default=1, gt=0) + expert_parallel_size: int = Field(default=1, gt=0) + context_length: int | None = Field(default=None, gt=0) + max_batch_size: int | None = Field(default=None, gt=0) + chunked_prefill_size: int | None = Field(default=4096, gt=0) + max_prefill_token_num: int | None = Field(default=4096, gt=0) + gpu_memory_utilization: float = Field(default=0.6, gt=0.0, le=1.0) + + @model_validator(mode="after") + def validate_parallel_sizes(self) -> OPDTeacherLaunchConfig: + if self.tensor_parallel_size > 1 and self.expert_parallel_size > 1: + raise ValueError("tensor_parallel_size and expert_parallel_size cannot both be greater than 1") + return self + + +class OPDTeacherConfig(BaseModel): + model_config = ConfigDict(extra="forbid") + + name: str + endpoint: str + api_key: str | None = None + request_timeout_s: float = Field(default=1200.0, gt=0.0) + max_retry_per_sample: int = Field(default=2, ge=0) + max_concurrency: int = Field(default=128, gt=0) + launch_config: OPDTeacherLaunchConfig | None = None + + +class OPDConfig(BaseModel): + model_config = ConfigDict(extra="forbid") + + mode: Literal["pg-opd"] = "pg-opd" + task_adv_weight: float = Field(default=0.0, ge=0.0) + opd_adv_weight: float = Field(default=1.0, ge=0.0) + teachers: list[OPDTeacherConfig] = Field(min_length=1) + data_source_teacher_map: dict[str, str] = Field(min_length=1) + + @model_validator(mode="after") + def validate_teacher_names(self) -> OPDConfig: + teacher_names = [teacher.name for teacher in self.teachers] + if len(teacher_names) != len(set(teacher_names)): + raise ValueError("OPD teacher names must be unique") + unknown_teachers = set(self.data_source_teacher_map.values()) - set(teacher_names) + if unknown_teachers: + raise ValueError(f"data_source_teacher_map references unknown teachers: {sorted(unknown_teachers)}") + return self + + +def validate_opd_sample_params(sample_params: SampleParams) -> None: + identity_sampling_params: dict[str, Any] = { + "temperature": 1.0, + "top_p": 1.0, + "top_k": 0, + "repetition_penalty": 1.0, + "presence_penalty": 0.0, + "frequency_penalty": 0.0, + "min_tokens": 0, + } + non_identity_params = { + name: getattr(sample_params, name) + for name, expected in identity_sampling_params.items() + if getattr(sample_params, name) != expected + } + if non_identity_params: + raise ValueError(f"PG-OPD requires identity student sampling, got {non_identity_params}") + if not sample_params.return_logprob or not sample_params.return_token_ids: + raise ValueError("PG-OPD requires return_logprob=True and return_token_ids=True") + + +class TeacherLogprobClient: + """Asynchronous teacher client scoped to one AgentLoop.""" + + def __init__(self, config: OPDTeacherConfig) -> None: + self.config = config + self.name = config.name + self.backend = self._resolve_backend_from_env() + self.url = f"{config.endpoint.rstrip('/')}/generate" + self._semaphore = asyncio.Semaphore(config.max_concurrency) + + headers = {"Content-Type": "application/json"} + if config.api_key is not None: + headers["Authorization"] = f"Bearer {config.api_key}" + self._client = httpx.AsyncClient(headers=headers, timeout=config.request_timeout_s) + + async def compute_logprobs(self, state: RolloutState) -> RolloutState: + start = time.perf_counter() + try: + prompt_ids = cast(list[int], state.prompt_ids) + response_ids = cast(list[int], state.response_ids) + if not prompt_ids or not response_ids: + state.status = Status.FAILED + state.error_msg = f"Teacher {self.name!r} scoring requires non-empty prompt_ids and response_ids" + return state + payload = self._construct_payload(prompt_ids, response_ids) + + retries = 0 + while True: + try: + async with self._semaphore: + response = await self._client.post(self.url, json=payload) + response.raise_for_status() + teacher_tokens, teacher_logprobs = self._parse_response(response, prompt_ids, response_ids) + state.teacher_tokens = teacher_tokens + state.teacher_logprobs = teacher_logprobs + return state + except (httpx.HTTPStatusError, httpx.RequestError, ValueError) as exc: + if retries >= self.config.max_retry_per_sample: + state.status = Status.FAILED + state.error_msg = f"Teacher {self.name!r} scoring failed after {retries + 1} attempts: {exc}" + return state + retries += 1 + await asyncio.sleep(0.1) + finally: + state.extra_fields["teacher_score_time_s"] = time.perf_counter() - start + + @staticmethod + def _resolve_backend_from_env() -> Literal["sglang", "lmdeploy"]: + use_sglang = os.environ.get("XTUNER_USE_SGLANG", "0") == "1" + use_lmdeploy = os.environ.get("XTUNER_USE_LMDEPLOY", "0") == "1" + use_vllm = os.environ.get("XTUNER_USE_VLLM", "0") == "1" + + if use_vllm: + raise RuntimeError("TeacherLogprobClient supports only SGLang or LMDeploy, not vLLM") + if use_sglang == use_lmdeploy: + raise RuntimeError("Exactly one of XTUNER_USE_SGLANG and XTUNER_USE_LMDEPLOY must be set to 1") + return "sglang" if use_sglang else "lmdeploy" + + def _construct_payload(self, prompt_ids: list[int], response_ids: list[int]) -> dict[str, Any]: + if self.backend == "sglang": + return self._construct_sglang_payload(prompt_ids, response_ids) + if self.backend == "lmdeploy": + return self._construct_lmdeploy_payload(prompt_ids, response_ids) + raise RuntimeError(f"Unsupported teacher backend: {self.backend}") + + @staticmethod + def _construct_sglang_payload(prompt_ids: list[int], response_ids: list[int]) -> dict[str, Any]: + return { + "input_ids": prompt_ids + response_ids, + "sampling_params": { + "max_new_tokens": 0, + "temperature": 0, + "skip_special_tokens": False, + }, + "return_logprob": True, + "logprob_start_len": 0, + "top_logprobs_num": 0, + "stream": False, + } + + @staticmethod + def _construct_lmdeploy_payload(prompt_ids: list[int], response_ids: list[int]) -> dict[str, Any]: + return { + "input_ids": prompt_ids + response_ids, + "return_logprob": True, + "logprob_start_len": 0, + "max_tokens": 0, + "stream": False, + } + + def _parse_response( + self, + response: httpx.Response, + prompt_ids: list[int], + response_ids: list[int], + ) -> tuple[list[int], list[float]]: + if self.backend == "sglang": + response_logprobs = self._parse_sglang_response(response, prompt_ids, response_ids) + else: + response_logprobs = self._parse_lmdeploy_response(response, prompt_ids, response_ids) + return self._validate_response_logprobs(response_logprobs, response_ids) + + @staticmethod + def _parse_sglang_response( + response: httpx.Response, + prompt_ids: list[int], + response_ids: list[int], + ) -> list[Any]: + raw_logprobs = TeacherLogprobClient._get_input_token_logprobs(response) + expected_length = len(prompt_ids) + len(response_ids) + if len(raw_logprobs) != expected_length: + raise ValueError( + "SGLang teacher logprob length mismatch: " + f"expected {expected_length} rows for the full input, got {len(raw_logprobs)}" + ) + return raw_logprobs[-len(response_ids) :] + + @staticmethod + def _parse_lmdeploy_response( + response: httpx.Response, + prompt_ids: list[int], + response_ids: list[int], + ) -> list[Any]: + raw_logprobs = TeacherLogprobClient._get_input_token_logprobs(response) + expected_length = len(prompt_ids) + len(response_ids) - 1 + if len(raw_logprobs) != expected_length: + raise ValueError( + "LMDeploy teacher logprob length mismatch: " + f"expected {expected_length} rows after the boundary token, got {len(raw_logprobs)}" + ) + return raw_logprobs[-len(response_ids) :] + + @staticmethod + def _get_input_token_logprobs(response: httpx.Response) -> list[Any]: + try: + raw_logprobs = response.json()["meta_info"]["input_token_logprobs"] + except (KeyError, TypeError, ValueError) as exc: + raise ValueError("Invalid teacher response") from exc + if not isinstance(raw_logprobs, list): + raise ValueError("Invalid teacher response") + return raw_logprobs + + @staticmethod + def _validate_response_logprobs( + response_logprobs: list[Any], + response_ids: list[int], + ) -> tuple[list[int], list[float]]: + try: + teacher_tokens = [item[1] for item in response_logprobs] + teacher_logprobs = [float(item[0]) for item in response_logprobs] + except (TypeError, IndexError, ValueError) as exc: + raise ValueError("Invalid teacher response") from exc + + if len(teacher_logprobs) != len(response_ids): + raise ValueError("Teacher logprob length mismatch") + if teacher_tokens != response_ids: + raise ValueError("Teacher token ids mismatch") + if not all(math.isfinite(logprob) for logprob in teacher_logprobs): + raise ValueError("Teacher logprobs contain NaN or Inf") + return teacher_tokens, teacher_logprobs + + +def route_teacher_client( + state: RolloutState, + *, + data_source_teacher_map: dict[str, str], + teacher_clients: dict[str, TeacherLogprobClient], +) -> TeacherLogprobClient: + data_source = state.extra_fields["origin_data_source"] + teacher_name = data_source_teacher_map[data_source] + return teacher_clients[teacher_name] + + +def apply_opd_kl_to_advantages( + loss_ctx: BaseRLLossContext, + *, + config: OPDConfig, +) -> torch.Tensor: + """Apply the OPD reverse-KL penalty and return its valid-token sum.""" + loss_kwargs = loss_ctx.loss_kwargs + old_logprobs = cast(torch.Tensor, loss_kwargs.old_logprobs) + teacher_logprobs = cast(torch.Tensor, loss_kwargs.teacher_logprobs) + response_mask = loss_kwargs.shifted_labels != loss_ctx.loss_cfg.ignore_idx + reverse_kl = old_logprobs - teacher_logprobs + reverse_kl_sum = (reverse_kl * response_mask).sum().detach() + loss_kwargs.advantages = torch.where( + response_mask, + loss_kwargs.advantages - config.opd_adv_weight * reverse_kl, + loss_kwargs.advantages, + ) + loss_kwargs.teacher_logprobs = None + return reverse_kl_sum diff --git a/xtuner/v1/rl/trainer/controller.py b/xtuner/v1/rl/trainer/controller.py index b87da88ced..e7e8b06e24 100644 --- a/xtuner/v1/rl/trainer/controller.py +++ b/xtuner/v1/rl/trainer/controller.py @@ -22,6 +22,7 @@ class ColateItem(TypedDict): shifted_labels: torch.Tensor advantage: float rollout_logprobs: torch.Tensor | None + teacher_logprobs: torch.Tensor | None class TrainingController: @@ -98,6 +99,10 @@ def _packing(self, data_batches, pack_max_length, language_cfg): if "rollout_logprobs" in data_batches[0] and data_batches[0]["rollout_logprobs"] is not None: rollout_logprobs_list = [data_batches[i]["rollout_logprobs"] for i in indices] + teacher_logprobs_list = None + if "teacher_logprobs" in data_batches[0] and data_batches[0]["teacher_logprobs"] is not None: + teacher_logprobs_list = [data_batches[i]["teacher_logprobs"] for i in indices] + if pad_len > 0: # Reduce the attn calculation time by using multiple short sequence packs pad_tokens = tuple( @@ -139,6 +144,14 @@ def _packing(self, data_batches, pack_max_length, language_cfg): device=data_batches[0]["shifted_labels"].device, ) rollout_logprobs_list.append(pad_rollout_logprobs) + if teacher_logprobs_list is not None: + pad_teacher_logprobs = torch.zeros( + 1, + pad_len, + dtype=data_batches[0]["teacher_logprobs"].dtype, + device=data_batches[0]["shifted_labels"].device, + ) + teacher_logprobs_list.append(pad_teacher_logprobs) seq_ctx = SequenceContext.cat(seq_ctx_list) shifted_labels = torch.cat(label_list, dim=1) # (1, max_len) @@ -149,12 +162,17 @@ def _packing(self, data_batches, pack_max_length, language_cfg): if rollout_logprobs_list is not None: rollout_logprobs = torch.cat(rollout_logprobs_list, dim=1) # (1, max_len) + teacher_logprobs = None + if teacher_logprobs_list is not None: + teacher_logprobs = torch.cat(teacher_logprobs_list, dim=1) # (1, max_len) + packed_data_batches.append( { "seq_ctx": seq_ctx, "shifted_labels": shifted_labels, "advantages": advantages, "rollout_logprobs": rollout_logprobs, + "teacher_logprobs": teacher_logprobs, } ) return packed_data_batches @@ -237,11 +255,17 @@ def fit(self, data_batches: list[ColateItem], pack_max_length: int, rollout_idx: pad_rollout_logprobs = torch.zeros( 1, pack_max_length, dtype=packed_data_batches[0]["rollout_logprobs"].dtype, device="cpu" ) + pad_teacher_logprobs = None + if "teacher_logprobs" in packed_data_batches[0] and packed_data_batches[0]["teacher_logprobs"] is not None: + pad_teacher_logprobs = torch.zeros( + 1, pack_max_length, dtype=packed_data_batches[0]["teacher_logprobs"].dtype, device="cpu" + ) pad_data = { "seq_ctx": pad_seq_ctx, "shifted_labels": pad_shifted_labels, "advantages": pad_advantages, "rollout_logprobs": pad_rollout_logprobs, + "teacher_logprobs": pad_teacher_logprobs, } pad_data_samples = [pad_data for _ in range(pad_num)] packed_data_batches = packed_data_batches + pad_data_samples diff --git a/xtuner/v1/rl/trainer/worker.py b/xtuner/v1/rl/trainer/worker.py index f7e0bc4871..1848866d1c 100644 --- a/xtuner/v1/rl/trainer/worker.py +++ b/xtuner/v1/rl/trainer/worker.py @@ -45,6 +45,7 @@ from xtuner.v1.model.utils.misc import ModelForwardExtraLogInfo from xtuner.v1.profiler import profiling_memory, profiling_time from xtuner.v1.rl.loss import BaseRLLossConfig, BaseRLLossContext, finalize_train_policy_metrics, kl_penalty +from xtuner.v1.rl.on_policy_distillation import OPDConfig, apply_opd_kl_to_advantages from xtuner.v1.rl.utils import SingleAcceleratorWorker from xtuner.v1.rl.weight_update import UpdateWeighter from xtuner.v1.train.trainer import LoadCheckpointConfig @@ -151,6 +152,7 @@ class WorkerConfig(BaseModel): profile_memory: bool = False free_rollout_routed_experts_in_worker: bool = True # 默认不需要用户配置 offload_rollout_routed_experts: bool = False + opd_config: OPDConfig | None = None # sft config sft_dataloader_cfg: DataloaderConfig | None = None @@ -184,6 +186,7 @@ class WorkerInputItem(TypedDict): shifted_labels: torch.LongTensor advantages: torch.Tensor rollout_logprobs: torch.Tensor | None + teacher_logprobs: torch.Tensor | None class WorkerTrainLogItem(TypedDict, total=False): @@ -196,6 +199,7 @@ class WorkerTrainLogItem(TypedDict, total=False): class WorkerLogItem(TypedDict): train_entropy: float + opd_reverse_kl: NotRequired[float] rollout_entropy: NotRequired[float] mismatch_metrics: NotRequired[dict[str, float]] rollout_is_metrics: NotRequired[dict[str, float]] @@ -586,6 +590,7 @@ def fit(self, data_batches: list[WorkerInputItem], rollout_idx: int) -> WorkerLo "shifted_labels": shifted_labels, "advantages": advantages, "rollout_logprobs": rollout_logprobs, + "teacher_logprobs": data.get("teacher_logprobs", None), }, sp_mesh=self.sp_mesh, ) @@ -625,8 +630,14 @@ def fit(self, data_batches: list[WorkerInputItem], rollout_idx: int) -> WorkerLo # compute old logprobs old_logprobs_list = self.compute_actor_logprobs(seq_ctx_list, shifted_labels_list) + rank_opd_reverse_kl_sum: torch.Tensor | None = None for old_logprobs, loss_ctx in zip(old_logprobs_list, loss_ctx_list): loss_ctx.loss_kwargs.old_logprobs = old_logprobs + if self.config.opd_config is not None: + reverse_kl_sum = apply_opd_kl_to_advantages(loss_ctx, config=self.config.opd_config) + rank_opd_reverse_kl_sum = ( + reverse_kl_sum if rank_opd_reverse_kl_sum is None else rank_opd_reverse_kl_sum + reverse_kl_sum + ) worker_log_item: WorkerLogItem = {"train_entropy": 0.0, "train_metrics": [], "sft_train_metrics": {}} logger_msg = f"Rollout {rollout_idx}: " @@ -651,6 +662,17 @@ def fit(self, data_batches: list[WorkerInputItem], rollout_idx: int) -> WorkerLo worker_log_item["rollout_entropy"] = avg_rollout_entropy.item() logger_msg += f", avg rollout entropy: {avg_rollout_entropy:.4f}" + if rank_opd_reverse_kl_sum is not None: + global_opd_reverse_kl_sum = rank_opd_reverse_kl_sum + dist.all_reduce(global_opd_reverse_kl_sum, op=dist.ReduceOp.SUM) + avg_opd_reverse_kl = ( + global_opd_reverse_kl_sum / global_grad_tokens + if global_grad_tokens > 0 + else global_opd_reverse_kl_sum.new_zeros(()) + ) + worker_log_item["opd_reverse_kl"] = avg_opd_reverse_kl.item() + logger_msg += f", OPD reverse KL: {avg_opd_reverse_kl:.4f}" + # compute rollout importance sampling metrics all_rollout_is_metrics = [] all_mismatch_metrics = [] diff --git a/xtuner/v1/train/rl_trainer.py b/xtuner/v1/train/rl_trainer.py index 9e810ad162..89e4f89404 100644 --- a/xtuner/v1/train/rl_trainer.py +++ b/xtuner/v1/train/rl_trainer.py @@ -33,6 +33,7 @@ ) from xtuner.v1.rl.agent_loop_manager.produce_utils import default_should_continue_fn from xtuner.v1.rl.evaluator import EvaluatorConfig +from xtuner.v1.rl.on_policy_distillation import OPDConfig from xtuner.v1.rl.replay_buffer import ( AsyncReplayBufferConfig, SyncReplayBufferConfig, @@ -342,6 +343,7 @@ class BaseRLTrainerConfig(BaseModel): total_epochs: int | None = None train_batch_size: int advantage_estimator_config: BaseAdvantageConfig = Field(default_factory=GRPOAdvantageConfig) + opd_config: OPDConfig | None = None sync_weights_interval: int = 1 enable_evaluate: bool = True @@ -594,6 +596,7 @@ def _init_common(self, cfg: BaseRLTrainerConfig, *, meta_path: str, logger_tag: self._init_rollout_config(cfg, log_dir) self._ensure_rollout_proxy_config(cfg) self._init_runtime_flags(cfg) + self._opd_config = cfg.opd_config self._advantage_estimator = cfg.advantage_estimator_config.build() self._cpu_resource_manager: CPUResourceManager | None = None self._num_workers = 1.0 @@ -675,6 +678,7 @@ def _init_train_worker_config(self, cfg: BaseRLTrainerConfig, log_dir: Path) -> cfg.train_worker_cfg.free_rollout_routed_experts_in_worker = False cfg.train_worker_cfg.load_from = cfg.load_from cfg.train_worker_cfg.log_dir = log_dir + cfg.train_worker_cfg.opd_config = cfg.opd_config self._train_worker_cfg = cfg.train_worker_cfg def _init_rollout_config(self, cfg: BaseRLTrainerConfig, log_dir: Path) -> None: @@ -708,6 +712,7 @@ def _build_agent_loop_components(self, cfg: BaseRLTrainerConfig, replay_buffer) replay_buffer=replay_buffer, logger=self.logger, sync_weights_interval=cfg.sync_weights_interval, + opd_config=self._opd_config, ) self.agent_loop_manager = cast(AgentLoopManager | DisaggAgentLoopManager, agent_loop_manager) @@ -722,6 +727,7 @@ def _build_agent_loop_components(self, cfg: BaseRLTrainerConfig, replay_buffer) replay_buffer=replay_buffer, logger=self.logger, sync_weights_interval=cfg.sync_weights_interval, + opd_config=None, ), ) @@ -1014,13 +1020,14 @@ def _prepare_train_data( rewards_list = [] # Per-session rewards for distribution metrics. Agentic sessions may split into several # trainable segments that share one reward; counting that reward once per session keeps - # rewards/* from being weighted by segment count. rewards_list stays per-segment for counts. + # rewards/* from being weighted by segment count. cluster_rewards_list: list[float] = [] advantages_list = [] prompt_len_list = [] response_len_list = [] tool_turns_list: list[int] = [] training_tokens = 0 + training_samples = 0 data_batches = [] @@ -1028,7 +1035,11 @@ def _prepare_train_data( if not is_valid_for_training(group, self.logger): self.logger.error(f"Skip one data group {group} due to rollout failed or empty response.") continue + training_samples += len(group) + # When opd_config is not None, PG-OPD currently supports only single-turn rollouts. + opd_config = self._opd_config + task_adv_weight = opd_config.task_adv_weight if opd_config is not None else 1.0 prompt_ids = None if any(data.input_ids is None for data in group): is_vlm_model = "train_prompt_ids" in group[0].extra_fields @@ -1040,39 +1051,47 @@ def _prepare_train_data( assert prompt_ids is not None and len(prompt_ids) > 0, ( f"Prompt ids cannot be None or empty in data: {group[0]}" ) - rewards = [] - # Agentic rollouts may split one model session into multiple trainable segments. - # Compute the group advantage once per session, then broadcast it back to each segment. - cluster_index_by_key: dict[Any, int] = {} - cluster_rewards: list[float] = [] - cluster_representatives: list[RolloutState] = [] - sample_cluster_indices: list[int] = [] for data in group: - assert data.reward is not None and "score" in data.reward, ( - f"Reward is missing or does not contain 'score' key in data: {data}" - ) - reward = float(data.reward["score"]) - rewards.append(reward) - # session_id is only set by agentic loops / XTUNER_DETERMINISTIC; plain RL falls back - # to rollout_id, which the sampler always assigns. Segments of one session share a key. - cluster_key = data.session_id if data.session_id is not None else data.rollout_id - cluster_index = cluster_index_by_key.get(cluster_key) - if cluster_index is None: - cluster_index = len(cluster_rewards) - cluster_index_by_key[cluster_key] = cluster_index - cluster_rewards.append(reward) - cluster_representatives.append(data) - sample_cluster_indices.append(cluster_index) # 有可能有重复,但是没有其他更好办法 turns = data.extra_fields.get("agent_tool_turns") if isinstance(turns, int): tool_turns_list.append(turns) - rewards_list.extend(rewards) - cluster_rewards_list.extend(cluster_rewards) - rewards_tensor = torch.tensor(cluster_rewards, dtype=torch.float32) - cluster_advantages = self._advantage_estimator.compute(rewards_tensor, cluster_representatives) - sample_advantages = [cluster_advantages[cluster_index].item() for cluster_index in sample_cluster_indices] + if task_adv_weight == 0: + sample_advantages = [0.0] * len(group) + else: + rewards = [] + # Agentic rollouts may split one model session into multiple trainable segments. + # Compute the group advantage once per session, then broadcast it back to each segment. + cluster_index_by_key: dict[Any, int] = {} + cluster_rewards: list[float] = [] + cluster_representatives: list[RolloutState] = [] + sample_cluster_indices: list[int] = [] + for data in group: + assert data.reward is not None and "score" in data.reward, ( + f"Reward is missing or does not contain 'score' key in data: {data}" + ) + reward = float(data.reward["score"]) + rewards.append(reward) + # session_id is only set by agentic loops / XTUNER_DETERMINISTIC; plain RL falls back + # to rollout_id, which the sampler always assigns. Segments of one session share a key. + cluster_key = data.session_id if data.session_id is not None else data.rollout_id + cluster_index = cluster_index_by_key.get(cluster_key) + if cluster_index is None: + cluster_index = len(cluster_rewards) + cluster_index_by_key[cluster_key] = cluster_index + cluster_rewards.append(reward) + cluster_representatives.append(data) + sample_cluster_indices.append(cluster_index) + + rewards_list.extend(rewards) + cluster_rewards_list.extend(cluster_rewards) + rewards_tensor = torch.tensor(cluster_rewards, dtype=torch.float32) + cluster_advantages = self._advantage_estimator.compute(rewards_tensor, cluster_representatives) + sample_advantages = [ + cluster_advantages[cluster_index].item() for cluster_index in sample_cluster_indices + ] + sample_advantages = [task_adv_weight * sample_advantage for sample_advantage in sample_advantages] prompt_repeat_k = len(group) for i in range(prompt_repeat_k): @@ -1177,12 +1196,9 @@ def _prepare_train_data( shifted_labels = [-100] * (len(prompt_ids) - 1) + response_labels shifted_labels_t = torch.tensor(shifted_labels, dtype=torch.int64).unsqueeze(0) - # 根据 response_mask 计算新的 advantages - advatnages_val = sample_advantages[i] - actual_advantages = [advatnages_val] * len(prompt_ids) + [ - 0.0 if mask == 0 else advatnages_val for mask in response_mask - ] - advantages_list.extend(actual_advantages[:-1]) + base_advantage = sample_advantages[i] + actual_advantages = [0.0 if label == -100 else base_advantage for label in shifted_labels] + advantages_list.extend(actual_advantages) assert len(input_ids) <= pack_max_length, f"{len(input_ids)} vs {pack_max_length}" training_tokens += len(input_ids) @@ -1207,6 +1223,12 @@ def _prepare_train_data( "advantage": actual_advantages, "rollout_logprobs": rollout_logprobs, } + if opd_config is not None: + response_teacher_logprobs = cast(list[float], group[i].teacher_logprobs) + data_dict["teacher_logprobs"] = torch.tensor( + [0.0] * (len(prompt_ids) - 1) + response_teacher_logprobs, + dtype=torch.float32, + ).unsqueeze(0) seq_ctx.rollout_routed_experts = group[i].routed_experts # n,layer*expert @@ -1214,8 +1236,8 @@ def _prepare_train_data( if not XTUNER_DETERMINISTIC: random.shuffle(data_batches) - # rewards/* report the per-session reward distribution; batch_size/training_samples below - # still use rewards_list (per-segment) so counts reflect the actual training samples. + # rewards/* report the per-session reward distribution; batch_size/training_samples + # count the valid rollout segments included in training. rewards_t = torch.tensor(cluster_rewards_list).float() if cluster_rewards_list else torch.tensor([0.0]).float() advantages_t = torch.tensor(advantages_list).float() if advantages_list else torch.tensor([0.0]).float() prompt_len_t = torch.tensor(prompt_len_list).float() if prompt_len_list else torch.tensor([0.0]).float() @@ -1223,8 +1245,8 @@ def _prepare_train_data( raw_rewards_mean = raw_rewards_sum / raw_rewards_count if raw_rewards_count > 0 else rewards_t.mean().item() info_dict = { - "batch_size": len(rewards_list), - "training_samples": len(rewards_list), + "batch_size": training_samples, + "training_samples": training_samples, "training_tokens": training_tokens, "rewards/mean": rewards_t.mean().item(), "rewards/min": rewards_t.min().item(), @@ -1363,6 +1385,8 @@ def _log_step( all_scalars.update({f"{k}": v for k, v in rank0_mismatch_metrics.items()}) all_scalars.update({"entropy/rollout": rank0_rollout_entropy}) all_scalars.update({"entropy/train": rank0_log_item["train_entropy"]}) + if "opd_reverse_kl" in rank0_log_item: + all_scalars["opd_reverse_kl"] = rank0_log_item["opd_reverse_kl"] for worker_idx, log_item in enumerate(train_info["workers_log_item"]): if not self._display_all_workers_log and worker_idx > 0: break @@ -1400,8 +1424,9 @@ def _save_trajectories(self, data_groups: list[list[RolloutState]], save_path: P if not is_valid_for_training(group, self.logger): continue for data in group: - assert data.reward is not None - rewards.append(data.reward["score"]) + reward = data.reward.get("score") if data.reward is not None else None + if reward is not None: + rewards.append(reward) response_ids = self._get_trajectory_response_ids(data) response = data.response if response is None and response_ids: @@ -1422,7 +1447,7 @@ def _save_trajectories(self, data_groups: list[list[RolloutState]], save_path: P "prompt": data.message, "label": ground_truth, "response": response, - "reward": data.reward["score"], + "reward": reward, "prompt_len": data.num_tokens, "response_len": response_len, "reward_payload": data.reward, @@ -1447,14 +1472,14 @@ def _save_trajectories(self, data_groups: list[list[RolloutState]], save_path: P with open(save_path, "w", encoding="utf-8") as f: summary = { "reward_mean": rewards_tensor.mean().item(), - "reward_std": rewards_tensor.std().item(), + "reward_std": rewards_tensor.std(unbiased=False).item(), "reward_max": rewards_tensor.max().item(), "reward_min": rewards_tensor.min().item(), "response_len_mean": response_lens.mean().item(), - "response_len_std": response_lens.std().item(), + "response_len_std": response_lens.std(unbiased=False).item(), "response_len_max": response_lens.max().item(), "response_len_min": response_lens.min().item(), - "total_len": len(rewards), + "total_len": len(trajectory_items), } json.dump(summary, f, ensure_ascii=False, separators=(",", ":")) f.write("\n")