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
283 changes: 234 additions & 49 deletions examples/geo3k_vlm_multi_turn/rollout.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,19 +8,25 @@

import torch

# When executed as a module: python -m examples.vlm_multi_turn.rollout
from examples.geo3k_vlm_multi_turn.base_env import BaseInteractionEnv
from vime.rollout.vllm_rollout import (
GenerateState,
_build_inference_sampling_params,
_inference_generate_tokens_and_logprobs,
_mm_render_response_to_generate_body,
)
from vime.utils.http_utils import post
from vime.utils.processing_utils import encode_image_for_rollout_engine
from vime.utils.processing_utils import build_processor_kwargs, encode_image_for_rollout_engine
from vime.utils.types import Sample

DEFAULT_ENV_MODULE = "examples.vlm_multi_turn.env_geo3k"

# Dummy messages used for calculating trim length in chat template encoding.
DUMMY_MESSAGES = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "I am a user."},
]


def _load_env_module(env_path: str | None):
"""Load the interaction environment module from a module path or a file path."""
Expand Down Expand Up @@ -62,17 +68,18 @@ def _content_to_render_format(content: list[dict]) -> list[dict]:


def _build_initial_user_message(sample: Sample) -> dict:
"""Build the initial user message from sample.prompt + sample.multimodal_inputs.images."""
content: list[dict] = []
"""Build the initial user message for the render route (same layout as vllm_rollout MM render)."""
images = (sample.multimodal_inputs or {}).get("images") or []
if not images:
return {"role": "user", "content": sample.prompt}
content: list[dict] = [{"type": "text", "text": sample.prompt}]
for image in images:
content.append({"type": "image", "image": image})
content.append({"type": "text", "text": sample.prompt})
return {"role": "user", "content": content}


def _messages_for_render(messages: list[dict]) -> list[dict]:
"""Normalize per-turn messages to the render-route shape (image image_url)."""
"""Normalize per-turn messages to the render-route shape (image -> image_url)."""
out: list[dict] = []
for msg in messages:
content = msg.get("content")
Expand All @@ -83,23 +90,76 @@ def _messages_for_render(messages: list[dict]) -> list[dict]:
return out


def _processor_features_from_message(processor, tokenizer, message: dict) -> dict | None:
"""Run the HF processor on a single user message containing images, returning train-side multimodal inputs."""
if processor is None:
return None
content = message.get("content")
if not isinstance(content, list):
return None
def _prepare_initial_inputs(sample: Sample, processor, tokenizer) -> tuple[list[int], dict | None]:
"""Initial train-side features from dataset-rendered ``sample.prompt`` (no re-template)."""
multimodal_train_inputs = None
raw_multimodal_inputs = sample.multimodal_inputs or {}
has_multimodal = any(value is not None for value in raw_multimodal_inputs.values())
if processor and has_multimodal:
processor_output = processor(text=sample.prompt, **build_processor_kwargs(raw_multimodal_inputs))
prompt_ids = processor_output["input_ids"][0]
multimodal_train_inputs = {
k: v for k, v in processor_output.items() if k not in ("input_ids", "attention_mask")
} or None
else:
prompt_ids = tokenizer.encode(sample.prompt, add_special_tokens=False)
return list(prompt_ids), multimodal_train_inputs


def _encode_observation_for_generation(
tokenizer,
processor,
message: dict,
metadata: dict | None,
apply_chat_template: bool,
apply_chat_template_kwargs: dict | None,
) -> tuple[list[int], dict | None, dict | None]:
"""Encode a fresh env observation message (may include images/videos)."""
tools = metadata.get("tools") if metadata else None
apply_kwargs = apply_chat_template_kwargs or {}
trim_length = 0

if apply_chat_template:
dummy_prompt = tokenizer.apply_chat_template(
DUMMY_MESSAGES,
tools=tools,
tokenize=False,
add_generation_prompt=False,
**apply_kwargs,
)
formatted_prompt = tokenizer.apply_chat_template(
DUMMY_MESSAGES + [message],
tools=tools,
tokenize=False,
add_generation_prompt=True,
**apply_kwargs,
)
trim_length = len(tokenizer.encode(dummy_prompt, add_special_tokens=False))
else:
formatted_prompt = [message]

multimodal_inputs = None
multimodal_train_inputs = None
if processor:
from qwen_vl_utils import process_vision_info

images, videos = process_vision_info([message])
if images or videos:
multimodal_inputs = {"images": images, "videos": videos}
processor_output = processor(text=formatted_prompt, **build_processor_kwargs(multimodal_inputs))
prompt_ids = processor_output["input_ids"][0]
multimodal_train_inputs = {
k: v for k, v in processor_output.items() if k not in ("input_ids", "attention_mask")
} or None
else:
prompt_ids = tokenizer.encode(formatted_prompt, add_special_tokens=False)
else:
prompt_ids = tokenizer.encode(formatted_prompt, add_special_tokens=False)

