Skip to content
Closed
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 12 additions & 3 deletions miles/rollout/sglang_rollout.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@
from tqdm import tqdm

from miles.backends.megatron_utils.lora_utils import LORA_ADAPTER_NAME, is_lora_enabled
from miles.rollout.base_types import RolloutFnEvalOutput, RolloutFnTrainOutput
from miles.rollout.base_types import GenerateFnInput, RolloutFnEvalOutput, RolloutFnTrainOutput
from miles.rollout.filter_hub.base_types import MetricGatherer, call_dynamic_filter
from miles.utils import dumper_utils
from miles.utils.async_utils import run
Expand Down Expand Up @@ -262,8 +262,17 @@ async def generate_and_rm(

if custom_func_path is not None:
custom_generate_func = load_function(custom_func_path)
# if signature has evaluation, pass evaluation
if "evaluation" in inspect.signature(custom_generate_func).parameters:
sig = inspect.signature(custom_generate_func)
params = list(sig.parameters.values())
# Support GenerateFnInput-style generate functions (single-arg with typed input)
if len(params) == 1:
output = await custom_generate_func(
GenerateFnInput(
state=state, sample=sample, sampling_params=sampling_params, evaluation=evaluation
)
)
sample = output.samples
elif "evaluation" in sig.parameters:
sample = await custom_generate_func(args, sample, sampling_params, evaluation=evaluation)
else:
sample = await custom_generate_func(args, sample, sampling_params)
Expand Down
Loading