[vla] fix: pi05 models does not apply scale factor for language embeddings (#33367)
Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
@@ -426,7 +426,7 @@ The benchmark reports:
|
|||||||
- SGLang Python in-process latency when `--sglang-api python` is set. This loads the native Pi0.5 pipeline in the benchmark process and avoids HTTP, websocket, scheduler, and serialization overhead. Use `--sglang-python-batch-mode grouped` to exercise the conservative native grouped-batch path. The Python path also reports actual SGLang module parameter dtype counts and example parameter names.
|
- SGLang Python in-process latency when `--sglang-api python` is set. This loads the native Pi0.5 pipeline in the benchmark process and avoids HTTP, websocket, scheduler, and serialization overhead. Use `--sglang-python-batch-mode grouped` to exercise the conservative native grouped-batch path. The Python path also reports actual SGLang module parameter dtype counts and example parameter names.
|
||||||
- OpenPI single-request latency through `Policy.infer`.
|
- OpenPI single-request latency through `Policy.infer`.
|
||||||
- Batch latency for grouped robot streams. SGLang uses concurrent HTTP requests in HTTP mode and persistent multi-connection msgpack calls in websocket mode. The Python backend can use true grouped model execution for fresh-prefix requests. OpenPI defaults to its internal direct model batch path because the public `Policy.infer` API is single-observation.
|
- Batch latency for grouped robot streams. SGLang uses concurrent HTTP requests in HTTP mode and persistent multi-connection msgpack calls in websocket mode. The Python backend can use true grouped model execution for fresh-prefix requests. OpenPI defaults to its internal direct model batch path because the public `Policy.infer` API is single-observation.
|
||||||
- Action difference on the common output prefix. Use `--deterministic-noise` for strict debugging when the SGLang and OpenPI horizons match. The `aloha` profile supports this directly; the LIBERO profile compares the common prefix because OpenPI's released LIBERO config uses a shorter output horizon than the LeRobot Pi0.5 checkpoint metadata.
|
- Action difference in normalized model space from identical OpenPI-transformed model inputs and noise. The check requires `--deterministic-noise` and fails when either `--action-max-abs-diff` or `--action-mean-abs-diff` is exceeded. This mode isolates model parity from robot-specific normalization and action postprocessing. The LIBERO policy returns only 10 actions after policy postprocessing, but its flow-matching model still generates a 50-step chunk; the benchmark compares that model output with SGLang before OpenPI unnormalization and horizon slicing.
|
||||||
|
|
||||||
For one-sided 16GB-class checks, run each backend separately under the same VRAM pressure. The SGLang Python path accepts the same pipeline config override as serving:
|
For one-sided 16GB-class checks, run each backend separately under the same VRAM pressure. The SGLang Python path accepts the same pipeline config override as serving:
|
||||||
|
|
||||||
@@ -455,7 +455,7 @@ The following checks were run on H100 GPUs with the native SGLang Pi0.5 path:
|
|||||||
| Check | Result |
|
| Check | Result |
|
||||||
| --- | --- |
|
| --- | --- |
|
||||||
| `lerobot/pi05_base` direct end-to-end | Prefix length `968`, output shape `[1, 50, 32]`, peak allocated memory `12.817 GiB`. |
|
| `lerobot/pi05_base` direct end-to-end | Prefix length `968`, output shape `[1, 50, 32]`, peak allocated memory `12.817 GiB`. |
|
||||||
| LeRobot reference parity | One-step velocity max absolute difference `1.17e-6`; final 10-step action max absolute difference `1.17e-7`. |
|
| Official OpenPI parity | Against OpenPI PyTorch revision `15a9616`, with the same LeRobot checkpoint revision, observation, and noise: first-step velocity max/mean absolute difference `0.02677` / `0.00344`; production 10-step normalized action `0.00813` / `0.00092`. |
|
||||||
| Action denoise CUDA graph | Eager 10-step denoise `125.4 ms`; steady graph replay `50.8 ms`; max output difference `0`. |
|
| Action denoise CUDA graph | Eager 10-step denoise `125.4 ms`; steady graph replay `50.8 ms`; max output difference `0`. |
|
||||||
| Exact full-prefix cache | First prefix pass about `203 ms`; exact cache hit prefix stage about `0.2 ms`. |
|
| Exact full-prefix cache | First prefix pass about `203 ms`; exact cache hit prefix stage about `0.2 ms`. |
|
||||||
| `lerobot/pi05_libero_base` direct end-to-end | Image keys `image`, `image2`, `empty_camera_0`; state dim `8`; output action dim `7`; output tensor shape `[1, 50, 32]`. |
|
| `lerobot/pi05_libero_base` direct end-to-end | Image keys `image`, `image2`, `empty_camera_0`; state dim `8`; output action dim `7`; output tensor shape `[1, 50, 32]`. |
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import argparse
|
import argparse
|
||||||
import asyncio
|
import asyncio
|
||||||
|
import copy
|
||||||
import dataclasses
|
import dataclasses
|
||||||
import json
|
import json
|
||||||
import sys
|
import sys
|
||||||
@@ -53,7 +54,7 @@ PROFILES = {
|
|||||||
openpi_checkpoint="gs://openpi-assets/checkpoints/pi05_libero",
|
openpi_checkpoint="gs://openpi-assets/checkpoints/pi05_libero",
|
||||||
prompt="pick up the object",
|
prompt="pick up the object",
|
||||||
sglang_action_horizon=50,
|
sglang_action_horizon=50,
|
||||||
openpi_action_horizon=10,
|
openpi_action_horizon=50,
|
||||||
action_dim=32,
|
action_dim=32,
|
||||||
output_action_dim=7,
|
output_action_dim=7,
|
||||||
),
|
),
|
||||||
@@ -231,15 +232,13 @@ def _make_aloha_observation(
|
|||||||
},
|
},
|
||||||
"prompt": prompt,
|
"prompt": prompt,
|
||||||
}
|
}
|
||||||
sglang_state = np.zeros((32,), dtype=np.float32)
|
|
||||||
sglang_state[: state.shape[0]] = state
|
|
||||||
sglang_observation = {
|
sglang_observation = {
|
||||||
"images": {
|
"images": {
|
||||||
"base_0_rgb": np.transpose(cam_high, (1, 2, 0)),
|
"base_0_rgb": np.transpose(cam_high, (1, 2, 0)),
|
||||||
"left_wrist_0_rgb": np.transpose(cam_left, (1, 2, 0)),
|
"left_wrist_0_rgb": np.transpose(cam_left, (1, 2, 0)),
|
||||||
"right_wrist_0_rgb": np.transpose(cam_right, (1, 2, 0)),
|
"right_wrist_0_rgb": np.transpose(cam_right, (1, 2, 0)),
|
||||||
},
|
},
|
||||||
"state": sglang_state,
|
"state": state,
|
||||||
}
|
}
|
||||||
return openpi_obs, sglang_observation
|
return openpi_obs, sglang_observation
|
||||||
|
|
||||||
@@ -264,6 +263,33 @@ def build_observations(
|
|||||||
return openpi_observations, sglang_observations
|
return openpi_observations, sglang_observations
|
||||||
|
|
||||||
|
|
||||||
|
def build_openpi_model_inputs(
|
||||||
|
policy,
|
||||||
|
observations: list[dict[str, Any]],
|
||||||
|
) -> list[dict[str, Any]]:
|
||||||
|
"""Convert raw robot observations to the exact inputs consumed by OpenPI."""
|
||||||
|
model_inputs = []
|
||||||
|
for observation in observations:
|
||||||
|
transformed = policy._input_transform(copy.deepcopy(observation))
|
||||||
|
images = {
|
||||||
|
name: np.asarray(value) for name, value in transformed["image"].items()
|
||||||
|
}
|
||||||
|
model_inputs.append(
|
||||||
|
{
|
||||||
|
"images": images,
|
||||||
|
"image_masks": {
|
||||||
|
name: bool(value)
|
||||||
|
for name, value in transformed["image_mask"].items()
|
||||||
|
},
|
||||||
|
"camera_order": tuple(images),
|
||||||
|
"state": np.asarray(transformed["state"]),
|
||||||
|
"tokens": np.asarray(transformed["tokenized_prompt"]),
|
||||||
|
"token_masks": np.asarray(transformed["tokenized_prompt_mask"]),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
return model_inputs
|
||||||
|
|
||||||
|
|
||||||
def _json_tensor(array: np.ndarray) -> dict[str, Any]:
|
def _json_tensor(array: np.ndarray) -> dict[str, Any]:
|
||||||
return {
|
return {
|
||||||
"dtype": str(array.dtype),
|
"dtype": str(array.dtype),
|
||||||
@@ -290,6 +316,13 @@ def build_sglang_payload(
|
|||||||
"images": encoded_images,
|
"images": encoded_images,
|
||||||
"state": _json_tensor(np.asarray(observation["state"], dtype=np.float32)),
|
"state": _json_tensor(np.asarray(observation["state"], dtype=np.float32)),
|
||||||
}
|
}
|
||||||
|
if "tokens" in observation:
|
||||||
|
encoded_observation["tokens"] = np.asarray(observation["tokens"]).tolist()
|
||||||
|
encoded_observation["token_masks"] = np.asarray(
|
||||||
|
observation["token_masks"]
|
||||||
|
).tolist()
|
||||||
|
encoded_observation["image_masks"] = observation["image_masks"]
|
||||||
|
encoded_observation["camera_order"] = list(observation["camera_order"])
|
||||||
if noise is not None:
|
if noise is not None:
|
||||||
encoded_observation["noise"] = _json_tensor(noise.astype(np.float32))
|
encoded_observation["noise"] = _json_tensor(noise.astype(np.float32))
|
||||||
return {
|
return {
|
||||||
@@ -326,6 +359,11 @@ def build_sglang_python_payload(
|
|||||||
},
|
},
|
||||||
"state": np.asarray(observation["state"], dtype=np.float32),
|
"state": np.asarray(observation["state"], dtype=np.float32),
|
||||||
}
|
}
|
||||||
|
if "tokens" in observation:
|
||||||
|
encoded_observation["tokens"] = np.asarray(observation["tokens"])
|
||||||
|
encoded_observation["token_masks"] = np.asarray(observation["token_masks"])
|
||||||
|
encoded_observation["image_masks"] = observation["image_masks"]
|
||||||
|
encoded_observation["camera_order"] = observation["camera_order"]
|
||||||
if noise is not None:
|
if noise is not None:
|
||||||
encoded_observation["noise"] = noise.astype(np.float32)
|
encoded_observation["noise"] = noise.astype(np.float32)
|
||||||
return {
|
return {
|
||||||
@@ -366,6 +404,11 @@ def build_sglang_openpi_ws_payload(
|
|||||||
}
|
}
|
||||||
for key, value in observation["images"].items():
|
for key, value in observation["images"].items():
|
||||||
payload[f"observation.images.{key}"] = np.asarray(value)
|
payload[f"observation.images.{key}"] = np.asarray(value)
|
||||||
|
if "tokens" in observation:
|
||||||
|
payload["tokens"] = np.asarray(observation["tokens"])
|
||||||
|
payload["token_masks"] = np.asarray(observation["token_masks"])
|
||||||
|
payload["image_masks"] = observation["image_masks"]
|
||||||
|
payload["camera_order"] = observation["camera_order"]
|
||||||
if noise is not None:
|
if noise is not None:
|
||||||
payload["observation.noise"] = noise.astype(np.float32)
|
payload["observation.noise"] = noise.astype(np.float32)
|
||||||
return payload
|
return payload
|
||||||
@@ -791,11 +834,11 @@ def _openpi_infer(policy, observation: dict[str, Any], noise: np.ndarray | None)
|
|||||||
return policy.infer(observation, noise=noise)
|
return policy.infer(observation, noise=noise)
|
||||||
|
|
||||||
|
|
||||||
def _openpi_direct_batch(
|
def _openpi_model_batch(
|
||||||
policy,
|
policy,
|
||||||
observations: list[dict[str, Any]],
|
observations: list[dict[str, Any]],
|
||||||
noises: list[np.ndarray] | None,
|
noises: list[np.ndarray] | None,
|
||||||
):
|
) -> tuple[np.ndarray, np.ndarray]:
|
||||||
import jax
|
import jax
|
||||||
import numpy as onp
|
import numpy as onp
|
||||||
from openpi.models import model as openpi_model
|
from openpi.models import model as openpi_model
|
||||||
@@ -845,6 +888,15 @@ def _openpi_direct_batch(
|
|||||||
actions_np = onp.asarray(actions)
|
actions_np = onp.asarray(actions)
|
||||||
states_np = onp.asarray(inputs["state"])
|
states_np = onp.asarray(inputs["state"])
|
||||||
|
|
||||||
|
return actions_np, states_np
|
||||||
|
|
||||||
|
|
||||||
|
def _openpi_direct_batch(
|
||||||
|
policy,
|
||||||
|
observations: list[dict[str, Any]],
|
||||||
|
noises: list[np.ndarray] | None,
|
||||||
|
):
|
||||||
|
actions_np, states_np = _openpi_model_batch(policy, observations, noises)
|
||||||
outputs = []
|
outputs = []
|
||||||
for idx in range(actions_np.shape[0]):
|
for idx in range(actions_np.shape[0]):
|
||||||
outputs.append(
|
outputs.append(
|
||||||
@@ -858,6 +910,15 @@ def _openpi_direct_batch(
|
|||||||
return outputs
|
return outputs
|
||||||
|
|
||||||
|
|
||||||
|
def _openpi_model_actions(
|
||||||
|
policy,
|
||||||
|
observation: dict[str, Any],
|
||||||
|
noise: np.ndarray,
|
||||||
|
) -> np.ndarray:
|
||||||
|
actions, _ = _openpi_model_batch(policy, [observation], [noise])
|
||||||
|
return actions[0]
|
||||||
|
|
||||||
|
|
||||||
def run_openpi_policy(
|
def run_openpi_policy(
|
||||||
policy,
|
policy,
|
||||||
observations: list[dict[str, Any]],
|
observations: list[dict[str, Any]],
|
||||||
@@ -922,6 +983,12 @@ def run_openpi_policy(
|
|||||||
precision["output_action_dtype"] = str(first_actions.dtype)
|
precision["output_action_dtype"] = str(first_actions.dtype)
|
||||||
precision["output_action_shape"] = list(first_actions.shape)
|
precision["output_action_shape"] = list(first_actions.shape)
|
||||||
|
|
||||||
|
first_model_actions = (
|
||||||
|
_openpi_model_actions(policy, observations[0], noise)
|
||||||
|
if observations and noise is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"single": _stats_ms(single_latencies),
|
"single": _stats_ms(single_latencies),
|
||||||
"batch": _stats_ms(batch_latencies),
|
"batch": _stats_ms(batch_latencies),
|
||||||
@@ -930,6 +997,7 @@ def run_openpi_policy(
|
|||||||
key: _stats_ms(values) for key, values in policy_timings.items()
|
key: _stats_ms(values) for key, values in policy_timings.items()
|
||||||
},
|
},
|
||||||
"first_output": single_outputs[0] if single_outputs else None,
|
"first_output": single_outputs[0] if single_outputs else None,
|
||||||
|
"first_model_actions": first_model_actions,
|
||||||
"batch_mode": batch_mode,
|
"batch_mode": batch_mode,
|
||||||
"precision": precision,
|
"precision": precision,
|
||||||
}
|
}
|
||||||
@@ -951,10 +1019,18 @@ def _openpi_actions(output: dict[str, Any]) -> np.ndarray | None:
|
|||||||
|
|
||||||
def compare_first_actions(
|
def compare_first_actions(
|
||||||
sglang_output: dict[str, Any] | None,
|
sglang_output: dict[str, Any] | None,
|
||||||
openpi_output: dict[str, Any] | None,
|
openpi_output: dict[str, Any] | np.ndarray | None,
|
||||||
) -> dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
sglang_actions = _sglang_actions(sglang_output)
|
sglang_actions = _sglang_actions(sglang_output)
|
||||||
openpi_actions = _openpi_actions(openpi_output)
|
openpi_actions = (
|
||||||
|
_openpi_actions(openpi_output)
|
||||||
|
if isinstance(openpi_output, dict)
|
||||||
|
else (
|
||||||
|
np.asarray(openpi_output, dtype=np.float32)
|
||||||
|
if openpi_output is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
)
|
||||||
if sglang_actions is None or openpi_actions is None:
|
if sglang_actions is None or openpi_actions is None:
|
||||||
return {"available": False}
|
return {"available": False}
|
||||||
horizon = min(sglang_actions.shape[0], openpi_actions.shape[0])
|
horizon = min(sglang_actions.shape[0], openpi_actions.shape[0])
|
||||||
@@ -1051,6 +1127,8 @@ def parse_args() -> argparse.Namespace:
|
|||||||
parser.add_argument("--disable-prefix-cache", action="store_true")
|
parser.add_argument("--disable-prefix-cache", action="store_true")
|
||||||
parser.add_argument("--disable-cuda-graph", action="store_true")
|
parser.add_argument("--disable-cuda-graph", action="store_true")
|
||||||
parser.add_argument("--deterministic-noise", action="store_true")
|
parser.add_argument("--deterministic-noise", action="store_true")
|
||||||
|
parser.add_argument("--action-max-abs-diff", type=float, default=0.05)
|
||||||
|
parser.add_argument("--action-mean-abs-diff", type=float, default=0.005)
|
||||||
parser.add_argument("--skip-sglang", action="store_true")
|
parser.add_argument("--skip-sglang", action="store_true")
|
||||||
parser.add_argument("--skip-openpi", action="store_true")
|
parser.add_argument("--skip-openpi", action="store_true")
|
||||||
parser.add_argument("--output-file", default="")
|
parser.add_argument("--output-file", default="")
|
||||||
@@ -1096,6 +1174,24 @@ def main() -> None:
|
|||||||
dtype=np.float32,
|
dtype=np.float32,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
openpi_policy = None
|
||||||
|
if not args.skip_openpi:
|
||||||
|
openpi_policy = create_openpi_policy(
|
||||||
|
openpi_config,
|
||||||
|
openpi_checkpoint,
|
||||||
|
pytorch_device=args.openpi_device,
|
||||||
|
num_inference_steps=args.num_inference_steps,
|
||||||
|
pytorch_compile_mode=(
|
||||||
|
None
|
||||||
|
if args.openpi_pytorch_compile_mode == "none"
|
||||||
|
else args.openpi_pytorch_compile_mode
|
||||||
|
),
|
||||||
|
)
|
||||||
|
if not args.skip_sglang and args.deterministic_noise:
|
||||||
|
sglang_observations = build_openpi_model_inputs(
|
||||||
|
openpi_policy,
|
||||||
|
openpi_observations,
|
||||||
|
)
|
||||||
payloads = []
|
payloads = []
|
||||||
if args.skip_sglang:
|
if args.skip_sglang:
|
||||||
pass
|
pass
|
||||||
@@ -1138,20 +1234,6 @@ def main() -> None:
|
|||||||
for observation in sglang_observations
|
for observation in sglang_observations
|
||||||
]
|
]
|
||||||
|
|
||||||
openpi_policy = None
|
|
||||||
if not args.skip_openpi:
|
|
||||||
openpi_policy = create_openpi_policy(
|
|
||||||
openpi_config,
|
|
||||||
openpi_checkpoint,
|
|
||||||
pytorch_device=args.openpi_device,
|
|
||||||
num_inference_steps=args.num_inference_steps,
|
|
||||||
pytorch_compile_mode=(
|
|
||||||
None
|
|
||||||
if args.openpi_pytorch_compile_mode == "none"
|
|
||||||
else args.openpi_pytorch_compile_mode
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
sglang_result = None
|
sglang_result = None
|
||||||
if args.skip_sglang:
|
if args.skip_sglang:
|
||||||
pass
|
pass
|
||||||
@@ -1210,9 +1292,20 @@ def main() -> None:
|
|||||||
"repeats": args.repeats,
|
"repeats": args.repeats,
|
||||||
"warmup": args.warmup,
|
"warmup": args.warmup,
|
||||||
"deterministic_noise": args.deterministic_noise,
|
"deterministic_noise": args.deterministic_noise,
|
||||||
"action_diff": compare_first_actions(
|
"action_diff": (
|
||||||
|
compare_first_actions(
|
||||||
None if sglang_result is None else sglang_result.get("first_output"),
|
None if sglang_result is None else sglang_result.get("first_output"),
|
||||||
None if openpi_result is None else openpi_result.get("first_output"),
|
(
|
||||||
|
None
|
||||||
|
if openpi_result is None
|
||||||
|
else openpi_result.get("first_model_actions")
|
||||||
|
),
|
||||||
|
)
|
||||||
|
if args.deterministic_noise
|
||||||
|
else {
|
||||||
|
"available": False,
|
||||||
|
"reason": "use --deterministic-noise for action comparison",
|
||||||
|
}
|
||||||
),
|
),
|
||||||
"sglang": (
|
"sglang": (
|
||||||
None
|
None
|
||||||
@@ -1229,7 +1322,7 @@ def main() -> None:
|
|||||||
else {
|
else {
|
||||||
key: value
|
key: value
|
||||||
for key, value in openpi_result.items()
|
for key, value in openpi_result.items()
|
||||||
if key not in ("first_output",)
|
if key not in ("first_output", "first_model_actions")
|
||||||
}
|
}
|
||||||
),
|
),
|
||||||
}
|
}
|
||||||
@@ -1238,6 +1331,19 @@ def main() -> None:
|
|||||||
json.dump(result, f, indent=2, sort_keys=True)
|
json.dump(result, f, indent=2, sort_keys=True)
|
||||||
print_summary(result)
|
print_summary(result)
|
||||||
|
|
||||||
|
action_diff = result["action_diff"]
|
||||||
|
if action_diff.get("available") and (
|
||||||
|
action_diff["max_abs_diff"] > args.action_max_abs_diff
|
||||||
|
or action_diff["mean_abs_diff"] > args.action_mean_abs_diff
|
||||||
|
):
|
||||||
|
raise AssertionError(
|
||||||
|
"Pi0.5 action mismatch: "
|
||||||
|
f"max_abs_diff={action_diff['max_abs_diff']:.6f} "
|
||||||
|
f"(threshold {args.action_max_abs_diff:.6f}), "
|
||||||
|
f"mean_abs_diff={action_diff['mean_abs_diff']:.6f} "
|
||||||
|
f"(threshold {args.action_mean_abs_diff:.6f})"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
main()
|
main()
|
||||||
|
|||||||
@@ -1336,6 +1336,8 @@ class Pi05CoreModel(nn.Module):
|
|||||||
att_masks += [0] * num_image_embs
|
att_masks += [0] * num_image_embs
|
||||||
|
|
||||||
lang_emb = self.paligemma_with_expert.embed_language_tokens(tokens)
|
lang_emb = self.paligemma_with_expert.embed_language_tokens(tokens)
|
||||||
|
# Match OpenPI's Pi0.5 prefix embedding semantics.
|
||||||
|
lang_emb = lang_emb * math.sqrt(lang_emb.shape[-1])
|
||||||
embs.append(lang_emb)
|
embs.append(lang_emb)
|
||||||
pad_masks.append(token_masks)
|
pad_masks.append(token_masks)
|
||||||
att_masks += [0] * lang_emb.shape[1]
|
att_masks += [0] * lang_emb.shape[1]
|
||||||
|
|||||||
+33
-11
@@ -14,7 +14,7 @@ from sglang.multimodal_gen.configs.pipeline_configs.pi05 import Pi05PipelineConf
|
|||||||
from sglang.multimodal_gen.runtime.vla.observation import VLAObservationBatch
|
from sglang.multimodal_gen.runtime.vla.observation import VLAObservationBatch
|
||||||
|
|
||||||
|
|
||||||
def _tensor_from_image(value: Any) -> torch.Tensor:
|
def _tensor_from_image(value: Any) -> tuple[torch.Tensor, bool, bool]:
|
||||||
if isinstance(value, torch.Tensor):
|
if isinstance(value, torch.Tensor):
|
||||||
tensor = value.detach()
|
tensor = value.detach()
|
||||||
if tensor.ndim == 4:
|
if tensor.ndim == 4:
|
||||||
@@ -33,14 +33,16 @@ def _tensor_from_image(value: Any) -> torch.Tensor:
|
|||||||
)
|
)
|
||||||
is_integer = not tensor.is_floating_point()
|
is_integer = not tensor.is_floating_point()
|
||||||
tensor = tensor.to(dtype=torch.float32)
|
tensor = tensor.to(dtype=torch.float32)
|
||||||
if is_integer or tensor.max() > 2.0:
|
is_byte_scaled = is_integer or tensor.max() > 2.0
|
||||||
|
is_normalized = not is_byte_scaled and tensor.min() < 0.0
|
||||||
|
if is_byte_scaled:
|
||||||
tensor = tensor / 255.0
|
tensor = tensor / 255.0
|
||||||
return tensor
|
return tensor, is_byte_scaled, is_normalized
|
||||||
|
|
||||||
if isinstance(value, Image.Image):
|
if isinstance(value, Image.Image):
|
||||||
image = value.convert("RGB")
|
image = value.convert("RGB")
|
||||||
arr = np.asarray(image, dtype=np.float32) / 255.0
|
arr = np.asarray(image, dtype=np.float32) / 255.0
|
||||||
return torch.from_numpy(arr).permute(2, 0, 1)
|
return torch.from_numpy(arr).permute(2, 0, 1), True, False
|
||||||
|
|
||||||
if isinstance(value, (np.ndarray, list)):
|
if isinstance(value, (np.ndarray, list)):
|
||||||
arr = np.asarray(value)
|
arr = np.asarray(value)
|
||||||
@@ -57,15 +59,21 @@ def _tensor_from_image(value: Any) -> torch.Tensor:
|
|||||||
)
|
)
|
||||||
is_integer = not tensor.is_floating_point()
|
is_integer = not tensor.is_floating_point()
|
||||||
tensor = tensor.to(dtype=torch.float32)
|
tensor = tensor.to(dtype=torch.float32)
|
||||||
if is_integer or tensor.max() > 2.0:
|
is_byte_scaled = is_integer or tensor.max() > 2.0
|
||||||
|
is_normalized = not is_byte_scaled and tensor.min() < 0.0
|
||||||
|
if is_byte_scaled:
|
||||||
tensor = tensor / 255.0
|
tensor = tensor / 255.0
|
||||||
return tensor
|
return tensor, is_byte_scaled, is_normalized
|
||||||
|
|
||||||
raise TypeError(f"Unsupported Pi05 image type: {type(value)}")
|
raise TypeError(f"Unsupported Pi05 image type: {type(value)}")
|
||||||
|
|
||||||
|
|
||||||
def _resize_with_pad_image_tensor(
|
def _resize_with_pad_image_tensor(
|
||||||
tensor: torch.Tensor, size: tuple[int, int]
|
tensor: torch.Tensor,
|
||||||
|
size: tuple[int, int],
|
||||||
|
*,
|
||||||
|
round_to_uint8: bool = False,
|
||||||
|
pad_value: float = 0.0,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
height, width = size
|
height, width = size
|
||||||
if tensor.shape[-2:] == (height, width):
|
if tensor.shape[-2:] == (height, width):
|
||||||
@@ -80,16 +88,32 @@ def _resize_with_pad_image_tensor(
|
|||||||
mode="bilinear",
|
mode="bilinear",
|
||||||
align_corners=False,
|
align_corners=False,
|
||||||
)[0]
|
)[0]
|
||||||
|
if round_to_uint8:
|
||||||
|
# openpi rounds resized byte images before mapping them to [-1, 1]
|
||||||
|
tensor = torch.round(tensor * 255.0).clamp_(0.0, 255.0) / 255.0
|
||||||
pad_h0, rem_h = divmod(height - resized_height, 2)
|
pad_h0, rem_h = divmod(height - resized_height, 2)
|
||||||
pad_w0, rem_w = divmod(width - resized_width, 2)
|
pad_w0, rem_w = divmod(width - resized_width, 2)
|
||||||
return F.pad(
|
return F.pad(
|
||||||
tensor,
|
tensor,
|
||||||
(pad_w0, pad_w0 + rem_w, pad_h0, pad_h0 + rem_h),
|
(pad_w0, pad_w0 + rem_w, pad_h0, pad_h0 + rem_h),
|
||||||
mode="constant",
|
mode="constant",
|
||||||
value=0.0,
|
value=pad_value,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _preprocess_image(value: Any, size: tuple[int, int]) -> torch.Tensor:
|
||||||
|
tensor, is_byte_scaled, is_normalized = _tensor_from_image(value)
|
||||||
|
tensor = _resize_with_pad_image_tensor(
|
||||||
|
tensor,
|
||||||
|
size,
|
||||||
|
round_to_uint8=is_byte_scaled,
|
||||||
|
pad_value=-1.0 if is_normalized else 0.0,
|
||||||
|
)
|
||||||
|
if is_normalized:
|
||||||
|
return tensor.clamp_(-1.0, 1.0)
|
||||||
|
return tensor * 2.0 - 1.0
|
||||||
|
|
||||||
|
|
||||||
class Pi05Preprocessor:
|
class Pi05Preprocessor:
|
||||||
def __init__(self, config: Pi05PipelineConfig):
|
def __init__(self, config: Pi05PipelineConfig):
|
||||||
self.config = config
|
self.config = config
|
||||||
@@ -142,9 +166,7 @@ class Pi05Preprocessor:
|
|||||||
value = raw_images.get(key)
|
value = raw_images.get(key)
|
||||||
is_present = value is not None and bool(image_masks_in.get(key, True))
|
is_present = value is not None and bool(image_masks_in.get(key, True))
|
||||||
if is_present:
|
if is_present:
|
||||||
tensor = _tensor_from_image(value)
|
tensor = _preprocess_image(value, self.config.image_size)
|
||||||
tensor = _resize_with_pad_image_tensor(tensor, self.config.image_size)
|
|
||||||
tensor = tensor * 2.0 - 1.0
|
|
||||||
else:
|
else:
|
||||||
channels = 3
|
channels = 3
|
||||||
height, width = self.config.image_size
|
height, width = self.config.image_size
|
||||||
|
|||||||
@@ -370,7 +370,7 @@ PI05_ACTION_CI_sampling_params = DiffusionSamplingParams(
|
|||||||
extras={
|
extras={
|
||||||
"action_horizon": 50,
|
"action_horizon": 50,
|
||||||
"action_dim": 32,
|
"action_dim": 32,
|
||||||
"state_dim": 32,
|
"state_dim": 14,
|
||||||
"image_size": 64,
|
"image_size": 64,
|
||||||
"num_inference_steps": 2,
|
"num_inference_steps": 2,
|
||||||
"seed": 0,
|
"seed": 0,
|
||||||
|
|||||||
@@ -58,7 +58,7 @@ def _action_request_kwargs(tag: str) -> dict:
|
|||||||
"state": np.linspace(
|
"state": np.linspace(
|
||||||
-0.5,
|
-0.5,
|
||||||
0.5,
|
0.5,
|
||||||
_env_int("SGLANG_PI05_E2E_STATE_DIM", 32),
|
_env_int("SGLANG_PI05_E2E_STATE_DIM", 14),
|
||||||
dtype=np.float32,
|
dtype=np.float32,
|
||||||
),
|
),
|
||||||
"noise": rng.standard_normal((action_horizon, action_dim)).astype(np.float32),
|
"noise": rng.standard_normal((action_horizon, action_dim)).astype(np.float32),
|
||||||
@@ -138,9 +138,9 @@ def test_pi05_python_action_e2e(pi05_generator):
|
|||||||
_assert_action_output(output)
|
_assert_action_output(output)
|
||||||
|
|
||||||
|
|
||||||
def test_pi05_python_action_consistency(pi05_generator):
|
def test_pi05_python_action_repeatability_and_cache(pi05_generator):
|
||||||
first = pi05_generator.generate_action(_action_request_kwargs("consistency"))
|
first = pi05_generator.generate_action(_action_request_kwargs("repeatability"))
|
||||||
second = pi05_generator.generate_action(_action_request_kwargs("consistency"))
|
second = pi05_generator.generate_action(_action_request_kwargs("repeatability"))
|
||||||
_assert_action_output(first, expect_cache_hit=False)
|
_assert_action_output(first, expect_cache_hit=False)
|
||||||
_assert_action_output(second, expect_cache_hit=True)
|
_assert_action_output(second, expect_cache_hit=True)
|
||||||
|
|
||||||
|
|||||||
@@ -39,7 +39,7 @@ logger = init_logger(__name__)
|
|||||||
# NPU/ascend) is read from sgl-project/ci-data-diffusion, where the GT-gen workflows
|
# NPU/ascend) is read from sgl-project/ci-data-diffusion, where the GT-gen workflows
|
||||||
# publish.
|
# publish.
|
||||||
SGL_TEST_FILES_CI_DATA_REPO = "sgl-project/ci-data-diffusion"
|
SGL_TEST_FILES_CI_DATA_REPO = "sgl-project/ci-data-diffusion"
|
||||||
SGL_TEST_FILES_CI_DATA_REVISION = "739c6c9b7cb972149cc3472cf19fe4bb29cf15c3"
|
SGL_TEST_FILES_CI_DATA_REVISION = "d05810e3ea3eff1d137dec723f6e66d9c11b470f"
|
||||||
|
|
||||||
# The NPU pin is kept as a separate branch so ascend GT can be bumped independently
|
# The NPU pin is kept as a separate branch so ascend GT can be bumped independently
|
||||||
# when it's regenerated on its own cadence.
|
# when it's regenerated on its own cadence.
|
||||||
|
|||||||
@@ -2,12 +2,14 @@
|
|||||||
|
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
from torch import nn
|
from torch import nn
|
||||||
|
|
||||||
import sglang.multimodal_gen.runtime.models.vlas.pi05_policy as pi05_policy_module
|
import sglang.multimodal_gen.runtime.models.vlas.pi05_policy as pi05_policy_module
|
||||||
from sglang.multimodal_gen.configs.pipeline_configs.pi05 import Pi05PipelineConfig
|
from sglang.multimodal_gen.configs.pipeline_configs.pi05 import Pi05PipelineConfig
|
||||||
from sglang.multimodal_gen.runtime.models.vlas.pi05_core import (
|
from sglang.multimodal_gen.runtime.models.vlas.pi05_core import (
|
||||||
|
Pi05CoreModel,
|
||||||
Pi05SiglipAttention,
|
Pi05SiglipAttention,
|
||||||
patch_siglip_vision_attention_to_native,
|
patch_siglip_vision_attention_to_native,
|
||||||
)
|
)
|
||||||
@@ -15,6 +17,10 @@ from sglang.multimodal_gen.runtime.models.vlas.pi05_policy import (
|
|||||||
Pi05CheckpointManifest,
|
Pi05CheckpointManifest,
|
||||||
Pi05PolicyModel,
|
Pi05PolicyModel,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.pi05_preprocess import (
|
||||||
|
_preprocess_image,
|
||||||
|
_resize_with_pad_image_tensor,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.vla.denoise_cuda_graph import (
|
from sglang.multimodal_gen.runtime.vla.denoise_cuda_graph import (
|
||||||
VLADenoiseGraphRunner,
|
VLADenoiseGraphRunner,
|
||||||
_CapturedDenoiseGraph,
|
_CapturedDenoiseGraph,
|
||||||
@@ -172,3 +178,64 @@ def test_siglip_attention_patch_uses_native_wrapper_once():
|
|||||||
|
|
||||||
assert isinstance(first, Pi05SiglipAttention)
|
assert isinstance(first, Pi05SiglipAttention)
|
||||||
assert layer.self_attn is first
|
assert layer.self_attn is first
|
||||||
|
|
||||||
|
|
||||||
|
def test_prefix_language_embedding_matches_openpi_scale():
|
||||||
|
image_embedding = torch.ones(1, 2, 8)
|
||||||
|
language_embedding = torch.full((1, 3, 8), 0.25)
|
||||||
|
model = SimpleNamespace(
|
||||||
|
paligemma_with_expert=SimpleNamespace(
|
||||||
|
embed_images=lambda images: [image_embedding],
|
||||||
|
embed_language_tokens=lambda tokens: language_embedding,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
embeddings, _, _ = Pi05CoreModel.embed_prefix(
|
||||||
|
model,
|
||||||
|
images=[torch.zeros(1, 3, 4, 4)],
|
||||||
|
image_masks=[torch.ones(1, dtype=torch.bool)],
|
||||||
|
tokens=torch.ones(1, 3, dtype=torch.long),
|
||||||
|
token_masks=torch.ones(1, 3, dtype=torch.bool),
|
||||||
|
)
|
||||||
|
|
||||||
|
torch.testing.assert_close(
|
||||||
|
embeddings[:, 2:],
|
||||||
|
language_embedding * (language_embedding.shape[-1] ** 0.5),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_uint8_resize_rounds_before_normalization():
|
||||||
|
image = torch.tensor([[[0.0, 1.0], [2.0, 3.0]]]) / 255.0
|
||||||
|
|
||||||
|
resized = _resize_with_pad_image_tensor(
|
||||||
|
image,
|
||||||
|
(3, 3),
|
||||||
|
round_to_uint8=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
expected = (
|
||||||
|
torch.round(
|
||||||
|
torch.nn.functional.interpolate(
|
||||||
|
image[None], size=(3, 3), mode="bilinear", align_corners=False
|
||||||
|
)[0]
|
||||||
|
* 255.0
|
||||||
|
)
|
||||||
|
/ 255.0
|
||||||
|
)
|
||||||
|
torch.testing.assert_close(resized, expected, rtol=0.0, atol=0.0)
|
||||||
|
|
||||||
|
|
||||||
|
def test_normalized_float_image_is_not_normalized_twice():
|
||||||
|
image = np.full((2, 4, 3), -0.5, dtype=np.float32)
|
||||||
|
|
||||||
|
preprocessed = _preprocess_image(image, (4, 4))
|
||||||
|
|
||||||
|
assert preprocessed.shape == (3, 4, 4)
|
||||||
|
torch.testing.assert_close(
|
||||||
|
preprocessed[:, 1:3],
|
||||||
|
torch.full((3, 2, 4), -0.5),
|
||||||
|
)
|
||||||
|
torch.testing.assert_close(
|
||||||
|
preprocessed[:, (0, 3)],
|
||||||
|
torch.full((3, 2, 4), -1.0),
|
||||||
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user