from qwen_vl_utils import process_vision_info
if trim_length:
prompt_ids = prompt_ids[trim_length:]

images, videos = process_vision_info([message])
if not images and not videos:
return None
# We only need processor-side features here; tokens are produced by the render route.
formatted_prompt = tokenizer.apply_chat_template([message], tokenize=False, add_generation_prompt=False)
processor_output = processor(text=formatted_prompt, images=images, videos=videos)
return {k: v for k, v in processor_output.items() if k not in ("input_ids", "attention_mask")} or None
return list(prompt_ids), multimodal_inputs, multimodal_train_inputs


def _merge_multimodal_train_inputs(chunks: list[dict | None]) -> dict | None:
Expand All @@ -121,18 +181,115 @@ def _merge_multimodal_train_inputs(chunks: list[dict | None]) -> dict | None:
return merged


async def _render_and_generate(
args,
base_url: str,
messages: list[dict],
sampling_params: dict,
) -> dict:
"""Render messages to engine-prompt body, then call /inference/v1/generate."""
def _append_to_sample(
sample: Sample,
response_tokens: list[int],
tokens_to_add: list[int],
logprobs: list[float],
loss_mask_val: int,
*,
track_response: bool = False,
) -> None:
sample.tokens.extend(tokens_to_add)
if track_response:
response_tokens.extend(tokens_to_add)
sample.loss_mask.extend([loss_mask_val] * len(tokens_to_add))
sample.rollout_log_probs.extend(logprobs)
if track_response:
sample.response_length = len(response_tokens)


def _append_rendered_prefix(
sample: Sample,
response_tokens: list[int],
rendered_ids: list[int],
) -> int:
"""Append render delta with ``loss_mask=0``; fail on prefix drift."""
existing = list(sample.tokens)
if len(rendered_ids) < len(existing):
raise ValueError(
f"render token_ids length ({len(rendered_ids)}) shorter than sample.tokens ({len(existing)})"
)
if rendered_ids[: len(existing)] != existing:
raise ValueError("render token_ids prefix mismatch with sample.tokens (chat template drift)")
delta = rendered_ids[len(existing) :]
if delta:
_append_to_sample(sample, response_tokens, delta, [0.0] * len(delta), loss_mask_val=0, track_response=False)
return len(delta)


def _update_multimodal_state(
sample: Sample,
obs_multimodal_inputs: dict | None,
obs_multimodal_train_inputs: dict | None,
multimodal_train_inputs_buffer: list[dict | None],
) -> None:
if obs_multimodal_inputs:
if not sample.multimodal_inputs:
sample.multimodal_inputs = obs_multimodal_inputs
elif isinstance(sample.multimodal_inputs, dict) and isinstance(obs_multimodal_inputs, dict):
for key, val in obs_multimodal_inputs.items():
if val is None:
continue
if (
key in sample.multimodal_inputs
and isinstance(sample.multimodal_inputs[key], list)
and isinstance(val, list)
):
sample.multimodal_inputs[key].extend(val)
else:
sample.multimodal_inputs[key] = val
else:
sample.multimodal_inputs = obs_multimodal_inputs

if obs_multimodal_train_inputs:
multimodal_train_inputs_buffer.append(obs_multimodal_train_inputs)


async def _render_messages(args, base_url: str, messages: list[dict]) -> tuple[dict, list[int]]:
"""Call ``/v1/chat/completions/render`` and return generate body + prompt token ids."""
render_payload = {"model": args.hf_checkpoint, "messages": _messages_for_render(messages)}
render_data = await post(f"{base_url}/v1/chat/completions/render", render_payload)
body = _mm_render_response_to_generate_body(render_data, args.hf_checkpoint)
body["sampling_params"] = sampling_params
return await post(f"{base_url}/inference/v1/generate", body)
token_ids = body["token_ids"]
if isinstance(token_ids, list) and token_ids and isinstance(token_ids[0], list):
token_ids = token_ids[0]
return body, [int(x) for x in token_ids]


async def _generate_from_render_body(base_url: str, body: dict, sampling_params: dict) -> dict:
generate_body = dict(body)
generate_body["sampling_params"] = sampling_params
return await post(f"{base_url}/inference/v1/generate", generate_body)


def _process_env_step(
env: BaseInteractionEnv,
response_text: str,
tokenizer,
processor,
args: Any,
sample_metadata: dict,
) -> tuple[dict | None, list[int] | None, dict | None, dict | None, bool]:
observation, done, _ = env.step(response_text)
if done:
return None, None, None, None, True

next_user_message = env.format_observation(observation)
obs_prompt_ids, obs_multimodal_inputs, obs_multimodal_train_inputs = _encode_observation_for_generation(
tokenizer,
processor,
next_user_message,
sample_metadata,
getattr(args, "apply_chat_template", False),
getattr(args, "apply_chat_template_kwargs", None),
)

