diff --git a/examples/v1/config/reasoning_rl_qwen3p5vl_mtp_ep.py b/examples/v1/config/reasoning_rl_qwen3p5vl_mtp_ep.py index 0d1b42039d..c27f3e410f 100644 --- a/examples/v1/config/reasoning_rl_qwen3p5vl_mtp_ep.py +++ b/examples/v1/config/reasoning_rl_qwen3p5vl_mtp_ep.py @@ -291,10 +291,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( @@ -306,6 +305,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..ba71b8e686 --- /dev/null +++ b/recipe/on_policy_distillation/build_teacher_server_commands.py @@ -0,0 +1,418 @@ +"""Build executable Teacher server commands from an OPD config. + +The NUL-delimited output starts with the Teacher replica count, resolved +endpoint mapping, total Student worker count, and current-node Student worker +count. Each replica record contains its placement, endpoint, health URLs, and +executable command arguments. An externally managed Teacher replica has a +target node rank of ``-1`` and a command-argument count of zero. +""" + +import argparse +import json +import os +import sys +from dataclasses import dataclass +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.distillation import RolloutTeacherConfig # noqa: E402 +from xtuner.v1.utils.config import Config # noqa: E402 + + +@dataclass(frozen=True) +class TeacherReplicaRequest: + teacher: RolloutTeacherConfig + replica_index: int + + @property + def key(self) -> tuple[str, int]: + return self.teacher.name, self.replica_index + + @property + def display_name(self) -> str: + return f"{self.teacher.name}[{self.replica_index}]" + + +def build_teacher_launch_server_commands( + config_path: str, + backend: Literal["sglang", "lmdeploy"], +) -> tuple[dict[str, list[str]], int, int, list[list[str]]]: + """Build executable Teacher server command records. + + Args: + config_path (str): Path to the XTuner Python config containing + ``distillation_config``. + backend (Literal["sglang", "lmdeploy"]): Teacher serving backend. + + Returns: + A four-item tuple containing: + + - ``endpoint_map``: Logical Teacher name to advertised replica + endpoints. + - ``student_num_workers``: Total number of Student GPUs in the cluster. + - ``student_local_num_workers``: Number of Student GPUs on the current + node. + - ``records``: One flattened string record per Teacher replica. Each + record is ``[display_name, target_node_rank, local_devices, endpoint, + health_url, model_info_url, command_arg_count, *command]``. + """ + config = Config.fromfile(config_path) + node_count = int(os.environ.get("NODE_COUNT", "1")) + node_rank = int(os.environ.get("NODE_RANK", "0")) + gpus_per_node = int(os.environ.get("PROC_PER_NODE", "8")) + node_addresses = tuple( + address.strip() + for address in os.environ.get( + "WORKER_ALL_SOCKET_ADDRS", "127.0.0.1" + ).split(",") + ) + assert node_count > 0 + assert 0 <= node_rank < node_count + assert gpus_per_node > 0 + assert len(node_addresses) == node_count + assert all(node_addresses) + + teachers = config.distillation_config.rollout_teachers + replica_requests = _expand_teacher_replicas(teachers) + health_path, model_info_path = _get_teacher_server_paths(backend) + teacher_placements, student_local_num_workers = _allocate_teacher_devices( + replica_requests, + node_count=node_count, + node_rank=node_rank, + gpus_per_node=gpus_per_node, + ) + teacher_num_workers = sum( + len(local_devices) + for _, local_devices in teacher_placements.values() + ) + student_num_workers = node_count * gpus_per_node - teacher_num_workers + endpoint_map: dict[str, list[str]] = {teacher.name: [] for teacher in teachers} + records: list[list[str]] = [] + used_node_ports: set[tuple[int, int]] = set() + + for replica in replica_requests: + teacher = replica.teacher + launch_config = teacher.launch_config + if launch_config is None: + target_node_rank = -1 + local_cuda_visible_devices = "" + endpoint = teacher.endpoints[replica.replica_index].rstrip("/") + command: list[str] = [] + else: + target_node_rank, local_devices = teacher_placements[replica.key] + server_port = _allocate_teacher_server_port( + launch_config.server_port, + target_node_rank=target_node_rank, + used_node_ports=used_node_ports, + replica_name=replica.display_name, + ) + + local_cuda_visible_devices = ",".join(str(device) for device in local_devices) + endpoint = f"http://{node_addresses[target_node_rank]}:{server_port}" + if backend == "sglang": + command = _build_sglang_command(teacher, server_port=server_port) + elif backend == "lmdeploy": + command = _build_lmdeploy_command(teacher, server_port=server_port) + else: + raise ValueError(f"Unsupported Teacher backend: {backend}") + + endpoint_map[teacher.name].append(endpoint) + records.append( + [ + replica.display_name, + str(target_node_rank), + local_cuda_visible_devices, + endpoint, + f"{endpoint}/{health_path}", + f"{endpoint}/{model_info_path}", + str(len(command)), + *command, + ] + ) + return endpoint_map, student_num_workers, student_local_num_workers, records + + +def _expand_teacher_replicas( + teachers: list[RolloutTeacherConfig], +) -> list[TeacherReplicaRequest]: + return [ + TeacherReplicaRequest( + teacher=teacher, + replica_index=replica_index, + ) + for teacher in teachers + for replica_index in range(teacher.num_replicas) + ] + + +def _allocate_teacher_devices( + replica_requests: list[TeacherReplicaRequest], + *, + node_count: int, + node_rank: int, + gpus_per_node: int, +) -> tuple[dict[tuple[str, int], tuple[int, list[int]]], int]: + """Allocate high-rank GPUs to local Teachers and leave the rest to Student. + + Teacher replicas are processed by descending ``num_workers``. Requests + with the same size keep their expanded config order. Each replica is placed + wholly on the highest-rank node that has enough free GPUs, using that + node's highest free local device ordinals. Replicas without + ``launch_config`` are externally managed and do not consume cluster GPUs. + + Args: + replica_requests: Teacher replicas in expanded config order. + node_count: Number of homogeneous nodes in the cluster. + node_rank: Rank of the current node. + gpus_per_node: Number of local GPUs available on every node. + + Returns: + A pair of ``(placements, student_local_num_workers)``: + + - ``placements`` maps each local ``(Teacher name, replica index)`` to + ``(target_node_rank, local_device_ordinals)``. + - ``student_local_num_workers`` is the number of remaining GPUs assigned + to Student on the current node. + + For inputs equivalent to: + + .. code-block:: python + + node_count = 4 + node_rank = 3 + gpus_per_node = 8 + replica_requests = [ + ("teacher1", 0, 4), + ("teacher2", 0, 2), + ] + + the returned values are: + + .. code-block:: python + + placements = { + ("teacher1", 0): (3, [4, 5, 6, 7]), + ("teacher2", 0): (3, [2, 3]), + } + student_local_num_workers = 2 + + Raises: + ValueError: If a Teacher cannot fit wholly on one node, or if Teacher + allocation leaves no GPU for Student. + """ + free_devices_by_node = [ + list(range(gpus_per_node)) for _ in range(node_count) + ] + local_teacher_requests: list[tuple[TeacherReplicaRequest, int]] = [] + for replica in replica_requests: + launch_config = replica.teacher.launch_config + if launch_config is not None: + local_teacher_requests.append((replica, launch_config.num_workers)) + local_teacher_requests.sort(key=lambda request: request[1], reverse=True) + + placements: dict[tuple[str, int], tuple[int, list[int]]] = {} + for replica, num_workers in local_teacher_requests: + if num_workers > gpus_per_node: + raise ValueError( + f"Teacher replica {replica.display_name!r} requests {num_workers} workers, " + f"but each node has only {gpus_per_node} GPUs" + ) + + target_node_rank = next( + ( + node_rank + for node_rank in range(node_count - 1, -1, -1) + if len(free_devices_by_node[node_rank]) >= num_workers + ), + None, + ) + if target_node_rank is None: + remaining_devices = sum( + len(local_devices) for local_devices in free_devices_by_node + ) + raise ValueError( + f"Teacher replica {replica.display_name!r} requests {num_workers} workers, " + "but no single node has enough remaining GPUs; " + f"{remaining_devices} GPUs remain across the cluster" + ) + + free_devices = free_devices_by_node[target_node_rank] + local_devices = free_devices[-num_workers:] + del free_devices[-num_workers:] + placements[replica.key] = (target_node_rank, local_devices) + + if not any(free_devices_by_node): + raise ValueError("Teacher allocation leaves no GPUs for Student workers") + student_local_num_workers = len(free_devices_by_node[node_rank]) + return placements, student_local_num_workers + + +def _allocate_teacher_server_port( + base_port: int, + *, + target_node_rank: int, + used_node_ports: set[tuple[int, int]], + replica_name: str, +) -> int: + server_port = base_port + while (target_node_rank, server_port) in used_node_ports: + server_port += 1 + if server_port > 65535: + raise ValueError( + f"Teacher replica {replica_name!r} cannot allocate a free port " + f"on node {target_node_rank} starting from {base_port}" + ) + used_node_ports.add((target_node_rank, server_port)) + return server_port + + +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 _build_sglang_command( + teacher: RolloutTeacherConfig, + *, + server_port: int, +) -> list[str]: + config = teacher.launch_config + assert config is not None + tensor_parallel_size = config.tensor_parallel_size + if config.expert_parallel_size > 1: + tensor_parallel_size = config.expert_parallel_size + + command = [ + sys.executable, + "-m", + "sglang.launch_server", + "--model-path", + str(config.model_path), + "--host", + "0.0.0.0", + "--port", + str(server_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.log_level is not None: + command.extend(["--log-level", config.log_level]) + 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: RolloutTeacherConfig, + *, + server_port: int, +) -> list[str]: + config = teacher.launch_config + assert config is not None + data_parallel_size = ( + config.expert_parallel_size + if config.expert_parallel_size > 1 + else 1 + ) + command = [ + sys.executable, + "-m", + "lmdeploy", + "serve", + "api_server", + str(config.model_path), + "--backend", + "pytorch", + "--role", + "Hybrid", + "--logprobs-mode", + "raw_logprobs", + "--server-name", + "0.0.0.0", + "--server-port", + str(server_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.log_level is not None: + command.extend(["--log-level", config.log_level.upper()]) + 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)]) + if teacher.enable_prefix_caching: + command.append("--enable-prefix-caching") + return command + + +def _write_teacher_records( + endpoint_map: dict[str, list[str]], + student_num_workers: int, + student_local_num_workers: int, + records: list[list[str]], +) -> None: + endpoint_map_json = json.dumps( + endpoint_map, + ensure_ascii=False, + separators=(",", ":"), + sort_keys=True, + ) + fields = [ + str(len(records)), + endpoint_map_json, + str(student_num_workers), + str(student_local_num_workers), + ] + 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() + endpoint_map, student_num_workers, student_local_num_workers, records = ( + build_teacher_launch_server_commands(args.config_path, args.backend) + ) + _write_teacher_records( + endpoint_map, + student_num_workers, + student_local_num_workers, + records, + ) + + +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..7f930090b7 --- /dev/null +++ b/recipe/on_policy_distillation/config/rl_dapo_math_mopd.py @@ -0,0 +1,399 @@ +import json +import os + +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 Qwen3VLDense2BConfig +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.distillation import ( + DistillationConfig, + RolloutTeacherConfig, + RolloutTeacherLaunchConfig, +) +from xtuner.v1.rl.loss import DistillationLossConfig +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] + + +EVAL_DATA_SOURCE_TAGS = { + "openai/gsm8k": "gsm8k", + "hiyouga/geometry3k": "geo3k", +} + + +def mopd_compute_metric(samples): + scores_by_source = {data_source: [] for data_source in EVAL_DATA_SOURCE_TAGS} + for sample in samples: + data_source = sample.extra_fields.get("origin_data_source") + if data_source not in scores_by_source: + raise ValueError(f"Unexpected evaluation data source: {data_source!r}") + + reward = sample.reward or {} + if "score" not in reward: + raise ValueError(f"Missing reward score for evaluation data source: {data_source!r}") + scores_by_source[data_source].append(float(reward["score"])) + + missing_sources = [data_source for data_source, scores in scores_by_source.items() if not scores] + if missing_sources: + raise ValueError(f"Missing evaluation samples for data sources: {missing_sources}") + + all_scores = [score for scores in scores_by_source.values() for score in scores] + metrics = {"score": sum(all_scores) / len(all_scores)} + for data_source, tag in EVAL_DATA_SOURCE_TAGS.items(): + source_scores = scores_by_source[data_source] + metrics[f"{tag}/score"] = sum(source_scores) / len(source_scores) + return metrics + + +# 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 +total_train_steps_env = os.environ.get("TOTAL_TRAIN_STEPS") +total_train_steps = int(total_train_steps_env) if total_train_steps_env is not None else None +train_batch_size = 256 +teacher_num_replicas = int(os.environ.get("MOPD_TEACHER_NUM_REPLICAS", "2")) +teacher_replica_num_workers = int(os.environ.get("MOPD_TEACHER_REPLICA_NUM_WORKERS", "1")) +enable_prefix_caching = os.environ.get("MOPD_ENABLE_PREFIX_CACHING", "0") == "1" +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: four colocated Student workers, plus one GPU per Teacher replica. +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 = Qwen3VLDense2BConfig() +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 = DistillationLossConfig( + 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, + loss_mode="k1", + use_policy_gradient=True, + task_adv_weight=0.0, + distillation_loss_weight=1.0, +) +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, + 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, + ), + ), + } + ), + 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(compute_metric_func=mopd_compute_metric) + +# 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. +distillation_config = DistillationConfig( + loss_config=loss_cfg, + teachers=[ + RolloutTeacherConfig( + name="gsm8k_teacher", + num_replicas=teacher_num_replicas, + enable_prefix_caching=enable_prefix_caching, + launch_config=RolloutTeacherLaunchConfig( + model_path=gsm8k_teacher_model_path, + num_workers=teacher_replica_num_workers, + server_port=13141, + tensor_parallel_size=teacher_replica_num_workers, + expert_parallel_size=1, + context_length=max_num_tokens, + max_batch_size=max_num_tokens, + gpu_memory_utilization=0.8, + ), + ), + RolloutTeacherConfig( + name="geo3k_teacher", + num_replicas=teacher_num_replicas, + enable_prefix_caching=enable_prefix_caching, + launch_config=RolloutTeacherLaunchConfig( + model_path=geo3k_teacher_model_path, + num_workers=teacher_replica_num_workers, + server_port=13142, + tensor_parallel_size=teacher_replica_num_workers, + 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), + distillation_config=distillation_config, + enable_evaluate=enable_evaluate, + enable_initial_evaluate=enable_evaluate, + evaluate_step=evaluate_step, + total_train_steps=total_train_steps, + total_epochs=total_epochs, + work_dir=work_dir, + seed=1234, + debug_rollout=False, +) diff --git a/recipe/on_policy_distillation/config/rl_dapo_math_mopd_async.py b/recipe/on_policy_distillation/config/rl_dapo_math_mopd_async.py new file mode 100644 index 0000000000..974f34af94 --- /dev/null +++ b/recipe/on_policy_distillation/config/rl_dapo_math_mopd_async.py @@ -0,0 +1,402 @@ +import json +import os + +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 Qwen3VLDense2BConfig +from xtuner.v1.rl.advantage import GRPOAdvantageConfig +from xtuner.v1.rl.agent_loop import SingleTurnAgentLoopConfig +from xtuner.v1.rl.agent_loop_manager import ( + AgentLoopManagerConfig, + AsyncProduceStrategyConfig, + SamplerConfig, + TaskSpecConfig, +) +from xtuner.v1.rl.evaluator import EvaluatorConfig +from xtuner.v1.rl.judger import ComposedJudgerConfig, GEO3KJudgerConfig, GSM8KJudgerConfig +from xtuner.v1.rl.distillation import ( + DistillationConfig, + RolloutTeacherConfig, + RolloutTeacherLaunchConfig, +) +from xtuner.v1.rl.loss import DistillationLossConfig +from xtuner.v1.rl.replay_buffer import AsyncReplayBufferConfig +from xtuner.v1.rl.rollout.worker import RolloutConfig +from xtuner.v1.rl.trainer import RolloutImportanceSampling, 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] + + +EVAL_DATA_SOURCE_TAGS = { + "openai/gsm8k": "gsm8k", + "hiyouga/geometry3k": "geo3k", +} + + +def mopd_compute_metric(samples): + scores_by_source = {data_source: [] for data_source in EVAL_DATA_SOURCE_TAGS} + for sample in samples: + data_source = sample.extra_fields.get("origin_data_source") + if data_source not in scores_by_source: + raise ValueError(f"Unexpected evaluation data source: {data_source!r}") + + reward = sample.reward or {} + if "score" not in reward: + raise ValueError(f"Missing reward score for evaluation data source: {data_source!r}") + scores_by_source[data_source].append(float(reward["score"])) + + missing_sources = [data_source for data_source, scores in scores_by_source.items() if not scores] + if missing_sources: + raise ValueError(f"Missing evaluation samples for data sources: {missing_sources}") + + all_scores = [score for scores in scores_by_source.values() for score in scores] + metrics = {"score": sum(all_scores) / len(all_scores)} + for data_source, tag in EVAL_DATA_SOURCE_TAGS.items(): + source_scores = scores_by_source[data_source] + metrics[f"{tag}/score"] = sum(source_scores) / len(source_scores) + return metrics + + +# 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_async" +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: four colocated Student workers, plus one GPU per Teacher. +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 = Qwen3VLDense2BConfig() +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 = DistillationLossConfig( + 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, + loss_mode="k1", + use_policy_gradient=True, + task_adv_weight=0.0, + distillation_loss_weight=1.0, + rollout_is=RolloutImportanceSampling( + rollout_is_level="token", + rollout_is_mode="both", + rollout_is_threshold=(5, 0.5), + rollout_is_mask_threshold=(5, 0.5), + rollout_is_veto_threshold=(20, 0), + ), +) +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 = AsyncProduceStrategyConfig( + over_sample_threshold=1.0, + enable_partial_rollout=True, + max_staleness=2, + max_token_staleness=0, + tail_batch_trigger_size=0 +) +agent_loop_manager_cfg = AgentLoopManagerConfig( + tasks=TaskSpecConfig( + task_name="train_task", + agent_loop_config=agent_loop_config, + 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, + ), + ), + } + ), + 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(compute_metric_func=mopd_compute_metric) + +# 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. +distillation_config = DistillationConfig( + loss_config=loss_cfg, + teachers=[ + RolloutTeacherConfig( + name="gsm8k_teacher", + launch_config=RolloutTeacherLaunchConfig( + model_path=gsm8k_teacher_model_path, + num_workers=1, + server_port=13141, + tensor_parallel_size=1, + expert_parallel_size=1, + context_length=max_num_tokens, + max_batch_size=max_num_tokens, + gpu_memory_utilization=0.8, + ), + ), + RolloutTeacherConfig( + name="geo3k_teacher", + launch_config=RolloutTeacherLaunchConfig( + model_path=geo3k_teacher_model_path, + num_workers=1, + server_port=13142, + 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=AsyncReplayBufferConfig(), + 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), + distillation_config=distillation_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..05e1e026bb --- /dev/null +++ b/recipe/on_policy_distillation/config/rl_dapo_math_opd.py @@ -0,0 +1,249 @@ +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.distillation import ( + DistillationConfig, + RolloutTeacherConfig, + RolloutTeacherLaunchConfig, +) +from xtuner.v1.rl.loss import DistillationLossConfig +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 = DistillationLossConfig( + 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, + loss_mode="k1", + use_policy_gradient=True, + task_adv_weight=0.0, + distillation_loss_weight=1.0, +) +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 +distillation_config = DistillationConfig( + loss_config=loss_cfg, + teachers=[ + RolloutTeacherConfig( + name="teacher", + enable_prefix_caching=True, + launch_config=RolloutTeacherLaunchConfig( + model_path=teacher_model_path, + num_workers=1, + server_port=13141, + ), + ) + ], + 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), + distillation_config=distillation_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..ca47c95f78 --- /dev/null +++ b/recipe/on_policy_distillation/scripts/launch_teacher_utils.sh @@ -0,0 +1,267 @@ +#!/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=() +STUDENT_CUDA_VISIBLE_DEVICES= +XTUNER_RL_NUM_WORKERS= + +start_single_teacher_server() { + _start_teacher_servers "$1" "$2" "$3" "1" +} + +start_teacher_servers() { + _start_teacher_servers "$1" "$2" "$3" "" +} + +wait_for_teacher_servers() { + local startup_timeout_s=$1 + local deadline=$((SECONDS + startup_timeout_s)) + local teacher_index + local pid + local log_file + local -a teacher_ready=() + local -a health_check_pids=() + 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=() + 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[teacher_index]="$!" + done + + for teacher_index in "${!health_check_pids[@]}"; do + if wait "${health_check_pids[teacher_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 teacher_count + local teacher_index + local teacher_field_offset=4 + local teacher_name + local teacher_target_node_rank + local teacher_local_devices + local teacher_endpoint + local teacher_command_arg_count + local teacher_log_file + local device + local student_local_num_workers + local -a teacher_fields=() + local -a teacher_command=() + + # Local non-rjob launches default to a single node. Preserve topology + # values supplied by distributed launchers. + export NODE_COUNT="${NODE_COUNT:-1}" + export NODE_RANK="${NODE_RANK:-0}" + + TEACHER_NAMES=() + TEACHER_ENDPOINTS=() + TEACHER_HEALTH_URLS=() + TEACHER_MODEL_INFO_URLS=() + TEACHER_LOG_FILES=() + TEACHER_PIDS=() + STUDENT_CUDA_VISIBLE_DEVICES= + XTUNER_RL_NUM_WORKERS= + mkdir -p "${work_dir}" + + # Command builder output: + # replica count, endpoint map JSON, Student worker count, current-node + # Student worker count, then repeated display name, target node rank, local devices, + # endpoint, health URL, model-info URL, + # command-argument count, and command argv. + unset XTUNER_OPD_TEACHER_ENDPOINTS_JSON + mapfile -d "" -t teacher_fields < <( + python "${OPD_RECIPE_DIR}/build_teacher_server_commands.py" \ + "${config_file}" "${backend}" + ) + export XTUNER_OPD_TEACHER_ENDPOINTS_JSON="${teacher_fields[1]}" + export XTUNER_RL_NUM_WORKERS="${teacher_fields[2]}" + student_local_num_workers=${teacher_fields[3]} + STUDENT_CUDA_VISIBLE_DEVICES= + for ((device = 0; device < student_local_num_workers; device++)); do + STUDENT_CUDA_VISIBLE_DEVICES+="${STUDENT_CUDA_VISIBLE_DEVICES:+,}${device}" + done + export STUDENT_CUDA_VISIBLE_DEVICES + + 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}" + echo "Student workers: ${XTUNER_RL_NUM_WORKERS}" + echo "Student local GPUs on node ${NODE_RANK}: ${STUDENT_CUDA_VISIBLE_DEVICES:-none}" + for ((teacher_index = 0; teacher_index < teacher_count; teacher_index++)); do + teacher_name=${teacher_fields[teacher_field_offset]} + teacher_target_node_rank=${teacher_fields[teacher_field_offset + 1]} + teacher_local_devices=${teacher_fields[teacher_field_offset + 2]} + teacher_endpoint=${teacher_fields[teacher_field_offset + 3]} + teacher_command_arg_count=${teacher_fields[teacher_field_offset + 6]} + + TEACHER_NAMES+=("${teacher_name}") + TEACHER_ENDPOINTS+=("${teacher_endpoint}") + TEACHER_HEALTH_URLS+=("${teacher_fields[teacher_field_offset + 4]}") + TEACHER_MODEL_INFO_URLS+=("${teacher_fields[teacher_field_offset + 5]}") + + teacher_field_offset=$((teacher_field_offset + 7)) + teacher_command=( + "${teacher_fields[@]:teacher_field_offset:teacher_command_arg_count}" + ) + teacher_field_offset=$((teacher_field_offset + teacher_command_arg_count)) + + if (( teacher_command_arg_count == 0 )); then + if [[ -n "${expected_teacher_count}" ]]; 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 (( NODE_RANK != teacher_target_node_rank )); then + echo \ + "Teacher ${teacher_name} is assigned to node ${teacher_target_node_rank};" \ + "skipping node ${NODE_RANK}." + 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}_node_${NODE_RANK}.log" + fi + TEACHER_LOG_FILES+=("${teacher_log_file}") + + echo "Starting Teacher ${teacher_name}" + echo "Teacher node: ${NODE_RANK}" + echo "Teacher local GPUs: ${teacher_local_devices}" + echo "Teacher endpoint: ${teacher_endpoint}" + echo "Teacher log: ${teacher_log_file}" + + setsid env \ + PYTHONUNBUFFERED=1 \ + CUDA_VISIBLE_DEVICES="${teacher_local_devices}" \ + "${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 +} 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..e9f7e3c927 --- /dev/null +++ b/recipe/on_policy_distillation/scripts/run_sampled_token_mopd.sh @@ -0,0 +1,74 @@ +#!/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 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..212be2a67d --- /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 MODEL_PATH="${STUDENT_MODEL_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 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_teacher_servers "${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..b9d3a98dfd --- /dev/null +++ b/tests/rl/test_on_policy_distillation.py @@ -0,0 +1,335 @@ +import tempfile +import unittest +from pathlib import Path +from unittest.mock import AsyncMock, patch + +import httpx +import torch + +from recipe.on_policy_distillation.build_teacher_server_commands import ( + build_teacher_launch_server_commands, +) +from xtuner.v1.data_proto.rl_data import RolloutState, Status +from xtuner.v1.data_proto.sequence_context import SequenceContext +from xtuner.v1.rl.distillation import RolloutTeacherClient, RolloutTeacherConfig +from xtuner.v1.rl.loss import DistillationLossConfig +from xtuner.v1.rl.trainer.controller import TrainingController + + +class TestDistillationRecipeConfig(unittest.TestCase): + def test_teacher_launcher_reads_distillation_config(self) -> None: + config_source = """ +from xtuner.v1.rl.distillation import ( + DistillationConfig, + RolloutTeacherConfig, + RolloutTeacherLaunchConfig, +) +from xtuner.v1.rl.loss import DistillationLossConfig + +loss_cfg = DistillationLossConfig(policy_loss_cfg={"loss_type": "vanilla"}) +distillation_config = DistillationConfig( + loss_config=loss_cfg, + teachers=[ + RolloutTeacherConfig( + name="teacher", + launch_config=RolloutTeacherLaunchConfig( + model_path="/models/teacher", + num_workers=1, + server_port=13141, + ), + ) + ], + data_source_teacher_map={"math": "teacher"}, +) +""" + with tempfile.TemporaryDirectory() as temp_dir: + config_path = Path(temp_dir) / "distillation_config.py" + config_path.write_text(config_source) + with patch.dict( + "os.environ", + { + "NODE_COUNT": "1", + "NODE_RANK": "0", + "PROC_PER_NODE": "2", + "WORKER_ALL_SOCKET_ADDRS": "127.0.0.1", + }, + ): + endpoint_map, student_num_workers, student_local_num_workers, records = ( + build_teacher_launch_server_commands(str(config_path), "lmdeploy") + ) + + self.assertEqual(endpoint_map, {"teacher": ["http://127.0.0.1:13141"]}) + self.assertEqual(student_num_workers, 1) + self.assertEqual(student_local_num_workers, 1) + self.assertEqual(len(records), 1) + self.assertEqual(records[0][:4], ["teacher[0]", "0", "1", "http://127.0.0.1:13141"]) + self.assertIn("lmdeploy", records[0][7:]) + + +class TestRolloutTeacherClient(unittest.IsolatedAsyncioTestCase): + @staticmethod + def _response(payload: dict) -> httpx.Response: + return httpx.Response( + 200, + request=httpx.Request("POST", "http://teacher/generate"), + json=payload, + ) + + def _build_topk_client(self, *, max_retry_per_sample: int = 0) -> RolloutTeacherClient: + loss_config = DistillationLossConfig( + policy_loss_cfg={"loss_type": "vanilla"}, + loss_mode="forward_kl_topk", + use_policy_gradient=False, + top_k=2, + ) + with patch.dict( + "os.environ", + {"XTUNER_USE_LMDEPLOY": "1", "XTUNER_USE_SGLANG": "0", "XTUNER_USE_VLLM": "0"}, + ): + client = RolloutTeacherClient( + RolloutTeacherConfig( + name="teacher", + endpoints=["http://teacher"], + max_retry_per_sample=max_retry_per_sample, + ), + loss_config, + ) + self.addAsyncCleanup(client._client.aclose) + return client + + @staticmethod + def _state() -> RolloutState: + return RolloutState( + group_id=1, + message=[], + prompt_ids=[10, 11, 12], + response_ids=[13, 14], + status=Status.COMPLETED, + extra_fields={"origin_data_source": "math"}, + ) + + async def test_compute_sampled_token_logprobs_uses_current_interface(self) -> None: + response = httpx.Response( + 200, + request=httpx.Request("POST", "http://teacher/generate"), + json={ + "meta_info": { + "prompt_tokens": 5, + "input_token_logprobs": [ + [-0.1, 11], + [-0.2, 12], + [-0.3, 13], + [-0.4, 14], + ], + } + }, + ) + loss_config = DistillationLossConfig(policy_loss_cfg={"loss_type": "vanilla"}) + with patch.dict( + "os.environ", + {"XTUNER_USE_LMDEPLOY": "1", "XTUNER_USE_SGLANG": "0", "XTUNER_USE_VLLM": "0"}, + ): + client = RolloutTeacherClient( + RolloutTeacherConfig(name="teacher", endpoints=["http://teacher"]), + loss_config, + ) + self.addAsyncCleanup(client._client.aclose) + client._client.post = AsyncMock(return_value=response) + state = RolloutState( + group_id=1, + message=[], + prompt_ids=[10, 11, 12], + response_ids=[13, 14], + status=Status.COMPLETED, + extra_fields={"origin_data_source": "math"}, + ) + + result = await client.compute_logprobs(state) + + self.assertEqual(result.status, Status.COMPLETED) + self.assertEqual(result.teacher_tokens, [13, 14]) + self.assertEqual(result.teacher_logprobs, [-0.3, -0.4]) + self.assertIn("teacher_score_time_s", result.extra_fields) + + async def test_malformed_sampled_response_becomes_failed_state(self) -> None: + response = self._response( + { + "meta_info": { + "input_token_logprobs": [ + [-0.1, 11], + [-0.2, 12], + [-0.3, 13], + [-0.4, 14], + ] + } + } + ) + loss_config = DistillationLossConfig(policy_loss_cfg={"loss_type": "vanilla"}) + with patch.dict( + "os.environ", + {"XTUNER_USE_LMDEPLOY": "1", "XTUNER_USE_SGLANG": "0", "XTUNER_USE_VLLM": "0"}, + ): + client = RolloutTeacherClient( + RolloutTeacherConfig( + name="teacher", + endpoints=["http://teacher"], + max_retry_per_sample=0, + ), + loss_config, + ) + self.addAsyncCleanup(client._client.aclose) + client._client.post = AsyncMock(return_value=response) + + result = await client.compute_logprobs(self._state()) + + self.assertEqual(result.status, Status.FAILED) + self.assertIn("prompt_tokens", result.error_msg or "") + + async def test_sampled_response_rejects_non_numeric_logprob(self) -> None: + response = self._response( + { + "meta_info": { + "prompt_tokens": 5, + "input_token_logprobs": [ + [-0.1, 11], + [-0.2, 12], + [True, 13], + [-0.4, 14], + ], + } + } + ) + loss_config = DistillationLossConfig(policy_loss_cfg={"loss_type": "vanilla"}) + with patch.dict( + "os.environ", + {"XTUNER_USE_LMDEPLOY": "1", "XTUNER_USE_SGLANG": "0", "XTUNER_USE_VLLM": "0"}, + ): + client = RolloutTeacherClient( + RolloutTeacherConfig( + name="teacher", + endpoints=["http://teacher"], + max_retry_per_sample=0, + ), + loss_config, + ) + self.addAsyncCleanup(client._client.aclose) + client._client.post = AsyncMock(return_value=response) + + result = await client.compute_logprobs(self._state()) + + self.assertEqual(result.status, Status.FAILED) + self.assertIn("non-numeric logprob", result.error_msg or "") + + async def test_malformed_topk_responses_become_failed_states(self) -> None: + valid_rows = [ + [[-0.1, 1], [-0.2, 2]], + [[-0.3, 3], [-0.4, 4]], + [[-0.5, 5], [-0.6, 6]], + [[-0.7, 7], [-0.8, 8]], + ] + malformed_payloads = { + "missing_meta_info": {}, + "missing_topk_field": {"meta_info": {"prompt_tokens": 5}}, + "wrong_topk_type": {"meta_info": {"prompt_tokens": 5, "input_top_logprobs": "invalid"}}, + "wrong_row_count": {"meta_info": {"prompt_tokens": 5, "input_top_logprobs": valid_rows[:-1]}}, + "ragged_k": {"meta_info": {"prompt_tokens": 5, "input_top_logprobs": [*valid_rows[:-1], [[-0.7, 7]]]}}, + "invalid_token_id": { + "meta_info": { + "prompt_tokens": 5, + "input_top_logprobs": [*valid_rows[:-1], [[-0.7, "7"], [-0.8, 8]]], + } + }, + } + + for case_name, payload in malformed_payloads.items(): + with self.subTest(case=case_name): + client = self._build_topk_client() + client._client.post = AsyncMock(return_value=self._response(payload)) + + result = await client.compute_logprobs(self._state()) + + self.assertEqual(result.status, Status.FAILED) + self.assertIsNone(result.teacher_tokens) + self.assertIsNone(result.teacher_logprobs) + self.assertIn("last_error=", result.error_msg or "") + client._client.post.assert_awaited_once() + + client = self._build_topk_client() + non_finite_response = httpx.Response( + 200, + request=httpx.Request("POST", "http://teacher/generate"), + headers={"content-type": "application/json"}, + content=( + b'{"meta_info":{"prompt_tokens":5,"input_top_logprobs":' + b"[[[-0.1,1],[-0.2,2]],[[-0.3,3],[-0.4,4]]," + b"[[0.5,5],[-0.6,6]],[[NaN,7],[-0.8,8]]]}}" + ), + ) + client._client.post = AsyncMock(return_value=non_finite_response) + + result = await client.compute_logprobs(self._state()) + + self.assertEqual(result.status, Status.FAILED) + self.assertIn("NaN or Inf", result.error_msg or "") + + async def test_invalid_topk_response_is_retried_before_success(self) -> None: + invalid_response = self._response({"meta_info": {"prompt_tokens": 5}}) + valid_response = self._response( + { + "meta_info": { + "prompt_tokens": 5, + "input_top_logprobs": [ + [[-0.1, 1], [-0.2, 2]], + [[-0.3, 3], [-0.4, 4]], + [[-0.5, 5], [-0.6, 6]], + [[-0.7, 7], [-0.8, 8]], + ], + } + } + ) + client = self._build_topk_client(max_retry_per_sample=1) + client._client.post = AsyncMock(side_effect=[invalid_response, valid_response]) + + with patch("xtuner.v1.rl.distillation.rollout_teacher_client.asyncio.sleep", new=AsyncMock()): + result = await client.compute_logprobs(self._state()) + + self.assertEqual(result.status, Status.COMPLETED) + self.assertEqual(result.teacher_tokens, [[5, 6], [7, 8]]) + self.assertEqual(result.teacher_logprobs, [[-0.5, -0.6], [-0.7, -0.8]]) + self.assertEqual(client._client.post.await_count, 2) + request_payload = client._client.post.await_args.kwargs["json"] + self.assertEqual(request_payload["input_ids"], [10, 11, 12, 13, 14]) + self.assertEqual(request_payload["top_logprobs_num"], 2) + + +class TestTopKTrainingController(unittest.TestCase): + def test_packs_targets_along_sequence_dimension(self) -> None: + controller = TrainingController(workers=[]) + first = { + "seq_ctx": SequenceContext.from_input_ids((torch.tensor([[1, 2]]),), device="cpu"), + "shifted_labels": torch.tensor([[-100, 2]]), + "advantage": [0.0, 0.0], + "rollout_logprobs": torch.zeros(1, 2), + "teacher_logprobs": torch.tensor([[[-0.1, -0.2], [-0.3, -0.4]]]), + "target_token_ids": torch.tensor([[[1, 2], [3, 4]]]), + } + second = { + "seq_ctx": SequenceContext.from_input_ids((torch.tensor([[3]]),), device="cpu"), + "shifted_labels": torch.tensor([[3]]), + "advantage": [0.0], + "rollout_logprobs": torch.zeros(1, 1), + "teacher_logprobs": torch.tensor([[[-0.5, -0.6]]]), + "target_token_ids": torch.tensor([[[5, 6]]]), + } + + packed = controller._packing([first, second], pack_max_length=4, language_cfg=None) + + self.assertEqual(len(packed), 1) + self.assertEqual(packed[0]["teacher_logprobs"].shape, (1, 4, 2)) + self.assertEqual(packed[0]["target_token_ids"].shape, (1, 4, 2)) + torch.testing.assert_close(packed[0]["teacher_logprobs"][0, 3], torch.zeros(2)) + torch.testing.assert_close(packed[0]["target_token_ids"][0, 3], torch.zeros(2, dtype=torch.long)) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/rl/test_prepare_train_data.py b/tests/rl/test_prepare_train_data.py index 3c178b4769..e4fe72ea2f 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 @@ -19,6 +16,8 @@ import torch from xtuner.v1.data_proto.rl_data import RolloutState, Status, reset_rollout_response +from xtuner.v1.rl.distillation import DistillationConfig, RolloutTeacherConfig +from xtuner.v1.rl.loss import DistillationLossConfig from xtuner.v1.train.rl_trainer import BaseRLTrainer @@ -36,6 +35,9 @@ class TestPrepareTrainData(unittest.TestCase): def _build_trainer(self, advantages: list[float]): trainer = BaseRLTrainer.__new__(BaseRLTrainer) trainer._advantage_estimator = _FakeAdvantageEstimator(advantages) + trainer._distillation_config = None + trainer._distillation_loss_cfg = None + trainer._train_teacher_config = None trainer.tokenizer = MagicMock(return_value={"input_ids": torch.tensor([[999]])}) trainer.logger = MagicMock() return trainer @@ -56,6 +58,10 @@ def _state( position_ids: np.ndarray | None = None, mm_info: dict | None = None, extra_fields: dict | None = None, + input_ids: list[int] | None = None, + labels: list[int] | None = None, + teacher_tokens: list[int] | list[list[int]] | None = None, + teacher_logprobs: list[float] | list[list[float]] | None = None, ) -> RolloutState: return RolloutState( rollout_id=uid, @@ -73,12 +79,25 @@ def _state( position_ids=position_ids, mm_info=mm_info, extra_fields=extra_fields or {}, + input_ids=input_ids, + labels=labels, + teacher_tokens=teacher_tokens, + teacher_logprobs=teacher_logprobs, ) def _prepare(self, trainer, data_groups, pack_max_length=128): with patch("xtuner.v1.train.rl_trainer.XTUNER_DETERMINISTIC", True): return trainer._prepare_train_data(data_groups, pack_max_length=pack_max_length) + @staticmethod + def _enable_rollout_distillation(trainer, loss_config: DistillationLossConfig) -> None: + trainer._distillation_config = DistillationConfig( + loss_config=loss_config, + teachers=[RolloutTeacherConfig(name="teacher", endpoints=["http://teacher"])], + data_source_teacher_map={"agent_math": "teacher"}, + ) + trainer._distillation_loss_cfg = loss_config + def test_text_path_builds_shifted_training_tensors(self): # 文本主路径固定 token 布局:input_ids 去掉 response 最后一个 token,label/logprob 对齐预测位置。 trainer = self._build_trainer([1.5]) @@ -102,8 +121,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) @@ -125,7 +144,7 @@ def test_rerolled_state_without_semantic_mask_uses_all_response_tokens(self): self.assertIsNone(state.response_mask) self.assertEqual(data_batches[0]["shifted_labels"].tolist(), [[-100, -100, 30, 31]]) - self.assertEqual(data_batches[0]["advantage"], [1.0, 1.0, 1.0, 1.0, 1.0]) + self.assertEqual(data_batches[0]["advantage"], [0.0, 0.0, 1.0, 1.0]) def test_multi_sample_group_uses_each_sample_reward_and_advantage(self): # 同一个 prompt 下的多个 response 要分别使用自己的 reward 和 advantage。 @@ -136,8 +155,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) @@ -171,6 +190,141 @@ def test_vlm_path_uses_train_prompt_ids_and_preserves_multimodal_fields(self): self.assertEqual(seq_ctx.image_grid_thw.dtype, torch.long) self.assertEqual(seq_ctx.image_grid_thw.tolist(), [[1, 2, 3]]) + def test_agentic_topk_targets_include_token_ids_and_logprobs(self): + loss_config = DistillationLossConfig( + policy_loss_cfg={ + "loss_type": "vanilla", + "cliprange_low": 0.2, + "cliprange_high": 0.2, + }, + loss_mode="reverse", + use_policy_gradient=False, + top_k=2, + ) + trainer = self._build_trainer([0.0]) + self._enable_rollout_distillation(trainer, loss_config) + state = self._state( + input_ids=[10, 11, 20, 21, 22], + labels=[-100, -100, 20, -100, 22], + logprobs=[0.0, -0.1, -0.2, -0.3, -0.4], + teacher_tokens=[[100, 101], [102, 103], [104, 105]], + teacher_logprobs=[[-0.5, -0.6], [-0.7, -0.8], [-0.9, -1.0]], + extra_fields={"origin_data_source": "agent_math"}, + ) + + data_batches, _ = self._prepare(trainer, [[state]]) + + self.assertEqual(len(data_batches), 1) + batch = data_batches[0] + self.assertEqual(batch["shifted_labels"].tolist(), [[-100, 20, -100, 22]]) + self.assertEqual( + batch["target_token_ids"].tolist(), + [[[0, 0], [100, 101], [102, 103], [104, 105]]], + ) + torch.testing.assert_close( + batch["teacher_logprobs"], + torch.tensor( + [[[0.0, 0.0], [-0.5, -0.6], [-0.7, -0.8], [-0.9, -1.0]]], + dtype=torch.float32, + ), + ) + with patch("xtuner.v1.rl.loss.distillation_loss.DEVICE", "cpu"): + loss_ctx = loss_config.build( + { + "shifted_labels": batch["shifted_labels"], + "advantages": torch.tensor([batch["advantage"]], dtype=torch.float32), + "old_logprobs": torch.zeros_like(batch["shifted_labels"], dtype=torch.float32), + "teacher_logprobs": batch["teacher_logprobs"], + "target_token_ids": batch["target_token_ids"], + } + ) + assert loss_ctx is not None + type(loss_ctx).build_batches([loss_ctx]) + loss, _ = loss_ctx.loss_fn( + hidden_states=torch.randn(1, 4, 8), + head_weight=torch.randn(128, 8), + head_bias=None, + loss_kwargs=loss_ctx.loss_kwargs, + ) + self.assertTrue(torch.isfinite(loss)) + + def test_plain_topk_targets_include_masked_response_rows(self): + loss_config = DistillationLossConfig( + policy_loss_cfg={"loss_type": "vanilla"}, + loss_mode="forward_kl_topk", + use_policy_gradient=False, + top_k=2, + ) + trainer = self._build_trainer([0.0]) + self._enable_rollout_distillation(trainer, loss_config) + state = self._state( + prompt_ids=[10, 11, 12], + response_ids=[20, 21, 22], + response_mask=[0, 1, 1], + teacher_tokens=[[100, 101], [102, 103], [104, 105]], + teacher_logprobs=[[-0.5, -0.6], [-0.7, -0.8], [-0.9, -1.0]], + extra_fields={"origin_data_source": "agent_math"}, + ) + + data_batches, _ = self._prepare(trainer, [[state]]) + + batch = data_batches[0] + self.assertEqual(batch["shifted_labels"].tolist(), [[-100, -100, -100, 21, 22]]) + self.assertEqual( + batch["target_token_ids"].tolist(), + [[[0, 0], [0, 0], [100, 101], [102, 103], [104, 105]]], + ) + torch.testing.assert_close( + batch["teacher_logprobs"], + torch.tensor( + [[[0.0, 0.0], [0.0, 0.0], [-0.5, -0.6], [-0.7, -0.8], [-0.9, -1.0]]], + dtype=torch.float32, + ), + ) + + def test_sampled_token_targets_align_for_plain_and_agentic_rollouts(self): + loss_config = DistillationLossConfig( + policy_loss_cfg={"loss_type": "vanilla"}, + loss_mode="k1", + use_policy_gradient=True, + ) + trainer = self._build_trainer([0.0, 0.0]) + self._enable_rollout_distillation(trainer, loss_config) + plain_state = self._state( + uid=1, + group_id=1, + prompt_ids=[10, 11, 12], + response_ids=[20, 21, 22], + response_mask=[0, 1, 1], + teacher_tokens=[20, 21, 22], + teacher_logprobs=[-0.5, -0.7, -0.9], + extra_fields={"origin_data_source": "agent_math"}, + ) + agentic_state = self._state( + uid=2, + group_id=2, + input_ids=[30, 31, 40, 41, 42], + labels=[-100, -100, 40, -100, 42], + logprobs=[0.0, -0.1, -0.2, -0.3, -0.4], + teacher_tokens=[40, 41, 42], + teacher_logprobs=[-1.1, -1.2, -1.3], + extra_fields={"origin_data_source": "agent_math"}, + ) + + data_batches, _ = self._prepare(trainer, [[plain_state], [agentic_state]]) + + self.assertEqual(len(data_batches), 2) + torch.testing.assert_close( + data_batches[0]["teacher_logprobs"], + torch.tensor([[0.0, 0.0, -0.5, -0.7, -0.9]], dtype=torch.float32), + ) + torch.testing.assert_close( + data_batches[1]["teacher_logprobs"], + torch.tensor([[0.0, -1.1, -1.2, -1.3]], dtype=torch.float32), + ) + self.assertNotIn("target_token_ids", data_batches[0]) + self.assertNotIn("target_token_ids", data_batches[1]) + def test_invalid_group_is_skipped(self): # FAILED/FILTERED/ABORTED group 不能进入训练 batch,也不能贡献训练样本数。 trainer = self._build_trainer([1.0]) diff --git a/tests/rl/test_producer.py b/tests/rl/test_producer.py index 89f8674ff5..6013829741 100644 --- a/tests/rl/test_producer.py +++ b/tests/rl/test_producer.py @@ -23,6 +23,7 @@ from unittest.mock import AsyncMock, MagicMock, patch 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), token_stale_threshold=getattr(strategy, "token_stale_threshold", None), expired_groups_retryable=getattr(strategy, "tail_batch_trigger_size", -1) >= 0, @@ -195,7 +199,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), token_stale_threshold=getattr(strategy, "token_stale_threshold", None), expired_groups_retryable=getattr(strategy, "tail_batch_trigger_size", -1) >= 0, @@ -331,8 +334,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 = [] @@ -340,16 +343,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) @@ -420,19 +425,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)) @@ -506,7 +513,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, @@ -517,6 +524,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) @@ -547,7 +555,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 95b46b6ac0..5ee02dd959 100644 --- a/tests/rl/test_rl_colocate_trainer.py +++ b/tests/rl/test_rl_colocate_trainer.py @@ -99,7 +99,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 ec00526144..6c6a96f65e 100644 --- a/xtuner/v1/data_proto/rl_data.py +++ b/xtuner/v1/data_proto/rl_data.py @@ -109,6 +109,10 @@ class RolloutState(BaseModel): tool_calls: list[RolloutToolCall] | None = None response_ids: list[int] | None = None logprobs: list[float] | None = None + # Sampled-token teachers store one value per response token. Top-K + # teachers store one K-wide row per response token. + teacher_tokens: list[int] | list[list[int]] | None = None + teacher_logprobs: list[float] | list[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 +248,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 = None diff --git a/xtuner/v1/loss/__init__.py b/xtuner/v1/loss/__init__.py index d2f20b3a16..51c7596daf 100644 --- a/xtuner/v1/loss/__init__.py +++ b/xtuner/v1/loss/__init__.py @@ -11,7 +11,7 @@ ZLossKwargs, ) from .mtp_loss import MTPLossContext -from .rl_loss import LogProbConfig, LogProbContext +from .rl_loss import LogProbConfig, LogProbContext, TopKLogProbConfig, TopKLogProbContext __all__ = [ @@ -33,6 +33,8 @@ "MTPLossContext", "LogProbConfig", "LogProbContext", + "TopKLogProbConfig", + "TopKLogProbContext", ] import torch diff --git a/xtuner/v1/loss/rl_loss.py b/xtuner/v1/loss/rl_loss.py index b6d15291d1..fd22b2c629 100644 --- a/xtuner/v1/loss/rl_loss.py +++ b/xtuner/v1/loss/rl_loss.py @@ -2,6 +2,7 @@ import torch import torch.nn.functional as F +from pydantic import Field from torch.distributed.device_mesh import DeviceMesh from xtuner.v1.rl.utils.misc import gather_logprobs @@ -91,3 +92,75 @@ def forward( else: logprobs, _ = self.loss_fn(hidden_states, head_weight, head_bias, self.loss_kwargs) return logprobs, (None, {}) + + +class TopKLogProbConfig(BaseLossConfig): + """Select model Top-K tokens and compute their exact full-softmax log + probabilities.""" + + top_k: int = Field(gt=0) + + @property + def loss_ctx_cls(self) -> type["TopKLogProbContext"]: + return TopKLogProbContext + + @property + def _loss_kwargs_cls(self) -> type[BaseLossKwargs]: + return BaseLossKwargs + + def build(self, data: dict, sp_mesh: DeviceMesh | None = None) -> "TopKLogProbContext": + del data, sp_mesh + return self.loss_ctx_cls(self, self._loss_kwargs_cls()) + + +class TopKLogProbContext(BaseLossContext): + """Select model Top-K IDs without storing full log-softmax.""" + + loss_cfg: TopKLogProbConfig + loss_kwargs: BaseLossKwargs + + def loss_fn( + self, + hidden_states: torch.Tensor, + head_weight: torch.Tensor, + head_bias: torch.Tensor | None, + loss_kwargs: BaseLossKwargs, + ) -> tuple[torch.Tensor, tuple[torch.Tensor, dict[str, Any]]]: + del loss_kwargs + logits = F.linear(hidden_states, head_weight, head_bias).float() + selected_logits, token_ids = torch.topk(logits, k=self.loss_cfg.top_k, dim=-1) + selected_logprobs = selected_logits - torch.logsumexp(logits, dim=-1, keepdim=True) + return selected_logprobs, (token_ids, {}) + + def chunk_mode( + self, + hidden_states: torch.Tensor, + head_weight: torch.Tensor, + head_bias: torch.Tensor | None, + loss_kwargs: BaseLossKwargs, + ) -> tuple[torch.Tensor, tuple[torch.Tensor, dict[str, Any]]]: + assert self.loss_cfg.chunk_size is not None, "chunk_size must be set in chunk mode" + + logprob_chunks = [] + token_id_chunks = [] + for hidden_states_chunk in torch.split(hidden_states, self.loss_cfg.chunk_size, dim=1): + logprobs, (token_ids, _) = self.loss_fn( + hidden_states_chunk, + head_weight, + head_bias, + loss_kwargs, + ) + logprob_chunks.append(logprobs) + token_id_chunks.append(token_ids) + return torch.cat(logprob_chunks, dim=1), (torch.cat(token_id_chunks, dim=1), {}) + + def forward( + self, + hidden_states: torch.Tensor, + head_weight: torch.Tensor, + head_bias: torch.Tensor | None = None, + ) -> tuple[torch.Tensor, tuple[torch.Tensor, dict[str, Any]]]: + if self.loss_cfg.mode == "chunk": + return self.chunk_mode(hidden_states, head_weight, head_bias, self.loss_kwargs) + else: + return self.loss_fn(hidden_states, head_weight, head_bias, self.loss_kwargs) diff --git a/xtuner/v1/model/utils/misc.py b/xtuner/v1/model/utils/misc.py index 70fbf2d2d2..67d970e5e6 100644 --- a/xtuner/v1/model/utils/misc.py +++ b/xtuner/v1/model/utils/misc.py @@ -108,6 +108,17 @@ def get(self): "reduced_train_policy_kl1_sum", "reduced_train_policy_kl3_sum", "reduced_train_policy_valid_count", + "reduced_distillation_kl_sum", + "reduced_distillation_abs_loss_sum", + "reduced_distillation_valid_count", + "reduced_opd_reverse_kl_sum", + "reduced_opd_abs_logprob_loss_sum", + "reduced_topk_opd_kl_sum", + "reduced_topk_opd_loss_sum", + "reduced_topk_opd_student_selected_mass_sum", + "reduced_topk_opd_teacher_selected_mass_sum", + "reduced_topk_opd_overlap_fraction_sum", + "reduced_topk_opd_valid_count", ) max_keys = ("max_ratio", "reduced_train_policy_ratio_max") min_keys = ("reduced_train_policy_ratio_min",) diff --git a/xtuner/v1/rl/agent_loop/agent_loop.py b/xtuner/v1/rl/agent_loop/agent_loop.py index 6a69d0f80c..06991fe690 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,7 +10,13 @@ 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.distillation import ( + DistillationConfig, + RolloutTeacherClient, + route_rollout_teacher_client, + validate_opd_sample_params, +) from xtuner.v1.rl.judger import Judger from xtuner.v1.rl.rollout import RolloutController from xtuner.v1.rl.rollout.constants import AGENT_LOOP_RAY_GENERATE_MAX_CONCURRENCY @@ -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, + *, + distillation_config: DistillationConfig | 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_distillation(distillation_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, + distillation_config=distillation_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, + distillation_config=distillation_config, ) @abstractmethod @@ -86,6 +104,7 @@ def _build_ray_actor( pg: PlacementGroup | None = None, judger: Judger | None = None, logger=None, + distillation_config: DistillationConfig | 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, + distillation_config=distillation_config, ), ) @@ -116,6 +136,7 @@ def _build_ray_actors( judger: Judger | None = None, logger=None, start_bundle_idx: int = 0, + distillation_config: DistillationConfig | 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, + distillation_config=distillation_config, ), ) @@ -147,6 +169,7 @@ def _build_router( judger: Judger | None = None, logger=None, start_bundle_idx: int = 0, + distillation_config: DistillationConfig | 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, + distillation_config=distillation_config, ), rollout_ctl=rollout_controller, ) @@ -184,6 +208,19 @@ def __init__( else: self.logger = logger self._judger_pause_event = asyncio.Event() + self.teacher_clients: dict[str, RolloutTeacherClient] = {} + self.data_source_teacher_map: dict[str, str] = {} + + def configure_distillation(self, distillation_config: DistillationConfig | None) -> None: + if distillation_config is None: + return + + validate_opd_sample_params(self.sample_params) + self.teacher_clients = { + teacher.name: RolloutTeacherClient(teacher, distillation_config.loss_config) + for teacher in distillation_config.rollout_teachers + } + self.data_source_teacher_map = dict(distillation_config.data_source_teacher_map) @abstractmethod async def generate_sample(self, rollout_state: RolloutState, **kwargs) -> RolloutState: ... @@ -201,6 +238,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_rollout_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_rollout_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 +366,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 +400,15 @@ def __init__( rollout_controller: RolloutController, judger: Judger | None = None, logger=None, + *, + distillation_config: DistillationConfig | None = None, ): self.agent_loop = agent_loop_config.build_local( rollout_controller=rollout_controller, judger=judger, logger=logger, ) + self.agent_loop.configure_distillation(distillation_config) @ray_method(concurrency_group=AGENT_LOOP_CONCURRENCY_GROUP_GENERATE) async def generate_sample(self, rollout_state: RolloutState, **kwargs) -> RolloutState: @@ -329,6 +418,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/sandbox_agent_loop/agent_in_sandbox_loop.py b/xtuner/v1/rl/agent_loop/sandbox_agent_loop/agent_in_sandbox_loop.py index 8c88a27fb2..52ae3726cb 100644 --- a/xtuner/v1/rl/agent_loop/sandbox_agent_loop/agent_in_sandbox_loop.py +++ b/xtuner/v1/rl/agent_loop/sandbox_agent_loop/agent_in_sandbox_loop.py @@ -6,11 +6,13 @@ import json import traceback import uuid +from collections.abc import Callable from typing import Any, Literal from lagent.utils import create_object -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.distillation import route_rollout_teacher_client from xtuner.v1.rl.judger import Judger from xtuner.v1.rl.rollout import RolloutController from xtuner.v1.rl.utils import create_task @@ -228,6 +230,34 @@ def __init__( self._sample_semaphore = asyncio.Semaphore(max_concurrent_samples) if max_concurrent_samples else None self.mode = mode + async def collect_rollout_group( + self, + rollout_state: list[RolloutState], + *, + is_valid_sample_func: Callable[[list[RolloutState]], bool] | None = None, + **kwargs, + ) -> list[RolloutState]: + """Generate agent traces and score every completed segment.""" + 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 not self.teacher_clients: + return group + + async def score_trace(state: RolloutState) -> RolloutState: + teacher = route_rollout_teacher_client( + state, + data_source_teacher_map=self.data_source_teacher_map, + teacher_clients=self.teacher_clients, + ) + return await teacher.compute_logprobs(state) + + return list(await asyncio.gather(*(create_task(score_trace(state)) for state in group))) + async def generate_group(self, rollout_state: list[RolloutState], **kwargs) -> list[RolloutState]: async def generate_one(state: RolloutState) -> list[RolloutState]: if self._sample_semaphore is None: @@ -287,6 +317,7 @@ async def _build_rollout_states(self, rollout_state: RolloutState, item: AgentRo response_message.get("finish_reason") or ("stop" if item.status == RolloutStatus.COMPLETED else "error") ) rollout_state.reward = _extract_reward_payload(item) + rollout_state.extra_fields["origin_data_source"] = item.data_source rollout_state.extra_fields["agent_status"] = item.status.value selected_agent = _selected_agent(item) if selected_agent is not None: @@ -366,6 +397,7 @@ def _fill_eval_rollout_state(self, rollout_state: RolloutState, item: AgentRollo rollout_state.routed_experts = None rollout_state.response_mask = None rollout_state.response_model_steps = None + rollout_state.extra_fields["origin_data_source"] = item.data_source rollout_state.extra_fields["agent_status"] = item.status.value selected_agent = _selected_agent(item) if selected_agent is not None: 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 07b17e98f4..f4d49a672a 100644 --- a/xtuner/v1/rl/agent_loop_manager/agent_loop_manager.py +++ b/xtuner/v1/rl/agent_loop_manager/agent_loop_manager.py @@ -9,6 +9,7 @@ from transformers import PreTrainedTokenizer, PreTrainedTokenizerFast from xtuner.v1.data_proto.rl_data import Status from xtuner.v1.rl.agent_loop import AgentLoopConfig +from xtuner.v1.rl.distillation import DistillationConfig from xtuner.v1.rl.judger import ComposedJudgerConfig, JudgerConfig, build_judger from xtuner.v1.rl.replay_buffer import ReplayBuffer from xtuner.v1.rl.rollout import RolloutController @@ -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, + distillation_config: DistillationConfig | 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, + distillation_config=distillation_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 f2d246989c..df70b351eb 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 @@ -10,6 +10,7 @@ from transformers import PreTrainedTokenizer, PreTrainedTokenizerFast from xtuner.v1.data_proto.rl_data import Status from xtuner.v1.rl.agent_loop import AgentLoopConfig +from xtuner.v1.rl.distillation import DistillationConfig from xtuner.v1.rl.judger import ComposedJudgerConfig, JudgerConfig, build_judger from xtuner.v1.rl.replay_buffer import ReplayBuffer from xtuner.v1.rl.rollout import RolloutController @@ -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, + distillation_config: DistillationConfig | 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, + distillation_config=distillation_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 a4118e8cdc..140e878804 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 @@ -307,7 +304,6 @@ def build( max_token_staleness=self.max_token_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, ) @@ -315,10 +311,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 @@ -347,10 +341,9 @@ def __init__( max_staleness: int, max_token_staleness: int | None, 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( @@ -431,7 +424,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_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 6639256c05..f95f3a5508 100644 --- a/xtuner/v1/rl/agent_loop_manager/produce_utils.py +++ b/xtuner/v1/rl/agent_loop_manager/produce_utils.py @@ -86,10 +86,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 @@ -121,7 +117,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 expired_groups_retryable: bool = True token_stale_threshold: int | None = None @@ -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,31 +170,22 @@ async def generate_group( async def put_generated_group(self, group: list[RolloutState]) -> bool: produced_tokens = sum(len(item.response_ids or []) - len(item.response_model_steps or []) for item in group) 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: - if item.reward is None or "score" not in item.reward: - logger.warning( - f"Missing reward score in item (rollout_id: {item.rollout_id}) of completed group for task {self.task_name}. This item will be skipped in reward statistics." - ) - continue # TODO: 在 agent 存在一拆多的情况下,这个 raw reward 统计会不准,但是考虑到在这区分有点 hard code,应该暂时不处理 - rewards_sum += float(item.reward["score"]) # type: ignore[index] - rewards_count += 1 + reward = item.reward.get("score") if item.reward is not None else None + if reward is not None: + rewards_sum += float(reward) + 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)) await release_and_discard_rollout_groups([group]) return False @@ -273,13 +262,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 88c7dde548..3fa918eff0 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,10 +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): @@ -191,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``. @@ -261,7 +246,6 @@ def build( max_token_staleness=self.max_token_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, ) @@ -269,10 +253,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 @@ -292,7 +274,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.") @@ -310,7 +292,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() @@ -325,7 +307,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() @@ -341,10 +323,9 @@ def __init__( max_staleness: int, max_token_staleness: int | None, 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逻辑 @@ -413,7 +394,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_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/distillation/__init__.py b/xtuner/v1/rl/distillation/__init__.py new file mode 100644 index 0000000000..85c604ac17 --- /dev/null +++ b/xtuner/v1/rl/distillation/__init__.py @@ -0,0 +1,30 @@ +from .config import ( + DistillationConfig, + RolloutTeacherConfig, + RolloutTeacherLaunchConfig, + TeacherConfig, + TrainTeacherConfig, +) +from .rollout_teacher_client import ( + RolloutTeacherClient, + RolloutTeacherReplicaRouter, + route_rollout_teacher_client, + validate_opd_sample_params, +) +from .train_teacher_manager import TrainTeacherManager, TrainTeacherOutputs, TrainTeacherTimings + + +__all__ = [ + "DistillationConfig", + "RolloutTeacherConfig", + "RolloutTeacherLaunchConfig", + "TeacherConfig", + "TrainTeacherConfig", + "RolloutTeacherClient", + "RolloutTeacherReplicaRouter", + "route_rollout_teacher_client", + "TrainTeacherManager", + "TrainTeacherOutputs", + "TrainTeacherTimings", + "validate_opd_sample_params", +] diff --git a/xtuner/v1/rl/distillation/config.py b/xtuner/v1/rl/distillation/config.py new file mode 100644 index 0000000000..f9557637db --- /dev/null +++ b/xtuner/v1/rl/distillation/config.py @@ -0,0 +1,137 @@ +from __future__ import annotations + +from pathlib import Path +from typing import Literal, cast + +from pydantic import BaseModel, ConfigDict, Field, PrivateAttr, model_validator + +from xtuner.v1.config.fsdp import FSDPConfig +from xtuner.v1.model.base import TransformerConfig +from xtuner.v1.model.compose.base import BaseComposeConfig +from xtuner.v1.rl.loss.distillation_loss import DistillationLossConfig + + +class RolloutTeacherLaunchConfig(BaseModel): + model_config = ConfigDict(extra="forbid") + + model_path: str | Path + num_workers: int = Field(default=1, gt=0) + server_port: int = Field(gt=0, le=65535) + 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) + log_level: Literal["critical", "error", "warning", "info", "debug"] | None = None + 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) + + +class RolloutTeacherConfig(BaseModel): + """Teacher served by an external rollout/inference engine.""" + + model_config = ConfigDict(extra="forbid") + + name: str + num_replicas: int = Field(default=1, gt=0) + endpoints: list[str] = Field(default_factory=list) + 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) + enable_prefix_caching: bool = False + launch_config: RolloutTeacherLaunchConfig | None = None + + +class TrainTeacherConfig(BaseModel): + """Frozen Teacher loaded and executed by each training worker.""" + + model_config = ConfigDict(extra="forbid", arbitrary_types_allowed=True) + + name: str + model_path: str | Path + model_cfg: TransformerConfig | BaseComposeConfig + fsdp_cfg: FSDPConfig | None = None + + +TeacherConfig = RolloutTeacherConfig | TrainTeacherConfig + + +class DistillationConfig(BaseModel): + """Distillation objective, Teacher runtimes, and data-source routing.""" + + model_config = ConfigDict(extra="forbid", arbitrary_types_allowed=True) + + loss_config: DistillationLossConfig + teachers: list[TeacherConfig] = Field(min_length=1) + data_source_teacher_map: dict[str, str] = Field(min_length=1) + _teacher_index_by_data_source: dict[str, int] = PrivateAttr(default_factory=dict) + + @model_validator(mode="after") + def validate_config(self) -> DistillationConfig: + teacher_names = [teacher.name for teacher in self.teachers] + if len(teacher_names) != len(set(teacher_names)): + raise ValueError("Distillation 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)}") + + teacher_types = {type(teacher) for teacher in self.teachers} + if len(teacher_types) != 1: + raise ValueError( + "Distillation requires all teachers to use the same runtime type; " + "mixing RolloutTeacherConfig and TrainTeacherConfig is not supported" + ) + teacher_index_by_name = {teacher.name: index for index, teacher in enumerate(self.teachers)} + self._teacher_index_by_data_source = { + data_source: teacher_index_by_name[teacher_name] + for data_source, teacher_name in self.data_source_teacher_map.items() + } + return self + + @property + def rollout_teachers(self) -> list[RolloutTeacherConfig]: + return [teacher for teacher in self.teachers if isinstance(teacher, RolloutTeacherConfig)] + + @property + def train_teachers(self) -> list[TrainTeacherConfig]: + return [teacher for teacher in self.teachers if isinstance(teacher, TrainTeacherConfig)] + + @property + def teacher_index_by_data_source(self) -> dict[str, int]: + return self._teacher_index_by_data_source + + def validate_student_model(self, student_model_cfg: TransformerConfig | BaseComposeConfig) -> None: + """Validate contracts that require both Student and TrainTeacher + configs.""" + if not self.train_teachers: + return + student_lm_cfg = ( + student_model_cfg.text_config if isinstance(student_model_cfg, BaseComposeConfig) else student_model_cfg + ) + student_vocab_size = cast(TransformerConfig, student_lm_cfg).vocab_size + for teacher in self.train_teachers: + teacher_lm_cfg = ( + teacher.model_cfg.text_config + if isinstance(teacher.model_cfg, BaseComposeConfig) + else teacher.model_cfg + ) + teacher_vocab_size = cast(TransformerConfig, teacher_lm_cfg).vocab_size + if teacher_vocab_size != student_vocab_size: + raise ValueError( + f"distillation teacher {teacher.name!r} vocab_size={teacher_vocab_size} does not match " + f"student vocab_size={student_vocab_size}" + ) + + def resolve_teacher_endpoints( + self, + endpoint_map: dict[str, list[str]], + ) -> DistillationConfig: + teachers = [ + teacher + if not isinstance(teacher, RolloutTeacherConfig) or teacher.launch_config is None + else teacher.model_copy(update={"endpoints": endpoint_map[teacher.name]}) + for teacher in self.teachers + ] + return self.model_copy(update={"teachers": teachers}) diff --git a/xtuner/v1/rl/distillation/rollout_teacher_client.py b/xtuner/v1/rl/distillation/rollout_teacher_client.py new file mode 100644 index 0000000000..bc3bff4a48 --- /dev/null +++ b/xtuner/v1/rl/distillation/rollout_teacher_client.py @@ -0,0 +1,471 @@ +from __future__ import annotations + +import asyncio +import math +import os +import time +from hashlib import blake2b +from typing import Any, Literal, cast + +import httpx + +from xtuner.v1.data_proto.rl_data import RolloutState, SampleParams, Status +from xtuner.v1.rl.loss import DistillationLossConfig +from xtuner.v1.utils import get_logger + +from .config import RolloutTeacherConfig + + +logger = get_logger() + + +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 RolloutTeacherReplicaRouter: + """Resolve a physical teacher replica for one scoring request.""" + + def __init__(self, num_replicas: int) -> None: + if num_replicas <= 0: + raise ValueError(f"num_replicas must be positive, got {num_replicas}") + self._num_replicas = num_replicas + + def resolve_replica_idx( + self, + *, + teacher_name: str, + data_source: str, + group_id: int | None, + ) -> int: + if self._num_replicas == 1: + return 0 + routing_key = f"{teacher_name}\0{data_source}\0{group_id}".encode() + digest = blake2b(routing_key, digest_size=8).digest() + return int.from_bytes(digest, byteorder="big") % self._num_replicas + + +class RolloutTeacherClient: + """Asynchronous teacher client scoped to one AgentLoop.""" + + def __init__(self, config: RolloutTeacherConfig, loss_config: DistillationLossConfig) -> None: + self.config = config + self.loss_config = loss_config + self.name = config.name + self.backend = self._resolve_backend_from_env() + if self.loss_config.uses_topk_targets and self.backend != "lmdeploy": + raise RuntimeError("Rollout Teacher Top-K targets currently require LMDeploy") + self.urls = [f"{endpoint.rstrip('/')}/generate" for endpoint in config.endpoints] + self._semaphores = [asyncio.Semaphore(config.max_concurrency) for _ in self.urls] + self._replica_router = RolloutTeacherReplicaRouter(len(self.urls)) + + 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: + scoring_input = self._prepare_scoring_input(state) + if scoring_input is None: + return state + prompt_ids, response_ids = scoring_input + routed_replica_idx = self._replica_router.resolve_replica_idx( + teacher_name=self.name, + data_source=str(state.extra_fields.get("origin_data_source", "")), + group_id=state.group_id, + ) + image_data = state.extra_fields.get("image_data") + expanded_prompt_len = None + # Recompute the final prompt token so the first response token's + # logprob remains available while the earlier prompt can be reused. + logprob_start_len = 0 + if self.config.enable_prefix_caching: + if not image_data: + logprob_start_len = len(prompt_ids) - 1 + else: + expanded_prompt_len = len(state.extra_fields["train_prompt_ids"]) + logprob_start_len = expanded_prompt_len - 1 + payload = self._construct_payload( + prompt_ids, + response_ids, + logprob_start_len=logprob_start_len, + image_data=image_data, + ) + + attempt_idx = 0 + while True: + replica_idx = (routed_replica_idx + attempt_idx) % len(self.urls) + url = self.urls[replica_idx] + try: + async with self._semaphores[replica_idx]: + response = await self._client.post(url, json=payload) + response.raise_for_status() + teacher_tokens: list[int] | list[list[int]] + teacher_logprobs: list[float] | list[list[float]] + if self.loss_config.uses_topk_targets: + teacher_tokens, teacher_logprobs = self._parse_topk_response( + response, + response_ids, + logprob_start_len=logprob_start_len, + expanded_prompt_len=expanded_prompt_len, + ) + else: + teacher_tokens, teacher_logprobs = self._parse_response( + response, + response_ids, + logprob_start_len=logprob_start_len, + expanded_prompt_len=expanded_prompt_len, + ) + state.teacher_tokens = teacher_tokens + state.teacher_logprobs = teacher_logprobs + return state + except (httpx.HTTPStatusError, httpx.RequestError, ValueError) as exc: + if attempt_idx >= self.config.max_retry_per_sample: + state.status = Status.FAILED + state.error_msg = ( + f"Teacher {self.name!r} logprobs calculation failed after {attempt_idx + 1} attempts; " + f"group_id={state.group_id}; routed_replica={routed_replica_idx}; " + f"replica={replica_idx}; endpoint={url}; last_error={exc}" + ) + logger.warning(state.error_msg) + return state + attempt_idx += 1 + await asyncio.sleep(0.1) + finally: + state.extra_fields["teacher_score_time_s"] = time.perf_counter() - start + + def _prepare_scoring_input(self, state: RolloutState) -> tuple[list[int], list[int]] | None: + if state.input_ids is not None: + input_ids = state.input_ids + labels = state.labels + if len(input_ids) < 2: + state.status = Status.FAILED + state.error_msg = f"Teacher {self.name!r} trace scoring requires at least two input_ids" + return None + if labels is None or len(labels) != len(input_ids): + state.status = Status.FAILED + state.error_msg = ( + f"Teacher {self.name!r} trace scoring requires input_ids and labels with equal lengths; " + f"got {len(input_ids)} and {None if labels is None else len(labels)}" + ) + return None + scoring_start = next((index for index, label in enumerate(labels[1:], start=1) if label != -100), None) + if scoring_start is None: + state.status = Status.FAILED + state.error_msg = f"Teacher {self.name!r} trace scoring requires at least one trainable label" + return None + # Keep later masked turns in the scored suffix so subsequent + # assistant turns retain their complete causal context. + return input_ids[:scoring_start], input_ids[scoring_start:] + + prompt_ids = cast(list[int] | None, state.prompt_ids) + response_ids = cast(list[int] | None, 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 None + return prompt_ids, response_ids + + @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("RolloutTeacherClient 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], + *, + logprob_start_len: int, + image_data: Any | None = None, + ) -> dict[str, Any]: + if self.backend == "sglang": + payload = self._construct_sglang_payload( + prompt_ids, + response_ids, + logprob_start_len=logprob_start_len, + ) + elif self.backend == "lmdeploy": + payload = self._construct_lmdeploy_payload( + prompt_ids, + response_ids, + logprob_start_len=logprob_start_len, + ) + else: + raise RuntimeError(f"Unsupported teacher backend: {self.backend}") + if image_data: + payload["image_data"] = image_data + if self.loss_config.uses_topk_targets: + payload["top_logprobs_num"] = cast(int, self.loss_config.top_k) + return payload + + @staticmethod + def _construct_sglang_payload( + prompt_ids: list[int], + response_ids: list[int], + *, + logprob_start_len: 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": logprob_start_len, + "top_logprobs_num": 0, + "stream": False, + } + + @staticmethod + def _construct_lmdeploy_payload( + prompt_ids: list[int], + response_ids: list[int], + *, + logprob_start_len: int, + ) -> dict[str, Any]: + return { + "input_ids": prompt_ids + response_ids, + "return_logprob": True, + "logprob_start_len": logprob_start_len, + "max_tokens": 0, + "stream": False, + } + + def _parse_response( + self, + response: httpx.Response, + response_ids: list[int], + *, + logprob_start_len: int, + expanded_prompt_len: int | None = None, + ) -> tuple[list[int], list[float]]: + if self.backend == "sglang": + response_logprobs = self._parse_sglang_response( + response, + response_ids, + logprob_start_len=logprob_start_len, + expanded_prompt_len=expanded_prompt_len, + ) + else: + response_logprobs = self._parse_lmdeploy_response( + response, + response_ids, + logprob_start_len=logprob_start_len, + expanded_prompt_len=expanded_prompt_len, + ) + return self._validate_response_logprobs(response_logprobs, response_ids) + + def _parse_topk_response( + self, + response: httpx.Response, + response_ids: list[int], + *, + logprob_start_len: int, + expanded_prompt_len: int | None = None, + ) -> tuple[list[list[int]], list[list[float]]]: + meta_info = self._get_response_meta_info(response) + prompt_token_count = self._get_prompt_token_count(meta_info) + if expanded_prompt_len is not None: + expected_prompt_token_count = expanded_prompt_len + len(response_ids) + if prompt_token_count != expected_prompt_token_count: + raise ValueError( + "LMDeploy expanded prompt length mismatch: " + f"expected {expected_prompt_token_count} total input tokens " + f"({expanded_prompt_len} prompt + {len(response_ids)} response), " + f"got {prompt_token_count}" + ) + + raw_topk = self._get_meta_list(meta_info, "input_top_logprobs") + expected_rows = prompt_token_count - logprob_start_len - 1 + if len(raw_topk) != expected_rows: + raise ValueError( + "LMDeploy teacher Top-K length mismatch: " + f"expected {expected_rows} rows for logprob_start_len={logprob_start_len}, got {len(raw_topk)}" + ) + if len(raw_topk) < len(response_ids): + raise ValueError( + "LMDeploy teacher Top-K response is shorter than the scored response: " + f"{len(raw_topk)} vs {len(response_ids)}" + ) + + top_k = cast(int, self.loss_config.top_k) + teacher_tokens: list[list[int]] = [] + teacher_logprobs: list[list[float]] = [] + for row_idx, row in enumerate(raw_topk[-len(response_ids) :]): + if not isinstance(row, list) or len(row) != top_k: + raise ValueError(f"Teacher Top-K row {row_idx} must contain exactly {top_k} entries") + + token_row: list[int] = [] + logprob_row: list[float] = [] + for entry_idx, entry in enumerate(row): + if not isinstance(entry, (list, tuple)) or len(entry) != 2: + raise ValueError(f"Teacher Top-K row {row_idx} entry {entry_idx} must be [logprob, token_id]") + raw_logprob, raw_token_id = entry + if isinstance(raw_logprob, bool) or not isinstance(raw_logprob, (int, float)): + raise ValueError(f"Teacher Top-K row {row_idx} entry {entry_idx} has a non-numeric logprob") + if isinstance(raw_token_id, bool) or not isinstance(raw_token_id, int) or raw_token_id < 0: + raise ValueError( + f"Teacher Top-K row {row_idx} entry {entry_idx} has an invalid non-negative token id" + ) + logprob = float(raw_logprob) + if not math.isfinite(logprob): + raise ValueError("Teacher Top-K logprobs contain NaN or Inf") + logprob_row.append(logprob) + token_row.append(raw_token_id) + if len(token_row) != len(set(token_row)): + raise ValueError(f"Teacher Top-K row {row_idx} contains duplicate token ids") + teacher_tokens.append(token_row) + teacher_logprobs.append(logprob_row) + return teacher_tokens, teacher_logprobs + + @staticmethod + def _parse_sglang_response( + response: httpx.Response, + response_ids: list[int], + *, + logprob_start_len: int, + expanded_prompt_len: int | None = None, + ) -> list[Any]: + meta_info = RolloutTeacherClient._get_response_meta_info(response) + raw_logprobs = RolloutTeacherClient._get_meta_list(meta_info, "input_token_logprobs") + prompt_token_count = RolloutTeacherClient._get_prompt_token_count(meta_info) + if expanded_prompt_len is not None: + expected_prompt_token_count = expanded_prompt_len + len(response_ids) + if prompt_token_count != expected_prompt_token_count: + raise ValueError( + "SGLang expanded prompt length mismatch: " + f"expected {expected_prompt_token_count} total input tokens " + f"({expanded_prompt_len} prompt + {len(response_ids)} response), " + f"got {prompt_token_count}" + ) + # SGLang includes an unscorable placeholder at the requested boundary, + # so N processed tokens with boundary S produce N-S rows. + expected_rows = prompt_token_count - logprob_start_len + if len(raw_logprobs) != expected_rows: + raise ValueError( + "SGLang teacher logprob length mismatch: " + f"expected {expected_rows} rows for logprob_start_len={logprob_start_len}, " + f"got {len(raw_logprobs)}" + ) + return raw_logprobs[-len(response_ids) :] + + @staticmethod + def _parse_lmdeploy_response( + response: httpx.Response, + response_ids: list[int], + *, + logprob_start_len: int, + expanded_prompt_len: int | None = None, + ) -> list[Any]: + meta_info = RolloutTeacherClient._get_response_meta_info(response) + raw_logprobs = RolloutTeacherClient._get_meta_list(meta_info, "input_token_logprobs") + prompt_token_count = RolloutTeacherClient._get_prompt_token_count(meta_info) + if expanded_prompt_len is not None: + expected_prompt_token_count = expanded_prompt_len + len(response_ids) + if prompt_token_count != expected_prompt_token_count: + raise ValueError( + "LMDeploy expanded prompt length mismatch: " + f"expected {expected_prompt_token_count} total input tokens " + f"({expanded_prompt_len} prompt + {len(response_ids)} response), " + f"got {prompt_token_count}" + ) + # LMDeploy omits the unscorable boundary row, hence one fewer row than + # SGLang for the same processed input and logprob boundary. + expected_rows = prompt_token_count - logprob_start_len - 1 + if len(raw_logprobs) != expected_rows: + raise ValueError( + "LMDeploy teacher logprob length mismatch: " + f"expected {expected_rows} rows for logprob_start_len={logprob_start_len}, " + f"got {len(raw_logprobs)}" + ) + return raw_logprobs[-len(response_ids) :] + + @staticmethod + def _get_response_meta_info(response: httpx.Response) -> dict[str, Any]: + try: + payload = response.json() + except (TypeError, ValueError) as exc: + raise ValueError("Invalid teacher response") from exc + if not isinstance(payload, dict) or not isinstance(payload.get("meta_info"), dict): + raise ValueError("Invalid teacher response") + return cast(dict[str, Any], payload["meta_info"]) + + @staticmethod + def _get_prompt_token_count(meta_info: dict[str, Any]) -> int: + prompt_token_count = meta_info.get("prompt_tokens") + if isinstance(prompt_token_count, bool) or not isinstance(prompt_token_count, int) or prompt_token_count < 1: + raise ValueError("Invalid teacher response prompt_tokens") + return prompt_token_count + + @staticmethod + def _get_meta_list(meta_info: dict[str, Any], field_name: str) -> list[Any]: + value = meta_info.get(field_name) + if not isinstance(value, list): + raise ValueError(f"Invalid teacher response {field_name}") + return value + + @staticmethod + def _validate_response_logprobs( + response_logprobs: list[Any], + response_ids: list[int], + ) -> tuple[list[int], list[float]]: + teacher_tokens: list[int] = [] + teacher_logprobs: list[float] = [] + for row_idx, item in enumerate(response_logprobs): + if not isinstance(item, (list, tuple)) or len(item) != 2: + raise ValueError(f"Teacher logprob row {row_idx} must be [logprob, token_id]") + raw_logprob, raw_token_id = item + if isinstance(raw_logprob, bool) or not isinstance(raw_logprob, (int, float)): + raise ValueError(f"Teacher logprob row {row_idx} has a non-numeric logprob") + if isinstance(raw_token_id, bool) or not isinstance(raw_token_id, int) or raw_token_id < 0: + raise ValueError(f"Teacher logprob row {row_idx} has an invalid non-negative token id") + teacher_tokens.append(raw_token_id) + teacher_logprobs.append(float(raw_logprob)) + + 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_rollout_teacher_client( + state: RolloutState, + *, + data_source_teacher_map: dict[str, str], + teacher_clients: dict[str, RolloutTeacherClient], +) -> RolloutTeacherClient: + data_source = state.extra_fields["origin_data_source"] + teacher_name = data_source_teacher_map[data_source] + return teacher_clients[teacher_name] diff --git a/xtuner/v1/rl/distillation/train_teacher_manager.py b/xtuner/v1/rl/distillation/train_teacher_manager.py new file mode 100644 index 0000000000..b441d5fbb6 --- /dev/null +++ b/xtuner/v1/rl/distillation/train_teacher_manager.py @@ -0,0 +1,247 @@ +from __future__ import annotations + +import time +from contextlib import contextmanager +from dataclasses import dataclass, field +from typing import Iterator, cast + +import torch + +from xtuner.v1.data_proto.sequence_context import DSATopKCacheState, SequenceContext +from xtuner.v1.loss import LogProbConfig, LogProbContext, TopKLogProbConfig +from xtuner.v1.model.compose.base import BaseComposeConfig +from xtuner.v1.rl.loss import DistillationLossConfig +from xtuner.v1.rl.model_utils import FrozenModel, build_frozen_model +from xtuner.v1.utils import get_device, get_torch_device_module + +from .config import DistillationConfig + + +DEVICE = get_device() +DEVICE_MODULE = get_torch_device_module() + + +@dataclass +class TrainTeacherTimings: + """Wall-clock seconds spent in the frozen Teacher lifecycle.""" + + compute: float = 0.0 + onload: float = 0.0 + offload: float = 0.0 + + +@dataclass +class TrainTeacherOutputs: + teacher_logprobs: list[torch.Tensor] + target_token_ids: list[torch.Tensor] | None = None + timings: TrainTeacherTimings = field(default_factory=TrainTeacherTimings) + + +class TrainTeacherManager: + """Execute training-side Teachers in one TrainingWorker process. + + The manager owns Teacher model construction, deterministic Teacher-major scheduling, CPU/device residency, and + sampled-token or top-k output calculation. The caller remains responsible for preparing distributed inputs and + swapping the Actor and optimizer around the Teacher phase. + """ + + def __init__(self, distillation_config: DistillationConfig, *, chunk_size: int | None) -> None: + self.loss_config = cast(DistillationLossConfig, distillation_config.loss_config) + mode = "chunk" if chunk_size is not None else "eager" + self.logprob_config = LogProbConfig(chunk_size=chunk_size, mode=mode) + self.topk_logprob_config: TopKLogProbConfig | None = None + if self.loss_config.uses_topk_targets: + self.topk_logprob_config = TopKLogProbConfig( + top_k=cast(int, self.loss_config.top_k), + chunk_size=chunk_size, + mode=mode, + ) + + # Build every frozen Teacher before the Actor so checkpoint/config + # errors fail during worker initialization. Each Teacher is offloaded + # immediately, preventing multiple full models from co-residing on GPU. + self._teachers: list[FrozenModel] = [] + for teacher_config in distillation_config.train_teachers: + teacher = build_frozen_model( + teacher_config.model_cfg, + teacher_config.model_path, + teacher_config.fsdp_cfg, + ) + self._teachers.append(teacher) + + self._teacher_is_composed = [ + isinstance(teacher.model_cfg, BaseComposeConfig) for teacher in distillation_config.train_teachers + ] + self._teacher_index_by_name = { + teacher_config.name: teacher_index + for teacher_index, teacher_config in enumerate(distillation_config.train_teachers) + } + + def compute_logprobs( + self, + *, + seq_ctx_list: list[SequenceContext], + shifted_labels_list: list[torch.Tensor], + teacher_indices_list: list[torch.Tensor], + ) -> TrainTeacherOutputs: + timings = TrainTeacherTimings() + if self.loss_config.uses_sampled_token_targets: + return TrainTeacherOutputs( + teacher_logprobs=self._compute_sampled_logprobs( + seq_ctx_list, + shifted_labels_list, + teacher_indices_list, + timings, + ), + timings=timings, + ) + + target_token_ids, teacher_logprobs = self._compute_topk_targets( + seq_ctx_list, + teacher_indices_list, + timings, + ) + return TrainTeacherOutputs( + teacher_logprobs=teacher_logprobs, + target_token_ids=target_token_ids, + timings=timings, + ) + + def offload_all_to_cpu(self) -> None: + for teacher in self._teachers: + self._offload_to_cpu(teacher) + + def offload_to_disk(self, teacher_name: str) -> None: + """Reserve the disk-offload lifecycle boundary for a later backend.""" + if teacher_name not in self._teacher_index_by_name: + raise KeyError(f"Unknown training Teacher: {teacher_name!r}") + raise NotImplementedError("Train Teacher disk offload is not implemented") + + @staticmethod + def _offload_to_cpu(teacher: FrozenModel) -> None: + teacher.to_device("cpu") + if hasattr(DEVICE_MODULE, "empty_cache"): + DEVICE_MODULE.empty_cache() + + @staticmethod + def _synchronize_device() -> None: + if str(DEVICE) != "cpu" and hasattr(DEVICE_MODULE, "synchronize"): + DEVICE_MODULE.synchronize() + + @contextmanager + def _teacher_on_device( + self, + teacher: FrozenModel, + timings: TrainTeacherTimings, + ) -> Iterator[None]: + onload_begin = time.perf_counter() + teacher.to_device(DEVICE) + timings.onload += time.perf_counter() - onload_begin + + compute_begin = time.perf_counter() + try: + yield + self._synchronize_device() + finally: + timings.compute += time.perf_counter() - compute_begin + offload_begin = time.perf_counter() + self._offload_to_cpu(teacher) + timings.offload += time.perf_counter() - offload_begin + + @staticmethod + def _teacher_seq_ctx(seq_ctx: SequenceContext, *, is_composed: bool) -> SequenceContext: + overrides = { + "rollout_routed_experts": None, + "offload_rollout_routed_experts": False, + "dsa_topk_cache": DSATopKCacheState(), + } + if not is_composed: + # A VLM tokenizer supplies 3D M-RoPE position ids for every sample, + # including text-only samples. Plain language-model Teachers need + # standard 2D packed positions. Passing ``None`` makes + # SequenceContext rebuild them from the packed sequence lengths. + # Visual fields are irrelevant to the text Teacher and may refer to + # a pack routed to another Teacher, so do not retain them here. + overrides.update( + position_ids=None, + image_grid_thw=None, + deepstack_visual_embeds=None, + visual_pos_masks=None, + pixel_values=None, + inputs_embeds=None, + num_img_tokens=None, + ) + return seq_ctx.copy(**overrides) + + def _compute_topk_targets( + self, + seq_ctx_list: list[SequenceContext], + teacher_indices_list: list[torch.Tensor], + timings: TrainTeacherTimings, + ) -> tuple[list[torch.Tensor], list[torch.Tensor]]: + top_k = cast(int, self.loss_config.top_k) + topk_logprob_config = cast(TopKLogProbConfig, self.topk_logprob_config) + target_ids = [ + torch.zeros((*teacher_indices.shape, top_k), dtype=torch.long, device=DEVICE) + for teacher_indices in teacher_indices_list + ] + target_logprobs = [ + torch.zeros((*teacher_indices.shape, top_k), dtype=torch.float32, device=DEVICE) + for teacher_indices in teacher_indices_list + ] + + # All ranks iterate Teachers and local packs in the same order so every + # FSDP rank enters the same collective sequence, even when a rank has no + # tokens routed to a particular Teacher. + for teacher_index, teacher in enumerate(self._teachers): + with self._teacher_on_device(teacher, timings): + # Every rank forwards every local pack before selecting routed + # tokens so all ranks enter the same FSDP collective sequence. + for batch_index, (seq_ctx, teacher_indices) in enumerate(zip(seq_ctx_list, teacher_indices_list)): + loss_ctx = topk_logprob_config.build(data={}) + assert loss_ctx is not None + with torch.no_grad(): + output = teacher( + seq_ctx=self._teacher_seq_ctx( + seq_ctx, + is_composed=self._teacher_is_composed[teacher_index], + ), + loss_ctx={"lm": loss_ctx}, + ) + selected = teacher_indices == teacher_index + target_ids[batch_index][selected] = cast(torch.Tensor, output.logits)[selected] + target_logprobs[batch_index][selected] = cast(torch.Tensor, output.loss)[selected] + return target_ids, target_logprobs + + def _compute_sampled_logprobs( + self, + seq_ctx_list: list[SequenceContext], + shifted_labels_list: list[torch.Tensor], + teacher_indices_list: list[torch.Tensor], + timings: TrainTeacherTimings, + ) -> list[torch.Tensor]: + target_logprobs = [ + torch.zeros_like(shifted_labels, dtype=torch.float32) for shifted_labels in shifted_labels_list + ] + + # Keep the Teacher-major schedule identical across ranks for FSDP. + for teacher_index, teacher in enumerate(self._teachers): + with self._teacher_on_device(teacher, timings): + for batch_index, (seq_ctx, shifted_labels, teacher_indices) in enumerate( + zip(seq_ctx_list, shifted_labels_list, teacher_indices_list) + ): + loss_ctx = cast( + LogProbContext, + self.logprob_config.build(data={"shifted_labels": shifted_labels}), + ) + with torch.no_grad(): + output = teacher( + seq_ctx=self._teacher_seq_ctx( + seq_ctx, + is_composed=self._teacher_is_composed[teacher_index], + ), + loss_ctx={"lm": loss_ctx}, + ) + selected = teacher_indices == teacher_index + target_logprobs[batch_index][selected] = cast(torch.Tensor, output.loss)[selected] + return target_logprobs diff --git a/xtuner/v1/rl/loss/__init__.py b/xtuner/v1/rl/loss/__init__.py index 371ea27bff..45d3e1d069 100644 --- a/xtuner/v1/rl/loss/__init__.py +++ b/xtuner/v1/rl/loss/__init__.py @@ -5,6 +5,13 @@ compute_kl_loss_weight, finalize_train_policy_metrics, ) +from .distillation_loss import ( + DistillationLossConfig, + DistillationLossContext, + DistillationLossKwargs, + compute_topk_distillation_kl, + finalize_distillation_metrics, +) from .grpo_loss import GRPOLossConfig, GRPOLossContext, GRPOLossKwargs from .loss_fn import check_config, get_policy_loss_fn, kl_penalty, pg_loss_fn, register_policy_loss, sft_loss_fn from .oreal_loss import OrealLossConfig, OrealLossContext, OrealLossKwargs diff --git a/xtuner/v1/rl/loss/base_loss.py b/xtuner/v1/rl/loss/base_loss.py index 108af41ed8..e2eccc9676 100644 --- a/xtuner/v1/rl/loss/base_loss.py +++ b/xtuner/v1/rl/loss/base_loss.py @@ -304,4 +304,5 @@ def reduce_values(keys, op): extra_info_dict["reduced_train_policy_ratio_min"] = ratio_min # legacy metric,keep here extra_info_dict["max_ratio"] = ratio_max + return extra_info_dict diff --git a/xtuner/v1/rl/loss/distillation_loss.py b/xtuner/v1/rl/loss/distillation_loss.py new file mode 100644 index 0000000000..e655c06ec7 --- /dev/null +++ b/xtuner/v1/rl/loss/distillation_loss.py @@ -0,0 +1,447 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from typing import Any, Literal, cast + +import torch +import torch.distributed as dist +import torch.nn.functional as F +from pydantic import Field, model_validator +from torch.distributed.device_mesh import DeviceMesh +from typing_extensions import Self + +from xtuner.v1.loss.utils import sp_split +from xtuner.v1.utils import get_logger +from xtuner.v1.utils.device import get_device + +from ..utils import gather_logprobs +from .base_loss import ( + BaseRLLossConfig, + BaseRLLossContext, + BaseRLLossKwargs, + compute_kl_loss_weight, +) +from .loss_fn import get_policy_loss_fn, kl_penalty + + +DEVICE = get_device() +logger = get_logger() + +TopKDistillationMode = Literal["forward", "reverse", "forward_kl_topk"] +DistillationLossMode = Literal["k1", "k3", "forward", "reverse", "forward_kl_topk"] + + +def compute_topk_distillation_kl( + student_logprobs: torch.Tensor, + teacher_logprobs: torch.Tensor, + loss_mode: TopKDistillationMode, + log_prob_min_clamp: float | None = None, + loss_max_clamp: float | None = None, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Compute a KL loss over teacher-selected Top-K tokens.""" + student_probs = student_logprobs.exp() + teacher_probs = teacher_logprobs.exp() + student_selected_mass = student_probs.sum(dim=-1) + teacher_selected_mass = teacher_probs.sum(dim=-1) + + if loss_mode == "forward_kl_topk": + if log_prob_min_clamp is not None: + student_logprobs = student_logprobs.clamp_min(log_prob_min_clamp) + teacher_logprobs = teacher_logprobs.clamp_min(log_prob_min_clamp) + teacher_probs = teacher_logprobs.exp() + loss = (teacher_probs * (teacher_logprobs - student_logprobs)).sum(dim=-1) + if loss_max_clamp is not None: + loss = loss.clamp(min=-loss_max_clamp, max=loss_max_clamp) + return loss, student_selected_mass, teacher_selected_mass + + eps = torch.finfo(student_logprobs.dtype).tiny + student_tail = (1.0 - student_selected_mass).clamp_min(eps) + teacher_tail = (1.0 - teacher_selected_mass).clamp_min(eps) + if loss_mode == "forward": + selected_kl = teacher_probs * (teacher_logprobs - student_logprobs) + tail_kl = teacher_tail * (teacher_tail.log() - student_tail.log()) + else: + selected_kl = student_probs * (student_logprobs - teacher_logprobs) + tail_kl = student_tail * (student_tail.log() - teacher_tail.log()) + loss = selected_kl.sum(dim=-1) + tail_kl + if loss_max_clamp is not None: + loss = loss.clamp(max=loss_max_clamp) + return loss, student_selected_mass, teacher_selected_mass + + +class DistillationLossConfig(BaseRLLossConfig): + """Configuration shared by sampled-token PG-OPD and direct GKD losses.""" + + loss_mode: DistillationLossMode = "k1" + use_policy_gradient: bool = True + task_adv_weight: float = Field(default=0.0, ge=0.0) + distillation_loss_weight: float = Field(default=1.0, ge=0.0) + top_k: int | None = Field(default=None, gt=0) + log_prob_min_clamp: float | None = None + loss_max_clamp: float | None = Field(default=None, gt=0.0) + + @model_validator(mode="after") + def validate_loss_options(self) -> Self: + if self.loss_mode in ("forward", "reverse", "forward_kl_topk"): + if self.use_policy_gradient: + raise ValueError(f"{self.loss_mode} only supports direct backpropagation") + if self.top_k is None: + raise ValueError(f"{self.loss_mode} requires top_k") + elif self.top_k is not None: + raise ValueError(f"top_k is not used by loss_mode={self.loss_mode}") + + if self.loss_mode == "k1" and not self.use_policy_gradient: + raise ValueError("k1 only supports the policy-gradient path") + return self + + @property + def uses_sampled_token_targets(self) -> bool: + return self.loss_mode in ("k1", "k3") + + @property + def uses_topk_targets(self) -> bool: + return not self.uses_sampled_token_targets + + @property + def loss_ctx_cls(self) -> type["DistillationLossContext"]: + return DistillationLossContext + + @property + def _loss_kwargs_cls(self) -> type["DistillationLossKwargs"]: + return DistillationLossKwargs + + def build( + self, + data: dict, + sp_mesh: DeviceMesh | None = None, + ) -> "DistillationLossContext | None": + if "shifted_labels" not in data or "advantages" not in data: + return None + + loss_kwargs = DistillationLossKwargs( + shifted_labels=data["shifted_labels"], + advantages=data["advantages"], + rollout_logprobs=data.get("rollout_logprobs"), + old_logprobs=data.get("old_logprobs"), + ref_logprobs=data.get("ref_logprobs"), + is_weights=data.get("rollout_is_weights"), + teacher_logprobs=data.get("teacher_logprobs"), + target_token_ids=data.get("target_token_ids"), + ).to(DEVICE) + if sp_mesh is not None and sp_mesh.size() > 1: + loss_kwargs = loss_kwargs.sp_split(sp_mesh) + return self.loss_ctx_cls(self, loss_kwargs) + + +class DistillationLossKwargs(BaseRLLossKwargs): + teacher_logprobs: torch.Tensor | None = None + target_token_ids: torch.Tensor | None = None + distillation_loss_weight: torch.Tensor | None = None + + def sp_split(self, sp_mesh: DeviceMesh) -> Self: + super().sp_split(sp_mesh) + 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.target_token_ids is not None: + self.target_token_ids = sp_split( + self.target_token_ids, + sp_mesh=sp_mesh, + split_dim=1, + padding_value=0, + ) + return self + + def to(self, device: torch.device | str) -> Self: + super().to(device) + if self.teacher_logprobs is not None: + self.teacher_logprobs = self.teacher_logprobs.to(device) + if self.target_token_ids is not None: + self.target_token_ids = self.target_token_ids.to(device) + if self.distillation_loss_weight is not None: + self.distillation_loss_weight = self.distillation_loss_weight.to(device) + return self + + +class DistillationLossContext(BaseRLLossContext): + loss_cfg: DistillationLossConfig + loss_kwargs: DistillationLossKwargs + + def __init__(self, loss_cfg: DistillationLossConfig, loss_kwargs: DistillationLossKwargs): + super().__init__(loss_cfg, loss_kwargs) + self.policy_loss_fn = get_policy_loss_fn(self.loss_cfg.policy_loss_cfg.get("loss_type", "vanilla")) + + @staticmethod + def build_batches( # type: ignore[override] + loss_ctx_list: list["DistillationLossContext"], + ) -> list["DistillationLossContext"]: + assert loss_ctx_list, "loss_ctx_list can not be empty" + loss_cfg = loss_ctx_list[0].loss_cfg + shifted_labels_list = [loss_ctx.loss_kwargs.shifted_labels for loss_ctx in loss_ctx_list] + rank_grad_tokens = sum((labels != loss_cfg.ignore_idx).sum() for labels in shifted_labels_list) + global_grad_tokens = cast(torch.Tensor, rank_grad_tokens) + if dist.is_initialized(): + dist.all_reduce(global_grad_tokens, op=dist.ReduceOp.SUM) + if global_grad_tokens == 0: + logger.warning("Global gradient tokens is 0; using one as the loss denominator") + global_grad_tokens.add_(1) + + for loss_ctx in loss_ctx_list: + loss_kwargs = loss_ctx.loss_kwargs + shifted_labels = loss_kwargs.shifted_labels + assert loss_kwargs.old_logprobs is not None, "old_logprobs can not be None" + + policy_loss_weight = torch.ones_like(shifted_labels, dtype=torch.float32) / global_grad_tokens + policy_loss_weight[shifted_labels == loss_cfg.ignore_idx] = 0.0 + if loss_kwargs.is_weights is not None: + policy_loss_weight = policy_loss_weight * loss_kwargs.is_weights + + if loss_cfg.use_kl_loss: + assert loss_kwargs.ref_logprobs is not None, "ref_logprobs can not be None" + kl_loss_weight = compute_kl_loss_weight( + shifted_labels, + global_grad_tokens, + loss_cfg.kl_loss_coef, + loss_cfg.ignore_idx, + ) + else: + kl_loss_weight = None + + distillation_loss_weight = None + if not loss_cfg.use_policy_gradient: + distillation_loss_weight = ( + torch.ones_like(shifted_labels, dtype=torch.float32) + / global_grad_tokens + * loss_cfg.distillation_loss_weight + ) + distillation_loss_weight[shifted_labels == loss_cfg.ignore_idx] = 0.0 + + loss_kwargs.policy_loss_weight = policy_loss_weight + loss_kwargs.kl_loss_weight = kl_loss_weight + loss_kwargs.distillation_loss_weight = distillation_loss_weight + loss_kwargs.global_grad_tokens = global_grad_tokens + return loss_ctx_list + + def _compute_sampled_distillation_loss( + self, + student_logprobs: torch.Tensor, + teacher_logprobs: torch.Tensor, + ) -> torch.Tensor: + if self.loss_cfg.loss_mode == "k1": + loss = student_logprobs - teacher_logprobs + else: + log_ratio = (teacher_logprobs - student_logprobs).clamp(min=-20.0, max=20.0) + loss = torch.exp(log_ratio) - log_ratio - 1.0 + if self.loss_cfg.loss_max_clamp is not None: + loss = loss.clamp( + min=-self.loss_cfg.loss_max_clamp, + max=self.loss_cfg.loss_max_clamp, + ) + return loss + + def _compute_topk_distillation_loss( + self, + logits: torch.Tensor, + loss_kwargs: DistillationLossKwargs, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + target_token_ids = cast(torch.Tensor, loss_kwargs.target_token_ids) + teacher_logprobs = cast(torch.Tensor, loss_kwargs.teacher_logprobs) + student_topk_logprobs = torch.gather(logits, dim=-1, index=target_token_ids) + student_topk_logprobs = student_topk_logprobs - torch.logsumexp(logits, dim=-1, keepdim=True) + + loss_mode = cast(TopKDistillationMode, self.loss_cfg.loss_mode) + + loss, student_mass, teacher_mass = compute_topk_distillation_kl( + student_topk_logprobs, + teacher_logprobs, + loss_mode, + self.loss_cfg.log_prob_min_clamp, + self.loss_cfg.loss_max_clamp, + ) + student_topk_ids = torch.topk(logits.detach(), k=target_token_ids.size(-1), dim=-1).indices + overlap = (student_topk_ids.unsqueeze(-1) == target_token_ids.detach().unsqueeze(-2)).any(dim=-1) + return loss, student_mass, teacher_mass, overlap.float().mean(dim=-1) + + def loss_fn( + self, + hidden_states: torch.Tensor, + head_weight: torch.Tensor, + head_bias: torch.Tensor | None, + loss_kwargs: DistillationLossKwargs, + ) -> tuple[torch.Tensor, tuple[torch.Tensor | None, dict[str, Any]]]: + logits = F.linear(hidden_states, head_weight, head_bias).float() + shifted_labels = loss_kwargs.shifted_labels + old_logprobs = cast(torch.Tensor, loss_kwargs.old_logprobs) + policy_loss_weight = cast(torch.Tensor, loss_kwargs.policy_loss_weight) + current_logprobs = gather_logprobs(logits, shifted_labels) + + topk_metrics: dict[str, torch.Tensor] = {} + if self.loss_cfg.uses_sampled_token_targets: + teacher_logprobs = cast(torch.Tensor, loss_kwargs.teacher_logprobs) + distillation_student_logprobs = old_logprobs if self.loss_cfg.use_policy_gradient else current_logprobs + per_token_distillation_loss = self._compute_sampled_distillation_loss( + distillation_student_logprobs, + teacher_logprobs, + ) + else: + ( + per_token_distillation_loss, + student_selected_mass, + teacher_selected_mass, + overlap_fraction, + ) = self._compute_topk_distillation_loss(logits, loss_kwargs) + topk_metrics = { + "reduced_topk_opd_student_selected_mass_sum": student_selected_mass.detach(), + "reduced_topk_opd_teacher_selected_mass_sum": teacher_selected_mass.detach(), + "reduced_topk_opd_overlap_fraction_sum": overlap_fraction.detach(), + } + + if self.loss_cfg.use_policy_gradient: + combined_advantages = ( + self.loss_cfg.task_adv_weight * loss_kwargs.advantages + - self.loss_cfg.distillation_loss_weight * per_token_distillation_loss.detach() + ) + loss = self.policy_loss_fn( + current_logprobs, + old_logprobs, + combined_advantages, + policy_loss_weight, + self.loss_cfg.policy_loss_cfg, + ) + else: + distillation_loss_weight = cast(torch.Tensor, loss_kwargs.distillation_loss_weight) + task_advantages = self.loss_cfg.task_adv_weight * loss_kwargs.advantages + task_loss = self.policy_loss_fn( + current_logprobs, + old_logprobs, + task_advantages, + policy_loss_weight, + self.loss_cfg.policy_loss_cfg, + ) + distillation_loss = (per_token_distillation_loss * distillation_loss_weight).sum() + loss = task_loss + distillation_loss + if self.loss_cfg.uses_topk_targets: + topk_metrics["reduced_topk_opd_loss_sum"] = distillation_loss.detach() + + valid_mask = shifted_labels != self.loss_cfg.ignore_idx + valid_float = valid_mask.float() + log_ratio = current_logprobs.detach() - old_logprobs.detach() + log_ratio_safe = torch.clamp(log_ratio, min=-20.0, max=20.0) + ratio = torch.exp(log_ratio_safe) + ratio_max = ratio.masked_fill(~valid_mask, 0.0).max() + ratio_min = ratio.masked_fill(~valid_mask, float("inf")).min() + extra_info = { + "max_ratio": ratio_max, + "reduced_train_policy_ratio_abs_dev_sum": ((ratio - 1.0).abs() * valid_float).sum(), + "reduced_train_policy_kl1_sum": (-log_ratio * valid_float).sum(), + "reduced_train_policy_kl3_sum": ((ratio - 1.0 - log_ratio_safe) * valid_float).sum(), + "reduced_train_policy_valid_count": valid_float.sum(), + "reduced_train_policy_ratio_max": ratio_max, + "reduced_train_policy_ratio_min": ratio_min, + "reduced_distillation_kl_sum": (per_token_distillation_loss.detach() * valid_float).sum(), + "reduced_distillation_abs_loss_sum": (per_token_distillation_loss.detach().abs() * valid_float).sum(), + "reduced_distillation_valid_count": valid_float.sum(), + **{ + key: (value * valid_float).sum() if key != "reduced_topk_opd_loss_sum" else value + for key, value in topk_metrics.items() + }, + } + if self.loss_cfg.loss_mode == "k1": + extra_info["reduced_opd_reverse_kl_sum"] = (per_token_distillation_loss.detach() * valid_float).sum() + extra_info["reduced_opd_abs_logprob_loss_sum"] = ( + per_token_distillation_loss.detach().abs() * valid_float + ).sum() + if self.loss_cfg.uses_topk_targets: + extra_info["reduced_topk_opd_kl_sum"] = (per_token_distillation_loss.detach() * valid_float).sum() + extra_info["reduced_topk_opd_valid_count"] = valid_float.sum() + + cliprange_low = self.loss_cfg.policy_loss_cfg.get("cliprange_low") + cliprange_high = self.loss_cfg.policy_loss_cfg.get("cliprange_high") + if cliprange_low is not None and cliprange_high is not None: + extra_info["reduced_train_policy_clip_low_count"] = ( + ((ratio < 1 - cliprange_low) & valid_mask).float().sum() + ) + extra_info["reduced_train_policy_clip_high_count"] = ( + ((ratio > 1 + cliprange_high) & valid_mask).float().sum() + ) + + if self.loss_cfg.use_kl_loss: + ref_logprobs = loss_kwargs.ref_logprobs + kl_loss_weight = loss_kwargs.kl_loss_weight + assert ref_logprobs is not None and kl_loss_weight is not None + loss = loss + kl_penalty( + current_logprobs, + ref_logprobs, + kl_loss_weight, + self.loss_cfg.kl_loss_type, + ) + + return loss, (logits, extra_info) + + +def finalize_distillation_metrics( + extra_info_dict: dict[str, Any], + device: str | torch.device, +) -> dict[str, Any]: + def reduce_values(keys: tuple[str, ...]) -> dict[str, float]: + values = torch.tensor( + [extra_info_dict.pop(key, 0.0) for key in keys], + dtype=torch.float32, + device=device, + ) + if dist.is_initialized(): + dist.all_reduce(values, op=dist.ReduceOp.SUM) + return dict(zip(keys, values.tolist())) + + if "reduced_distillation_valid_count" in extra_info_dict: + distillation_keys = [ + "reduced_distillation_kl_sum", + "reduced_distillation_abs_loss_sum", + "reduced_distillation_valid_count", + ] + has_opd_metrics = "reduced_opd_reverse_kl_sum" in extra_info_dict + if has_opd_metrics: + distillation_keys.extend( + [ + "reduced_opd_reverse_kl_sum", + "reduced_opd_abs_logprob_loss_sum", + ] + ) + values = reduce_values(tuple(distillation_keys)) + valid_count = values["reduced_distillation_valid_count"] + extra_info_dict["reduced_distillation_kl"] = ( + values["reduced_distillation_kl_sum"] / valid_count if valid_count > 0 else 0.0 + ) + extra_info_dict["reduced_distillation_abs_loss"] = ( + values["reduced_distillation_abs_loss_sum"] / valid_count if valid_count > 0 else 0.0 + ) + if has_opd_metrics: + extra_info_dict["opd_reverse_kl"] = ( + values["reduced_opd_reverse_kl_sum"] / valid_count if valid_count > 0 else 0.0 + ) + extra_info_dict["opd_abs_logprob_loss"] = ( + values["reduced_opd_abs_logprob_loss_sum"] / valid_count if valid_count > 0 else 0.0 + ) + + if "reduced_topk_opd_valid_count" in extra_info_dict: + topk_keys = ( + "reduced_topk_opd_kl_sum", + "reduced_topk_opd_loss_sum", + "reduced_topk_opd_student_selected_mass_sum", + "reduced_topk_opd_teacher_selected_mass_sum", + "reduced_topk_opd_overlap_fraction_sum", + "reduced_topk_opd_valid_count", + ) + values = reduce_values(topk_keys) + valid_count = values["reduced_topk_opd_valid_count"] + for output_key, sum_key in { + "reduced_topk_opd_kl": "reduced_topk_opd_kl_sum", + "reduced_topk_opd_student_selected_mass": "reduced_topk_opd_student_selected_mass_sum", + "reduced_topk_opd_teacher_selected_mass": "reduced_topk_opd_teacher_selected_mass_sum", + "reduced_topk_opd_overlap_fraction": "reduced_topk_opd_overlap_fraction_sum", + }.items(): + extra_info_dict[output_key] = values[sum_key] / valid_count if valid_count > 0 else 0.0 + extra_info_dict["reduced_topk_opd_loss"] = values["reduced_topk_opd_loss_sum"] + return extra_info_dict diff --git a/xtuner/v1/rl/model_utils.py b/xtuner/v1/rl/model_utils.py new file mode 100644 index 0000000000..404569d956 --- /dev/null +++ b/xtuner/v1/rl/model_utils.py @@ -0,0 +1,54 @@ +from pathlib import Path +from typing import cast + +import torch + +from xtuner.v1.config.fsdp import FSDPConfig +from xtuner.v1.float8.float8_handler import Float8Handler +from xtuner.v1.model.base import BaseModel as XtunerBaseModel +from xtuner.v1.model.base import TransformerConfig +from xtuner.v1.model.compose.base import BaseComposeConfig, BaseComposeModel +from xtuner.v1.utils import get_torch_device_module + + +DEVICE_MODULE = get_torch_device_module() + +FrozenModel = BaseComposeModel | XtunerBaseModel + + +def build_frozen_model( + model_cfg: TransformerConfig | BaseComposeConfig, + load_from: str | Path, + fsdp_cfg: FSDPConfig | None = None, +) -> FrozenModel: + """Build a frozen FSDP model and leave it resident on CPU.""" + with torch.device("meta"): + model = model_cfg.build() + + if isinstance(model_cfg, BaseComposeConfig): + assert model_cfg.text_config.float8_cfg is None, "BaseComposeConfig does not support float8" + if fsdp_cfg is None: + fsdp_cfg = FSDPConfig(recompute_ratio=0, cpu_offload=False, requires_grad=False) + model = model.fully_shard(fsdp_cfg) + model.from_hf(hf_path=load_from) + model.eval() # type: ignore + else: + model_cfg = cast(TransformerConfig, model_cfg) + if model_cfg.float8_cfg is not None and model_cfg.float8_cfg.enable_float8: + float8_handler = Float8Handler( + scaling_granularity_gemm=model_cfg.float8_cfg.scaling_granularity_gemm, + scaling_granularity_grouped_gemm=model_cfg.float8_cfg.scaling_granularity_grouped_gemm, + ) + else: + float8_handler = None + if fsdp_cfg is None: + fsdp_cfg = FSDPConfig(recompute_ratio=0, cpu_offload=False, requires_grad=False) + model = model.fully_shard(fsdp_cfg) # type: ignore + model.from_hf(hf_path=load_from) + model.eval() # type: ignore + if float8_handler is not None: + float8_handler.precompute_float8_dynamic_scale_for_fsdp(model) # type: ignore + + model.to_device("cpu") # type: ignore + DEVICE_MODULE.empty_cache() + return model diff --git a/xtuner/v1/rl/trainer/controller.py b/xtuner/v1/rl/trainer/controller.py index 3e005b1d22..8292e02454 100644 --- a/xtuner/v1/rl/trainer/controller.py +++ b/xtuner/v1/rl/trainer/controller.py @@ -4,6 +4,7 @@ import ray import torch +from typing_extensions import NotRequired from xtuner.v1.data_proto.sequence_context import SequenceContext from xtuner.v1.model.compose.base import BaseComposeConfig @@ -22,6 +23,9 @@ class ColateItem(TypedDict): shifted_labels: torch.Tensor advantage: float rollout_logprobs: torch.Tensor | None + teacher_logprobs: NotRequired[torch.Tensor | None] + target_token_ids: NotRequired[torch.Tensor | None] + teacher_indices: NotRequired[torch.Tensor] def _summarize_process_group_results(results: list[dict[str, Any]]) -> str: @@ -99,9 +103,12 @@ def _packing(self, data_batches, pack_max_length, language_cfg): ) packed_data_batches = [] - is_qwen3_vl = False - if len(data_batches[0]["seq_ctx"].position_ids.shape) == 3: - is_qwen3_vl = True + is_qwen3_vl = any(data["seq_ctx"].position_ids.ndim == 3 for data in data_batches) + if is_qwen3_vl: + for data in data_batches: + position_ids = data["seq_ctx"].position_ids + if position_ids.ndim == 2: + data["seq_ctx"].position_ids = position_ids.unsqueeze(0).expand(3, -1, -1) has_rollout_routed_experts = False if data_batches[0]["seq_ctx"].rollout_routed_experts is not None: @@ -121,6 +128,18 @@ 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] + + target_token_ids_list = None + if "target_token_ids" in data_batches[0] and data_batches[0]["target_token_ids"] is not None: + target_token_ids_list = [data_batches[i]["target_token_ids"] for i in indices] + + teacher_indices_list = None + if "teacher_indices" in data_batches[0] and data_batches[0]["teacher_indices"] is not None: + teacher_indices_list = [data_batches[i]["teacher_indices"] for i in indices] + if pad_len > 0: # Reduce the attn calculation time by using multiple short sequence packs pad_tokens = tuple( @@ -162,6 +181,30 @@ 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: + teacher_logprobs_shape = data_batches[0]["teacher_logprobs"].shape[2:] + pad_teacher_logprobs = torch.zeros( + (1, pad_len, *teacher_logprobs_shape), + dtype=data_batches[0]["teacher_logprobs"].dtype, + device=data_batches[0]["shifted_labels"].device, + ) + teacher_logprobs_list.append(pad_teacher_logprobs) + if target_token_ids_list is not None: + target_token_ids_shape = data_batches[0]["target_token_ids"].shape[2:] + pad_target_token_ids = torch.zeros( + (1, pad_len, *target_token_ids_shape), + dtype=data_batches[0]["target_token_ids"].dtype, + device=data_batches[0]["shifted_labels"].device, + ) + target_token_ids_list.append(pad_target_token_ids) + if teacher_indices_list is not None: + pad_teacher_indices = torch.full( + (1, pad_len), + -1, + dtype=data_batches[0]["teacher_indices"].dtype, + device=data_batches[0]["teacher_indices"].device, + ) + teacher_indices_list.append(pad_teacher_indices) seq_ctx = SequenceContext.cat(seq_ctx_list) shifted_labels = torch.cat(label_list, dim=1) # (1, max_len) @@ -172,12 +215,27 @@ 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) + + target_token_ids = None + if target_token_ids_list is not None: + target_token_ids = torch.cat(target_token_ids_list, dim=1) + + teacher_indices = None + if teacher_indices_list is not None: + teacher_indices = torch.cat(teacher_indices_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, + "target_token_ids": target_token_ids, + "teacher_indices": teacher_indices, } ) return packed_data_batches @@ -260,11 +318,38 @@ 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: + teacher_logprobs_shape = packed_data_batches[0]["teacher_logprobs"].shape[2:] + pad_teacher_logprobs = torch.zeros( + (1, pack_max_length, *teacher_logprobs_shape), + dtype=packed_data_batches[0]["teacher_logprobs"].dtype, + device="cpu", + ) + pad_target_token_ids = None + if "target_token_ids" in packed_data_batches[0] and packed_data_batches[0]["target_token_ids"] is not None: + target_token_ids_shape = packed_data_batches[0]["target_token_ids"].shape[2:] + pad_target_token_ids = torch.zeros( + (1, pack_max_length, *target_token_ids_shape), + dtype=packed_data_batches[0]["target_token_ids"].dtype, + device="cpu", + ) + pad_teacher_indices = None + if "teacher_indices" in packed_data_batches[0] and packed_data_batches[0]["teacher_indices"] is not None: + pad_teacher_indices = torch.full( + (1, pack_max_length), + -1, + dtype=packed_data_batches[0]["teacher_indices"].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, + "target_token_ids": pad_target_token_ids, + "teacher_indices": pad_teacher_indices, } 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 80929e9fff..d8de74dfed 100644 --- a/xtuner/v1/rl/trainer/worker.py +++ b/xtuner/v1/rl/trainer/worker.py @@ -35,16 +35,24 @@ from xtuner.v1.datasets.config import DataloaderConfig from xtuner.v1.datasets.dataloader import Dataloader from xtuner.v1.engine.train_engine import TrainEngine, TrainStepInfo -from xtuner.v1.float8.float8_handler import Float8Handler from xtuner.v1.loss import BaseLossContext, CELossConfig, LogProbConfig from xtuner.v1.loss.ce_loss import CELossContext, LMHeadLossContext from xtuner.v1.loss.mtp_loss import MTPLossConfig, MTPLossContext -from xtuner.v1.model.base import BaseModel as XtunerBaseModel +from xtuner.v1.loss.utils import sp_split from xtuner.v1.model.base import ModelItem, TransformerConfig -from xtuner.v1.model.compose.base import BaseComposeConfig, BaseComposeModel +from xtuner.v1.model.compose.base import BaseComposeConfig 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.distillation import DistillationConfig, TrainTeacherManager, TrainTeacherTimings +from xtuner.v1.rl.loss import ( + BaseRLLossConfig, + BaseRLLossContext, + DistillationLossConfig, + finalize_distillation_metrics, + finalize_train_policy_metrics, + kl_penalty, +) +from xtuner.v1.rl.model_utils import build_frozen_model from xtuner.v1.rl.utils import SingleAcceleratorWorker from xtuner.v1.rl.weight_update import WeightUpdater from xtuner.v1.train.trainer import LoadCheckpointConfig @@ -154,6 +162,7 @@ class WorkerConfig(BaseModel): profile_memory: bool = False free_rollout_routed_experts_in_worker: bool = True # 默认不需要用户配置 offload_rollout_routed_experts: bool = False + distillation_config: DistillationConfig | None = None # sft config sft_dataloader_cfg: DataloaderConfig | None = None @@ -187,6 +196,9 @@ class WorkerInputItem(TypedDict): shifted_labels: torch.LongTensor advantages: torch.Tensor rollout_logprobs: torch.Tensor | None + teacher_logprobs: NotRequired[torch.Tensor | None] + target_token_ids: NotRequired[torch.Tensor | None] + teacher_indices: NotRequired[torch.Tensor] class WorkerTrainLogItem(TypedDict, total=False): @@ -200,6 +212,9 @@ class WorkerTrainLogItem(TypedDict, total=False): class WorkerLogItem(TypedDict): train_entropy: float + teacher_compute_time: NotRequired[float] + teacher_onload_time: NotRequired[float] + teacher_offload_time: NotRequired[float] rollout_entropy: NotRequired[float] mismatch_metrics: NotRequired[dict[str, float]] rollout_is_metrics: NotRequired[dict[str, float]] @@ -252,10 +267,18 @@ def __init__( self._has_ref = True if worker_cfg.ref_load_from is None: worker_cfg.ref_load_from = worker_cfg.load_from - self._ref_model = self._build_ref_model( + self._ref_model = build_frozen_model( worker_cfg.model_cfg, worker_cfg.ref_load_from, worker_cfg.ref_model_fsdp_cfg ) + self._train_teacher_manager: TrainTeacherManager | None = None + distillation_config = worker_cfg.distillation_config + if distillation_config is not None and distillation_config.train_teachers: + self._train_teacher_manager = TrainTeacherManager( + distillation_config, + chunk_size=worker_cfg.loss_cfg.chunk_size, + ) + self._optimizer_steps = worker_cfg.optimizer_steps profile_step = worker_cfg.profile_step if isinstance(profile_step, int): @@ -387,46 +410,6 @@ def _build_engine(self, worker_cfg: WorkerConfig) -> TrainEngine: self.logger.info(f"The `compile_cfg` of model is {json.dumps(engine.model.compile_cfg, indent=4)}") return engine - def _build_ref_model( - self, - ref_model_cfg: TransformerConfig | BaseComposeConfig, - load_from: str | Path, - ref_model_fsdp_cfg: FSDPConfig | None = None, - ): - # TODO: 需要重构,使得能更优雅的兼容 mllm - model: BaseComposeModel | XtunerBaseModel - with torch.device("meta"): - model = ref_model_cfg.build() - - if isinstance(ref_model_cfg, BaseComposeConfig): - assert ref_model_cfg.text_config.float8_cfg is None, "BaseComposeConfig does not support float8" - if ref_model_fsdp_cfg is None: - ref_model_fsdp_cfg = FSDPConfig(recompute_ratio=0, cpu_offload=False, requires_grad=False) - model = model.fully_shard(ref_model_fsdp_cfg) - model.from_hf(hf_path=load_from) - model.eval() # type: ignore - else: - ref_model_cfg = cast(TransformerConfig, ref_model_cfg) - if ref_model_cfg.float8_cfg is not None and ref_model_cfg.float8_cfg.enable_float8: - float8_handler = Float8Handler( - scaling_granularity_gemm=ref_model_cfg.float8_cfg.scaling_granularity_gemm, - scaling_granularity_grouped_gemm=ref_model_cfg.float8_cfg.scaling_granularity_grouped_gemm, - ) - else: - float8_handler = None - if ref_model_fsdp_cfg is None: - ref_model_fsdp_cfg = FSDPConfig(recompute_ratio=0, cpu_offload=False, requires_grad=False) - model = model.fully_shard(ref_model_fsdp_cfg) # type: ignore - - model.from_hf(hf_path=load_from) - model.eval() # type: ignore - if float8_handler is not None: - # As the ref model is not updated, we only compute params' scales once - float8_handler.precompute_float8_dynamic_scale_for_fsdp(model) # type: ignore - model.to_device("cpu") # type: ignore - DEVICE_MODULE.empty_cache() - return model - def _init_data_mesh( self, sp_size: int, @@ -479,6 +462,44 @@ def compute_ref_logprobs( self._ref_model.to_device("cpu") return ref_logprobs_list + def _compute_train_teacher_outputs( + self, + seq_ctx_list: list[SequenceContext], + teacher_indices_list: list[torch.Tensor], + loss_ctx_list: list[BaseRLLossContext], + ) -> TrainTeacherTimings: + self._engine.put_model_to_device("cpu") + self._engine.put_optimizer_to_device("cpu") + if hasattr(DEVICE_MODULE, "empty_cache"): + DEVICE_MODULE.empty_cache() + try: + assert self._train_teacher_manager is not None + outputs = self._train_teacher_manager.compute_logprobs( + seq_ctx_list=seq_ctx_list, + shifted_labels_list=[loss_ctx.loss_kwargs.shifted_labels for loss_ctx in loss_ctx_list], + teacher_indices_list=teacher_indices_list, + ) + if outputs.target_token_ids is None: + for loss_ctx, teacher_logprobs in zip(loss_ctx_list, outputs.teacher_logprobs): + loss_ctx.loss_kwargs.teacher_logprobs = teacher_logprobs + else: + for loss_ctx, token_ids, teacher_logprobs in zip( + loss_ctx_list, + outputs.target_token_ids, + outputs.teacher_logprobs, + ): + loss_ctx.loss_kwargs.target_token_ids = token_ids + loss_ctx.loss_kwargs.teacher_logprobs = teacher_logprobs + return outputs.timings + finally: + if self._train_teacher_manager is not None: + self._train_teacher_manager.offload_all_to_cpu() + self._onload_actor_and_optimizer() + + def _onload_actor_and_optimizer(self) -> None: + self._engine.put_model_to_device(DEVICE) + self._engine.put_optimizer_to_device(DEVICE) + def _add_rollout_routed_experts( self, seq_ctx: SequenceContext, rollout_routed_experts: torch.Tensor | list[torch.Tensor | ray.ObjectRef] ): @@ -588,6 +609,8 @@ def _maybe_profiling(self, global_train_step: int, phase: str): def fit(self, data_batches: list[WorkerInputItem], rollout_idx: int) -> WorkerLogItem: # NOTE: sglang会清除logger handle, 重新创建 self.logger = get_logger(log_dir=self.log_dir, tag="TrainingWorker") + if self._train_teacher_manager is None: + self._onload_actor_and_optimizer() loss_cfg: BaseRLLossConfig = self.config.loss_cfg num_batches = len(data_batches) iters_per_step = math.ceil(num_batches / self._optimizer_steps) @@ -600,11 +623,13 @@ def fit(self, data_batches: list[WorkerInputItem], rollout_idx: int) -> WorkerLo # Init loss_ctx: shifted_labels, advantages, rollout_logprobs seq_ctx_list: list[SequenceContext] = [] loss_ctx_list: list[BaseRLLossContext] = [] + teacher_indices_list: list[torch.Tensor] = [] mtp_loss_ctx_list: list[list[MTPLossContext]] = [] prepare_inputs_begin = time.perf_counter() for data in data_batches: # update seq_ctx seq_ctx = data["seq_ctx"] + teacher_indices = data.get("teacher_indices") pixel_values = seq_ctx.pixel_values if pixel_values is not None: if not isinstance(pixel_values, np.ndarray): @@ -632,11 +657,23 @@ def fit(self, data_batches: list[WorkerInputItem], rollout_idx: int) -> WorkerLo advantages = data["advantages"].to(DEVICE) rollout_logprobs = data.get("rollout_logprobs", None) rollout_logprobs = rollout_logprobs.to(DEVICE) if rollout_logprobs is not None else None + if teacher_indices is not None: + teacher_indices = teacher_indices.to(DEVICE) + if self.sp_mesh.size() > 1: + teacher_indices = sp_split( + teacher_indices, + sp_mesh=self.sp_mesh, + split_dim=1, + padding_value=-1, + ) + teacher_indices_list.append(teacher_indices) loss_ctx = loss_cfg.build( data={ "shifted_labels": shifted_labels, "advantages": advantages, "rollout_logprobs": rollout_logprobs, + "teacher_logprobs": data.get("teacher_logprobs", None), + "target_token_ids": data.get("target_token_ids", None), }, sp_mesh=self.sp_mesh, ) @@ -674,12 +711,25 @@ def fit(self, data_batches: list[WorkerInputItem], rollout_idx: int) -> WorkerLo shifted_labels_list = [loss_ctx.loss_kwargs.shifted_labels for loss_ctx in loss_ctx_list] rollout_logprobs_list = [loss_ctx.loss_kwargs.rollout_logprobs for loss_ctx in loss_ctx_list] + worker_log_item: WorkerLogItem = {"train_entropy": 0.0, "train_metrics": [], "sft_train_metrics": {}} + if self._train_teacher_manager is not None: + # Training-side Teachers share the training workers' devices. Keep + # Actor and optimizer on CPU while each frozen Teacher produces its + # targets, then restore Actor state for old-logprob and train forward. + teacher_timings = self._compute_train_teacher_outputs( + seq_ctx_list, + teacher_indices_list, + loss_ctx_list, + ) + worker_log_item["teacher_compute_time"] = teacher_timings.compute + worker_log_item["teacher_onload_time"] = teacher_timings.onload + worker_log_item["teacher_offload_time"] = teacher_timings.offload + # compute old logprobs old_logprobs_list = self.compute_actor_logprobs(seq_ctx_list, shifted_labels_list) for old_logprobs, loss_ctx in zip(old_logprobs_list, loss_ctx_list): loss_ctx.loss_kwargs.old_logprobs = old_logprobs - worker_log_item: WorkerLogItem = {"train_entropy": 0.0, "train_metrics": [], "sft_train_metrics": {}} logger_msg = f"Rollout {rollout_idx}: " # compute entropy @@ -860,6 +910,8 @@ def fit(self, data_batches: list[WorkerInputItem], rollout_idx: int) -> WorkerLo if isinstance(v, (torch.Tensor, int, float)) } extra_info_dict = finalize_train_policy_metrics(extra_info_dict, DEVICE) + if isinstance(loss_cfg, DistillationLossConfig): + extra_info_dict = finalize_distillation_metrics(extra_info_dict, DEVICE) train_step_info.pop("total_loss") # type: ignore[misc] max_memory = DEVICE_MODULE.max_memory_allocated() / (1024**3) # type: ignore[attr-defined] reserved_memory = DEVICE_MODULE.max_memory_reserved() / (1024**3) # type: ignore[attr-defined] diff --git a/xtuner/v1/train/rl_trainer.py b/xtuner/v1/train/rl_trainer.py index e302082f2b..239ce10dbc 100644 --- a/xtuner/v1/train/rl_trainer.py +++ b/xtuner/v1/train/rl_trainer.py @@ -4,6 +4,7 @@ import random import re import time +from collections.abc import Sequence from dataclasses import asdict, dataclass from pathlib import Path from shutil import rmtree @@ -32,7 +33,9 @@ ProduceBatchStatus, ) from xtuner.v1.rl.agent_loop_manager.produce_utils import default_should_continue_fn +from xtuner.v1.rl.distillation import DistillationConfig from xtuner.v1.rl.evaluator import EvaluatorConfig +from xtuner.v1.rl.loss import DistillationLossConfig from xtuner.v1.rl.replay_buffer import ( AsyncReplayBufferConfig, SyncReplayBufferConfig, @@ -65,6 +68,13 @@ DEVICE = get_device() DEVICE_MODULE = get_torch_device_module() +_DISTILLATION_METRIC_PREFIXES = ( + "reduced_distillation_", + "reduced_topk_opd_", + "opd_", + "topk_opd_", +) + def _to_cpu_tensor(value: np.ndarray | None, *, dtype: torch.dtype | None = None) -> torch.Tensor | None: if value is None: @@ -73,6 +83,66 @@ def _to_cpu_tensor(value: np.ndarray | None, *, dtype: torch.dtype | None = None return torch.as_tensor(value, dtype=dtype, device="cpu") +def _align_rollout_teacher_targets( + state: RolloutState, + loss_config: DistillationLossConfig, + *, + shifted_labels: Sequence[int], + target_start: int, +) -> tuple[torch.Tensor, torch.Tensor | None]: + """Left-pad rollout Teacher targets to align with one training sequence.""" + + sequence_length = len(shifted_labels) + context = f"rollout_id={state.rollout_id}, group_id={state.group_id}" + if target_start < 0 or target_start > sequence_length: + raise ValueError( + f"Teacher target_start must be within the shifted sequence: {target_start} vs {sequence_length}; {context}" + ) + if any(label != loss_config.ignore_idx for label in shifted_labels[:target_start]): + raise ValueError(f"Teacher target prefix must contain only ignored labels; {context}") + + expected_target_rows = sequence_length - target_start + raw_teacher_logprobs = state.teacher_logprobs + if raw_teacher_logprobs is None or len(raw_teacher_logprobs) != expected_target_rows: + actual_rows = None if raw_teacher_logprobs is None else len(raw_teacher_logprobs) + raise ValueError( + "Teacher logprobs must align with the target suffix: " + f"expected {expected_target_rows} rows, got {actual_rows}; {context}" + ) + + if loss_config.uses_sampled_token_targets: + sampled_logprobs = cast(list[float], raw_teacher_logprobs) + teacher_logprobs = torch.tensor( + [0.0] * target_start + sampled_logprobs, + dtype=torch.float32, + ).unsqueeze(0) + return teacher_logprobs, None + + top_k = cast(int, loss_config.top_k) + raw_teacher_tokens = state.teacher_tokens + if raw_teacher_tokens is None or len(raw_teacher_tokens) != expected_target_rows: + actual_rows = None if raw_teacher_tokens is None else len(raw_teacher_tokens) + raise ValueError( + "Teacher token ids must align with the target suffix: " + f"expected {expected_target_rows} rows, got {actual_rows}; {context}" + ) + + topk_tokens = cast(list[list[int]], raw_teacher_tokens) + topk_logprobs = cast(list[list[float]], raw_teacher_logprobs) + + prompt_target_tokens = [[0] * top_k for _ in range(target_start)] + prompt_teacher_logprobs = [[0.0] * top_k for _ in range(target_start)] + teacher_logprobs = torch.tensor( + prompt_teacher_logprobs + topk_logprobs, + dtype=torch.float32, + ).unsqueeze(0) + target_token_ids = torch.tensor( + prompt_target_tokens + topk_tokens, + dtype=torch.int64, + ).unsqueeze(0) + return teacher_logprobs, target_token_ids + + def _agent_loop_manager_requires_rollout_proxy( cfg: AgentLoopManagerConfig | DisaggAgentLoopManagerConfig | None, ) -> bool: @@ -347,6 +417,7 @@ class BaseRLTrainerConfig(BaseModel): total_epochs: int | None = None train_batch_size: int advantage_estimator_config: BaseAdvantageConfig = Field(default_factory=GRPOAdvantageConfig) + distillation_config: DistillationConfig | None = None sync_weights_interval: int = 1 enable_evaluate: bool = True @@ -384,6 +455,10 @@ def _validate_sync_intervals(self): raise ValueError(f"total_train_steps must be positive, got {self.total_train_steps}.") if self.total_epochs is not None and self.total_epochs <= 0: raise ValueError(f"total_epochs must be positive, got {self.total_epochs}.") + if self.distillation_config is not None: + if self.train_worker_cfg.loss_cfg != self.distillation_config.loss_config: + raise ValueError("train_worker_cfg.loss_cfg must be distillation_config.loss_config") + self.distillation_config.validate_student_model(self.train_worker_cfg.model_cfg) _validate_sync_intervals( sync_weights_interval=self.sync_weights_interval, checkpoint_interval=self.checkpoint_interval, @@ -587,6 +662,25 @@ class BaseRLTrainer: _debug_train_files: dict[int, Path] def _init_common(self, cfg: BaseRLTrainerConfig, *, meta_path: str, logger_tag: str) -> None: + if cfg.distillation_config is not None and cfg.distillation_config.rollout_teachers: + endpoint_map = json.loads(os.environ.get("XTUNER_OPD_TEACHER_ENDPOINTS_JSON", "{}")) + cfg.distillation_config = cfg.distillation_config.resolve_teacher_endpoints(endpoint_map) + + self._distillation_config = cfg.distillation_config + self._distillation_loss_cfg = ( + self._distillation_config.loss_config if self._distillation_config is not None else None + ) + self._train_teacher_config = ( + self._distillation_config + if self._distillation_config is not None and self._distillation_config.train_teachers + else None + ) + self._rollout_teacher_config = ( + self._distillation_config + if self._distillation_config is not None and self._distillation_config.rollout_teachers + else None + ) + check_fa3() self._init_work_dir_and_meta(cfg, meta_path) self._init_load_source(cfg) @@ -689,6 +783,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.distillation_config = cfg.distillation_config self._train_worker_cfg = cfg.train_worker_cfg def _init_rollout_config(self, cfg: BaseRLTrainerConfig, log_dir: Path) -> None: @@ -716,12 +811,18 @@ def _init_runtime_flags(self, cfg: BaseRLTrainerConfig) -> None: def _build_agent_loop_components(self, cfg: BaseRLTrainerConfig, replay_buffer) -> None: self.tokenizer = AutoTokenizer.from_pretrained(cfg.tokenizer_path, trust_remote_code=True) + # Sampled-token objectives require identity Student sampling even when + # Teacher logprobs are produced later by the training workers. + agent_loop_distillation_config = self._rollout_teacher_config + if self._distillation_loss_cfg is not None and self._distillation_loss_cfg.uses_sampled_token_targets: + agent_loop_distillation_config = self._distillation_config agent_loop_manager = cfg.agent_loop_manager_cfg.build( rollout_controller=self.rollout_controller, tokenizer=self.tokenizer, replay_buffer=replay_buffer, logger=self.logger, sync_weights_interval=cfg.sync_weights_interval, + distillation_config=agent_loop_distillation_config, ) self.agent_loop_manager = cast(AgentLoopManager | DisaggAgentLoopManager, agent_loop_manager) @@ -736,6 +837,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, + distillation_config=None, ), ) @@ -941,7 +1043,7 @@ def _train_one_batch( step_timer_dict: dict, *, offload_rollout_before_train: bool = False, - onload_train_before_train: bool = False, + resume_train_before_train: bool = False, raw_rewards_sum: float = 0.0, raw_rewards_count: int = 0, ) -> TrainInfo: @@ -954,21 +1056,18 @@ def _train_one_batch( self._save_trajectories(train_batch, train_trajectory_path) self.logger.info(f"Train step {train_step} train trajectories saved to {train_trajectory_path}") - # 共卡训练前切换资源:检查 rollout -> offload rollout -> onload train。 + # 共卡训练前切换资源:检查 rollout -> offload rollout -> resume train NCCL。 if offload_rollout_before_train: ray.get( self.rollout_controller.check_and_shutdown_inactive_workers.remote(), timeout=RL_TRAINER_RAY_GET_TIMEOUT, ) ray.get(self.rollout_controller.offload.remote(), timeout=RL_TRAINER_RAY_GET_TIMEOUT) - if onload_train_before_train: + if resume_train_before_train: if getattr(self, "_train_nccl_suspended", False): with timer("resume_train_nccl", step_timer_dict): self.train_controller.resume_train_nccl_process_groups() self._train_nccl_suspended = False - with timer("onload", step_timer_dict): - self.train_controller.onload(target="all") - self.logger.info("Training controller loaded") with timer("prepare_data", step_timer_dict): data_batches, data_info = self._prepare_train_data( @@ -1056,24 +1155,32 @@ def _prepare_train_data( raw_rewards_sum: float = 0.0, raw_rewards_count: int = 0, ): - 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] = [] + teacher_rewards: dict[str, list[float]] = {} advantages_list = [] prompt_len_list = [] response_len_list = [] tool_turns_list: list[int] = [] training_tokens = 0 + training_samples = 0 data_batches = [] + teacher_index_by_data_source = ( + self._train_teacher_config.teacher_index_by_data_source if self._train_teacher_config is not None else None + ) for j, group in enumerate(data_groups): 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) + task_adv_weight = ( + self._distillation_loss_cfg.task_adv_weight if self._distillation_loss_cfg 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 @@ -1085,19 +1192,22 @@ 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. + for data in group: + # 有可能有重复,但是没有其他更好办法 + turns = data.extra_fields.get("agent_tool_turns") + if isinstance(turns, int): + tool_turns_list.append(turns) + # Collect rewards independently from task-advantage computation. Pure OPD may omit + # rewards entirely; when rewards are present they remain useful observability signals. 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}" - ) + if data.reward is None or "score" not in data.reward: + assert task_adv_weight == 0, f"Reward is missing or does not contain 'score' key in data: {data}" + continue 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 @@ -1108,19 +1218,34 @@ def _prepare_train_data( 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 self._distillation_config is not None: + for reward, representative in zip(cluster_rewards, cluster_representatives): + data_source = representative.extra_fields.get("origin_data_source") + if not isinstance(data_source, str): + continue + teacher_name = self._distillation_config.data_source_teacher_map.get(data_source) + if teacher_name is not None: + teacher_rewards.setdefault(teacher_name, []).append(reward) + + if task_adv_weight == 0: + sample_advantages = [0.0] * len(group) + else: + # 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. + 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 + ] prompt_repeat_k = len(group) for i in range(prompt_repeat_k): + teacher_index = None + if teacher_index_by_data_source is not None: + data_source = group[i].extra_fields["origin_data_source"] + teacher_index = teacher_index_by_data_source[data_source] if group[i].input_ids is not None: raw_input_ids = cast(list[int], group[i].input_ids) labels = cast(list[int] | None, group[i].labels) @@ -1142,6 +1267,23 @@ def _prepare_train_data( input_ids = raw_input_ids[:-1] shifted_labels = labels[1:] + teacher_logprobs = None + target_token_ids = None + if ( + self._distillation_config is not None + and self._distillation_loss_cfg is not None + and self._distillation_config.rollout_teachers + ): + teacher_response_start = next( + (index for index, label in enumerate(shifted_labels) if label != -100), + len(shifted_labels), + ) + teacher_logprobs, target_token_ids = _align_rollout_teacher_targets( + group[i], + self._distillation_loss_cfg, + shifted_labels=shifted_labels, + target_start=teacher_response_start, + ) prompt_len = sum(label == -100 for label in shifted_labels) response_len = len(shifted_labels) - prompt_len prompt_len_list.append(prompt_len) @@ -1172,6 +1314,15 @@ def _prepare_train_data( "advantage": actual_advantages, "rollout_logprobs": rollout_logprobs, } + if teacher_logprobs is not None: + data_dict["teacher_logprobs"] = teacher_logprobs + if target_token_ids is not None: + data_dict["target_token_ids"] = target_token_ids + if teacher_index is not None: + data_dict["teacher_indices"] = torch.full_like( + shifted_labels_t, + teacher_index, + ) seq_ctx.rollout_routed_experts = group[i].routed_experts data_batches.append(data_dict) @@ -1222,12 +1373,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) @@ -1252,6 +1400,25 @@ def _prepare_train_data( "advantage": actual_advantages, "rollout_logprobs": rollout_logprobs, } + if ( + self._distillation_config is not None + and self._distillation_loss_cfg is not None + and self._distillation_config.rollout_teachers + ): + teacher_logprobs, target_token_ids = _align_rollout_teacher_targets( + group[i], + self._distillation_loss_cfg, + shifted_labels=shifted_labels, + target_start=len(prompt_ids) - 1, + ) + data_dict["teacher_logprobs"] = teacher_logprobs + if target_token_ids is not None: + data_dict["target_token_ids"] = target_token_ids + if teacher_index is not None: + data_dict["teacher_indices"] = torch.full_like( + shifted_labels_t, + teacher_index, + ) seq_ctx.rollout_routed_experts = group[i].routed_experts # n,layer*expert @@ -1259,8 +1426,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() @@ -1268,8 +1435,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(), @@ -1286,6 +1453,9 @@ def _prepare_train_data( "prompt_len/min": prompt_len_t.min().item(), "prompt_len/max": prompt_len_t.max().item(), } + for teacher_name, rewards in teacher_rewards.items(): + if rewards: + info_dict[f"rewards/{teacher_name}/mean"] = torch.tensor(rewards, dtype=torch.float32).mean().item() if tool_turns_list: tool_turns_t = torch.tensor(tool_turns_list, dtype=torch.float32) info_dict["tool_turns/mean"] = tool_turns_t.mean().item() @@ -1408,6 +1578,12 @@ 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 "teacher_compute_time" in rank0_log_item: + all_scalars["time/train_teacher_compute"] = rank0_log_item["teacher_compute_time"] + if "teacher_onload_time" in rank0_log_item: + all_scalars["time/train_teacher_onload"] = rank0_log_item["teacher_onload_time"] + if "teacher_offload_time" in rank0_log_item: + all_scalars["time/train_teacher_offload"] = rank0_log_item["teacher_offload_time"] for worker_idx, log_item in enumerate(train_info["workers_log_item"]): if not self._display_all_workers_log and worker_idx > 0: break @@ -1419,6 +1595,10 @@ def _log_step( for key, value in mini_batch_metrics.items(): avg_value = sum(value) / len(value) all_scalars.update({f"train_metrics/worker_{worker_idx}/step_avg_{key}": avg_value}) + if worker_idx == 0 and key.startswith(_DISTILLATION_METRIC_PREFIXES): + all_scalars[f"distillation/{key}"] = avg_value + if key in ("opd_reverse_kl", "opd_abs_logprob_loss"): + all_scalars[key] = avg_value rank_sft_log = log_item["sft_train_metrics"] for k, v in rank_sft_log.items(): @@ -1447,8 +1627,9 @@ def _save_trajectories(self, data_groups: list[list[RolloutState]], save_path: P continue valid_groups.append(group) 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_len_list.append(len(response_ids)) @@ -1458,20 +1639,20 @@ 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(response_len_list), } json.dump(summary, f, ensure_ascii=False, separators=(",", ":")) f.write("\n") for group in valid_groups: for data in group: - assert data.reward is not None + reward = data.reward.get("score") if data.reward is not None else None response_ids = self._get_trajectory_response_ids(data) response = data.response if response is None and response_ids: @@ -1491,7 +1672,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": len(response_ids), "reward_payload": data.reward, @@ -1708,7 +1889,7 @@ def fit(self): self._exp_tracker.close() close_trace() - def _fit(self): + def _fit(self) -> None: self.logger.info("Start RL training") if self._cur_step >= self._total_train_steps: self.logger.info(f"Train steps {self._total_train_steps} reached, stop training") @@ -1734,7 +1915,7 @@ def _fit(self): model_step = self._get_colocate_rollout_model_step(init_train_step) for train_step in range(init_train_step, self._total_train_steps + 1): self.logger.info(f"Train step {train_step}/{self._total_train_steps} start") - step_timer_dict = {} + step_timer_dict: dict[str, float] = {} with timer("step", step_timer_dict): # 共卡一次调用内完成生产和消费。 self.logger.info( @@ -1761,7 +1942,7 @@ def _fit(self): train_step, step_timer_dict, offload_rollout_before_train=True, - onload_train_before_train=True, + resume_train_before_train=True, raw_rewards_sum=produce_result.raw_rewards_sum, raw_rewards_count=produce_result.raw_rewards_count, ) @@ -1801,7 +1982,7 @@ def _fit_debug_train(self) -> None: train_step, step_timer_dict, offload_rollout_before_train=False, - onload_train_before_train=False, + resume_train_before_train=False, ) eval_log_info: dict[str, float] = {} produce_result = ProduceBatchResult(rollout_states=train_batch)