From fe5e4e95bf77a2fe959e699b40251cdec712c281 Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Wed, 13 Aug 2025 17:17:07 +0800 Subject: [PATCH] Minor extract generate one sample (#168) * more * more * more * more * more * fmt * mv --- slime/rollout/components/__init__.py | 0 slime/rollout/components/sample_generator.py | 76 +++++++++++++++++++ slime/rollout/sglang_rollout.py | 79 +------------------- 3 files changed, 78 insertions(+), 77 deletions(-) create mode 100644 slime/rollout/components/__init__.py create mode 100644 slime/rollout/components/sample_generator.py diff --git a/slime/rollout/components/__init__.py b/slime/rollout/components/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/slime/rollout/components/sample_generator.py b/slime/rollout/components/sample_generator.py new file mode 100644 index 0000000000..57be64f51a --- /dev/null +++ b/slime/rollout/components/sample_generator.py @@ -0,0 +1,76 @@ +from slime.utils.http_utils import post +from slime.utils.types import Sample + + +async def generate_one_sample_vanilla(args, tokenizer, sample: Sample, sampling_params) -> Sample: + url = f"http://{args.sglang_router_ip}:{args.sglang_router_port}/generate" + + assert ( + sample.status == Sample.Status.PENDING or sample.status == Sample.Status.ABORTED + ), f"Sample status is {sample.status}" + + if len(sample.response) > 0: + sampling_params["max_new_tokens"] -= len(sample.tokens) - len( + tokenizer(sample.prompt, add_special_tokens=False)["input_ids"] + ) + + assert ( + sampling_params["max_new_tokens"] >= 0 + ), f"max_new_tokens: {sampling_params['max_new_tokens']} should not be less than 0" + if sampling_params["max_new_tokens"] == 0: + sample.status = Sample.Status.TRUNCATED + return sample + + # Prepare payload - shared structure + payload = { + "sampling_params": sampling_params, + "return_logprob": args.use_token_output, + } + + if args.use_token_output: + # Token-based mode: use tokens directly + if len(sample.response) > 0: + input_token_ids = sample.tokens + else: + # First turn: initialize with prompt tokens + prompt_token_ids = tokenizer(sample.prompt, add_special_tokens=False)["input_ids"] + input_token_ids = prompt_token_ids + # Initialize sample.tokens with prompt for subsequent turns + if not sample.tokens: # Only set if empty + sample.tokens = prompt_token_ids + payload["input_ids"] = input_token_ids + else: + # String-based mode: original implementation + input_text = sample.prompt + sample.response + payload["text"] = input_text + + output = await post(url, payload, use_http2=args.use_http2) + + if args.use_token_output: + # Extract new response tokens + assert ( + "meta_info" in output and "output_token_logprobs" in output["meta_info"] + ), "output_token_logprobs is not in the output" + new_response_tokens = [item[1] for item in output["meta_info"]["output_token_logprobs"]] + + # Update sample with tokens directly - avoiding re-tokenization + sample.tokens = sample.tokens + new_response_tokens + sample.response_length += len(new_response_tokens) + sample.response += tokenizer.decode(new_response_tokens, skip_special_tokens=False) + else: + # String-based processing + sample.response += output["text"] + prompt_tokens_ids = tokenizer(sample.prompt, add_special_tokens=False)["input_ids"] + response_token_ids = tokenizer(sample.response, add_special_tokens=False)["input_ids"] + sample.tokens = prompt_tokens_ids + response_token_ids + sample.response_length = len(response_token_ids) + + match output["meta_info"]["finish_reason"]["type"]: + case "length": + sample.status = Sample.Status.TRUNCATED + case "abort": + sample.status = Sample.Status.ABORTED + case "stop": + sample.status = Sample.Status.COMPLETED + + return sample diff --git a/slime/rollout/sglang_rollout.py b/slime/rollout/sglang_rollout.py index 62bb8ed9c2..f90edf3096 100644 --- a/slime/rollout/sglang_rollout.py +++ b/slime/rollout/sglang_rollout.py @@ -9,6 +9,7 @@ from slime.utils.data import Dataset from slime.utils.http_utils import get, post from slime.utils.misc import SingletonMeta, load_function from slime.utils.types import Sample +from slime.rollout.components.sample_generator import generate_one_sample_vanilla from .rm_hub import async_rm, batched_async_rm @@ -62,82 +63,6 @@ class GenerateState(metaclass=SingletonMeta): self.remaining_batch_size += len(samples) -async def generate(args, sample: Sample, sampling_params) -> Sample: - state = GenerateState(args) - - url = f"http://{args.sglang_router_ip}:{args.sglang_router_port}/generate" - - assert ( - sample.status == Sample.Status.PENDING or sample.status == Sample.Status.ABORTED - ), f"Sample status is {sample.status}" - - if len(sample.response) > 0: - sampling_params["max_new_tokens"] -= len(sample.tokens) - len( - state.tokenizer(sample.prompt, add_special_tokens=False)["input_ids"] - ) - - assert ( - sampling_params["max_new_tokens"] >= 0 - ), f"max_new_tokens: {sampling_params['max_new_tokens']} should not be less than 0" - if sampling_params["max_new_tokens"] == 0: - sample.status = Sample.Status.TRUNCATED - return sample - - # Prepare payload - shared structure - payload = { - "sampling_params": sampling_params, - "return_logprob": args.use_token_output, - } - - if args.use_token_output: - # Token-based mode: use tokens directly - if len(sample.response) > 0: - input_token_ids = sample.tokens - else: - # First turn: initialize with prompt tokens - prompt_token_ids = state.tokenizer(sample.prompt, add_special_tokens=False)["input_ids"] - input_token_ids = prompt_token_ids - # Initialize sample.tokens with prompt for subsequent turns - if not sample.tokens: # Only set if empty - sample.tokens = prompt_token_ids - payload["input_ids"] = input_token_ids - else: - # String-based mode: original implementation - input_text = sample.prompt + sample.response - payload["text"] = input_text - - output = await post(url, payload, use_http2=args.use_http2) - - if args.use_token_output: - # Extract new response tokens - assert ( - "meta_info" in output and "output_token_logprobs" in output["meta_info"] - ), "output_token_logprobs is not in the output" - new_response_tokens = [item[1] for item in output["meta_info"]["output_token_logprobs"]] - - # Update sample with tokens directly - avoiding re-tokenization - sample.tokens = sample.tokens + new_response_tokens - sample.response_length += len(new_response_tokens) - sample.response += state.tokenizer.decode(new_response_tokens, skip_special_tokens=False) - else: - # String-based processing - sample.response += output["text"] - prompt_tokens_ids = state.tokenizer(sample.prompt, add_special_tokens=False)["input_ids"] - response_token_ids = state.tokenizer(sample.response, add_special_tokens=False)["input_ids"] - sample.tokens = prompt_tokens_ids + response_token_ids - sample.response_length = len(response_token_ids) - - match output["meta_info"]["finish_reason"]["type"]: - case "length": - sample.status = Sample.Status.TRUNCATED - case "abort": - sample.status = Sample.Status.ABORTED - case "stop": - sample.status = Sample.Status.COMPLETED - - return sample - - async def generate_and_rm(args, sample: Sample, sampling_params: dict, evaluation=False) -> Sample: # For samples with existing response, check if they're complete if sample.status == Sample.Status.COMPLETED or sample.status == Sample.Status.TRUNCATED: @@ -158,7 +83,7 @@ async def generate_and_rm(args, sample: Sample, sampling_params: dict, evaluatio custom_generate_func = load_function(args.custom_generate_function_path) sample = await custom_generate_func(args, sample, sampling_params) else: - sample = await generate(args, sample, sampling_params) + sample = await generate_one_sample_vanilla(args, state.tokenizer, sample, sampling_params) if sample.status == Sample.Status.ABORTED: return sample