bos_id = tokenizer.bos_token_id
if bos_id is not None and obs_prompt_ids and obs_prompt_ids[0] == bos_id:
obs_prompt_ids = obs_prompt_ids[1:]

return next_user_message, obs_prompt_ids, obs_multimodal_inputs, obs_multimodal_train_inputs, False


async def generate(args: Any, sample: Sample, sampling_params) -> Sample:
Expand All @@ -148,17 +305,19 @@ async def generate(args: Any, sample: Sample, sampling_params) -> Sample:
sample.metadata = sample.metadata or {}
env = _build_env(env_module, sample, args)

messages: list[dict] = [_build_initial_user_message(sample)]
prompt_ids, init_mm_train = _prepare_initial_inputs(sample, state.processor, state.tokenizer)
multimodal_train_inputs_buffer: list[dict | None] = []
if init_mm_train:
multimodal_train_inputs_buffer.append(init_mm_train)

if not sample.tokens:
sample.tokens = []
response_tokens: list[int] = []
sample.loss_mask = sample.loss_mask or []
sample.rollout_log_probs = sample.rollout_log_probs or []
sample.tokens = list(sample.tokens) if sample.tokens else []
multimodal_train_inputs_buffer: list[dict | None] = []

initial_train_feats = _processor_features_from_message(state.processor, state.tokenizer, messages[0])
if initial_train_feats:
multimodal_train_inputs_buffer.append(initial_train_feats)
sample.response_length = len(response_tokens)

messages: list[dict] = [_build_initial_user_message(sample)]
sampling_params = sampling_params.copy()
inference_sampling_params = _build_inference_sampling_params(sampling_params)

Expand All @@ -170,14 +329,27 @@ async def generate(args: Any, sample: Sample, sampling_params) -> Sample:

try:
env.reset()
if budget is not None and budget <= 0:
sample.status = Sample.Status.TRUNCATED
return sample

for turn_idx in range(args.max_turns):
if budget is not None and budget <= 0:
sample.status = Sample.Status.TRUNCATED
break
if budget is not None:
inference_sampling_params["max_tokens"] = budget

output = await _render_and_generate(args, base_url, messages, inference_sampling_params)
render_body, rendered_ids = await _render_messages(args, base_url, messages)
prefix_added = _append_rendered_prefix(sample, response_tokens, rendered_ids)
if budget is not None:
budget -= prefix_added
if budget <= 0:
sample.status = Sample.Status.TRUNCATED
break
inference_sampling_params["max_tokens"] = budget

output = await _generate_from_render_body(base_url, render_body, inference_sampling_params)
choice = output["choices"][0]
finish_reason = choice.get("finish_reason") or "stop"
new_tokens, new_logprobs = _inference_generate_tokens_and_logprobs(choice)
Expand All @@ -188,11 +360,9 @@ async def generate(args: Any, sample: Sample, sampling_params) -> Sample:
break

response_text = state.tokenizer.decode(new_tokens, skip_special_tokens=False) if new_tokens else ""
sample.tokens.extend(new_tokens)
response_tokens.extend(new_tokens)
sample.loss_mask.extend([1] * len(new_tokens))
sample.rollout_log_probs.extend(new_logprobs)
sample.response_length = len(response_tokens)
_append_to_sample(
sample, response_tokens, new_tokens, new_logprobs, loss_mask_val=1, track_response=True
)
if budget is not None:
budget -= len(new_tokens)

Expand All @@ -205,18 +375,33 @@ async def generate(args: Any, sample: Sample, sampling_params) -> Sample:
sample.status = Sample.Status.ABORTED
break

observation, done, _ = env.step(response_text)
next_user_message, obs_prompt_ids, obs_multimodal_inputs, obs_multimodal_train_inputs, done = (
_process_env_step(env, response_text, state.tokenizer, state.processor, args, sample.metadata)
)
if done:
sample.status = Sample.Status.COMPLETED
break

next_user_message = env.format_observation(observation)
messages.append(next_user_message)
assert obs_prompt_ids is not None
_append_to_sample(
sample,
response_tokens,
obs_prompt_ids,
[0.0] * len(obs_prompt_ids),
loss_mask_val=0,
track_response=False,
)
if budget is not None:
budget -= len(obs_prompt_ids)

obs_train_feats = _processor_features_from_message(state.processor, state.tokenizer, next_user_message)
if obs_train_feats:
multimodal_train_inputs_buffer.append(obs_train_feats)
_update_multimodal_state(
sample, obs_multimodal_inputs, obs_multimodal_train_inputs, multimodal_train_inputs_buffer
)
messages.append(next_user_message)

if budget is not None and budget <= 0:
sample.status = Sample.Status.TRUNCATED
break
if turn_idx + 1 >= args.max_turns:
sample.status = Sample.Status.COMPLETED
break
Expand Down
Loading