Update Wall-X to 1.1.0 (#104)

This commit is contained in:
Starrick Liu
2026-06-15 11:40:00 +08:00
committed by GitHub
parent e23a586846
commit 72834e7de5
200 changed files with 33916 additions and 16771 deletions
+58 -362
View File
@@ -1,362 +1,58 @@
action_statistic_dof = {
"x2_normal": {
# water flowers
"follow_left_arm_joint_cur": {
"min": [-3.7121],
"delta": [7.6008],
},
"follow_right_arm_joint_cur": {
"min": [-3.6176],
"delta": [8.5015],
},
"follow_left_ee_cartesian_pos": {
"min": [-0.036, -0.3241, -0.1245],
"delta": [0.4389, 0.557, 0.479],
},
"follow_left_ee_rotation": {
"min": [-1.2373, -0.1929, -1.5182],
"delta": [2.2009, 1.5669, 2.0936],
},
"follow_left_gripper": {"min": [-0.1196], "delta": [4.5226]},
"follow_right_ee_cartesian_pos": {
"min": [-0.0326, -0.2273, -0.1377],
"delta": [0.4574, 0.5704, 0.4743],
},
"follow_right_ee_rotation": {
"min": [-1.2201, -0.2611, -0.7427],
"delta": [2.6623, 1.6622, 2.4186],
},
"follow_right_gripper": {"min": [-0.1208], "delta": [4.5261]},
"height": {"min": [-0.0001], "delta": [0.5051]},
"head_actions": {"min": [-1.5000, -1.4167], "delta": [2.5000, 1.8879]},
"base_velocity": {
"min": [-0.0359, -0.084, -0.0162],
"delta": [0.1539, 0.1848, 0.0322],
},
},
"DobbE": {
"follow_right_ee_cartesian_pos": {
"min": [-0.6107, -0.3272, -0.4282],
"delta": [1.2629, 1.5297, 0.8349],
},
"follow_right_ee_rotation": {
"min": [-1.7378, -1.4597, -1.8712],
"delta": [2.7031, 2.8182, 3.5921],
},
"follow_right_gripper": {"min": [0.0], "delta": [0.9983]},
},
"RH20T": {
"follow_right_ee_cartesian_pos": {
"min": [0.3646, -0.2722, 0.0066],
"delta": [0.3813, 0.5973, 0.3277],
},
"follow_right_ee_rotation": {
"min": [-1.8716, -0.4398, -3.1414],
"delta": [3.4145, 1.0225, 6.2828],
},
"follow_right_gripper": {"min": [0.0], "delta": [95.0]},
},
"agibotworld_alpha": {
"follow_left_ee_cartesian_pos": {
"min": [0.4954, 0.0166, 0.1729],
"delta": [0.3336, 0.5123, 0.9189],
},
"follow_left_ee_rotation": {
"min": [-3.1064, -1.2629, -3.1238],
"delta": [6.2127, 2.5923, 6.2496],
},
"follow_left_gripper": {"min": [34.6222], "delta": [86.1921]},
"follow_right_ee_cartesian_pos": {
"min": [0.4615, -0.5975, 0.1638],
"delta": [0.3823, 0.5577, 0.8873],
},
"follow_right_ee_rotation": {
"min": [-3.0891, -1.0739, -2.5091],
"delta": [6.1707, 2.3074, 3.8533],
},
"follow_right_gripper": {"min": [34.6222], "delta": [85.7635]},
"height": {"min": [0.0], "delta": [0.4535]},
"head_actions": {"min": [-0.1746, 0.0523], "delta": [0.2444, 0.4713]},
},
"austin_buds": {
"follow_right_ee_cartesian_pos": {
"min": [0.3496, -0.2855, 0.0105],
"delta": [0.3748, 0.492, 0.3116],
},
"follow_right_ee_rotation": {
"min": [-3.1405, -0.151, -0.0737],
"delta": [6.2813, 0.3218, 0.1536],
},
"follow_right_gripper": {"min": [0.0076], "delta": [0.0724]},
},
"austin_sailor": {
"follow_right_ee_cartesian_pos": {
"min": [0.387, -0.3165, 0.0244],
"delta": [0.2999, 0.5252, 0.2308],
},
"follow_right_ee_rotation": {
"min": [-3.1402, -0.1618, -1.5918],
"delta": [6.2804, 0.337, 2.9478],
},
"follow_right_gripper": {"min": [0.0005], "delta": [0.0773]},
},
"austin_sirius": {
"follow_right_ee_cartesian_pos": {
"min": [0.0, -0.1182, 0.0],
"delta": [0.5329, 0.3812, 0.2723],
},
"follow_right_ee_rotation": {
"min": [-3.1407, -0.1243, -1.7434],
"delta": [6.2823, 0.1975, 1.8073],
},
"follow_right_gripper": {"min": [0.0334], "delta": [0.046]},
},
"bc_z": {
"follow_right_ee_cartesian_pos": {
"min": [-0.3883, -0.1116, 0.6113],
"delta": [0.7199, 0.4288, 0.3709],
},
"follow_right_ee_rotation": {
"min": [-1.056, -1.0587, -2.6295],
"delta": [1.9142, 1.9455, 4.8064],
},
"follow_right_gripper": {"min": [0.2], "delta": [0.8]},
},
"berkeley_autolab_ur5": {
"follow_right_ee_cartesian_pos": {
"min": [0.3018, -0.2129, -0.1888],
"delta": [0.3121, 0.52, 0.3107],
},
"follow_right_ee_rotation": {
"min": [-3.1396, -0.2278, 1.1413],
"delta": [6.279, 0.454, 0.9841],
},
"follow_right_gripper": {"min": [0.0], "delta": [1.0]},
},
"berkeley_cable_routing": {
"follow_right_ee_cartesian_pos": {
"min": [0.4617, -0.28, 0.03],
"delta": [0.1838, 0.5665, 0.1272],
},
"follow_right_ee_rotation": {
"min": [-3.1413, -0.0299, -0.7665],
"delta": [6.2826, 0.0692, 3.322],
},
},
"berkeley_fanuc_manipulation": {
"follow_right_ee_cartesian_pos": {
"min": [0.3718, -0.4072, 0.0184],
"delta": [0.3483, 0.7201, 0.5229],
},
"follow_right_ee_rotation": {
"min": [-3.1399, -1.0166, -1.6988],
"delta": [6.2802, 1.4498, 3.2074],
},
"follow_right_gripper": {"min": [0.0], "delta": [1.0]},
},
"bridge_data_v2": {
"follow_right_ee_cartesian_pos": {
"min": [0.1498, -0.2178, -0.0901],
"delta": [0.3012, 0.469, 0.298],
},
"follow_right_ee_rotation": {
"min": [-0.3279, -0.6105, -1.0578],
"delta": [0.7378, 1.0353, 2.2552],
},
"follow_right_gripper": {"min": [0.0692], "delta": [0.9426]},
},
"dlr_edan_shared_control": {
"follow_right_ee_cartesian_pos": {
"min": [-0.8387, 0.1473, -0.3934],
"delta": [0.6579, 0.6025, 1.1566],
},
"follow_right_ee_rotation": {
"min": [-3.1217, -1.5197, -2.2516],
"delta": [6.2505, 1.5594, 4.2831],
},
"follow_right_gripper": {"min": [0.0], "delta": [1.0]},
},
"droid": {
"follow_right_ee_cartesian_pos": {
"min": [0.2667, -0.4396, -0.0472],
"delta": [0.5159, 0.8806, 0.8331],
},
"follow_right_ee_rotation": {
"min": [-3.1374, -1.216, -2.1741],
"delta": [6.2749, 2.1075, 4.2259],
},
"follow_right_gripper": {"min": [0.0], "delta": [0.9912]},
},
"fmb": {
"follow_right_ee_cartesian_pos": {
"min": [0.3554, -0.2844, 0.0354],
"delta": [0.336, 0.4961, 0.2943],
},
"follow_right_ee_rotation": {
"min": [-3.1404, -0.9302, -0.0599],
"delta": [6.2807, 1.724, 1.8284],
},
"follow_right_gripper": {"min": [0.0], "delta": [1.0]},
},
"fractal": {
"follow_right_ee_cartesian_pos": {
"min": [0.3242, -0.2836, 0.1405],
"delta": [0.5518, 0.4963, 0.9328],
},
"follow_right_ee_rotation": {
"min": [-3.1308, -0.2421, -2.9685],
"delta": [6.2609, 1.7343, 5.819],
},
"follow_right_gripper": {"min": [0.0], "delta": [1.0]},
},
"furniture_bench": {
"follow_right_ee_cartesian_pos": {
"min": [0.3691, -0.181, 0.0058],
"delta": [0.2962, 0.3582, 0.1775],
},
"follow_right_ee_rotation": {
"min": [-3.1394, -0.6121, -1.9958],
"delta": [6.2786, 1.6114, 3.7748],
},
"follow_right_gripper": {"min": [0.0035], "delta": [0.0762]},
},
"jaco_play": {
"follow_right_ee_cartesian_pos": {
"min": [-0.3787, -0.6294, 0.1682],
"delta": [0.5898, 0.3587, 0.2183],
},
"follow_right_ee_rotation": {
"min": [0.9792, -0.0668, -0.0498],
"delta": [0.0175, 0.1277, 0.0686],
},
"follow_right_gripper": {"min": [0.0791], "delta": [0.1033]},
},
"nyu_rot": {
"follow_right_ee_cartesian_pos": {
"min": [0.25, -1.0, -0.2],
"delta": [0.75, 2.0, 1.2],
},
"follow_right_ee_rotation": {
"min": [-3.1416, -3.1416, 6.2831],
"delta": [9.4248, 4.1416, 0.0],
},
"follow_right_gripper": {"min": [0.0], "delta": [1.0]},
},
"stanford_hydra": {
"follow_right_ee_cartesian_pos": {
"min": [0.2068, -0.274, 0.1317],
"delta": [0.4929, 0.4981, 0.4588],
},
"follow_right_ee_rotation": {
"min": [-3.1321, -0.7496, -3.0269],
"delta": [6.2658, 1.5176, 5.8261],
},
"follow_right_gripper": {"min": [0.0], "delta": [0.0811]},
},
"stanford_kuka_multimodal": {
"follow_right_ee_cartesian_pos": {
"min": [0.4781, -0.0659, 0.3424],
"delta": [0.0868, 0.0864, 0.1863],
},
"follow_right_ee_rotation": {
"min": [-3.136, -0.0521, -3.1413],
"delta": [6.2727, 0.1109, 6.2825],
},
"follow_right_gripper": {"min": [-0.4713], "delta": [0.9485]},
},
"taco_play": {
"follow_right_ee_cartesian_pos": {
"min": [0.1375, -0.4291, 0.2052],
"delta": [0.5327, 1.0237, 0.3913],
},
"follow_right_ee_rotation": {
"min": [-3.1391, -0.6946, -1.2808],
"delta": [6.2784, 0.8196, 3.0856],
},
"follow_right_gripper": {"min": [0.0001], "delta": [0.0806]},
},
"utaustin_mutex": {
"follow_right_ee_cartesian_pos": {
"min": [0.3213, -0.4734, 0.0141],
"delta": [0.2108, 0.8471, 0.5644],
},
"follow_right_ee_rotation": {
"min": [-3.1404, -0.2202, -1.5489],
"delta": [6.2805, 0.582, 1.9282],
},
"follow_right_gripper": {"min": [0.0019], "delta": [0.0738]},
},
"viola": {
"follow_right_ee_cartesian_pos": {
"min": [0.4011, -0.2521, 0.0103],
"delta": [0.2444, 0.4305, 0.4355],
},
"follow_right_ee_rotation": {
"min": [-3.1403, -0.2737, -1.8626],
"delta": [6.2804, 0.4901, 2.0618],
},
"follow_right_gripper": {"min": [0.0002], "delta": [0.0773]},
},
"kuka": {
"follow_right_ee_cartesian_pos": {
"min": [0.3914, -0.4901, 0.0175],
"delta": [0.3339, 0.8357, 0.9064],
},
"follow_right_ee_rotation": {
"min": [-3.1416, -0.9903, -3.1416],
"delta": [6.2832, 2.2421, 6.2832],
},
"follow_right_gripper": {
"min": [0.0000],
"delta": [1.0000],
},
},
"UMI-biarm": {
"follow_left_ee_cartesian_pos": {
"min": [-0.2917, -0.4926, 0.0063],
"delta": [0.9028, 0.8168, 0.3473],
},
"follow_left_ee_rotation": {
"min": [-2.5309, -1.5706, -1.3309],
"delta": [0.8758, 2.3315, 1.7076],
},
"follow_left_gripper": {
"min": [0.0029],
"delta": [0.0812],
},
"follow_right_ee_cartesian_pos": {
"min": [-0.0023, -0.5191, -0.0358],
"delta": [0.7474, 0.8668, 0.351],
},
"follow_right_ee_rotation": {
"min": [-2.4945, -2.0149, -0.8088],
"delta": [1.0941, 3.2628, 2.0018],
},
"follow_right_gripper": {
"min": [0.0019],
"delta": [0.0814],
},
},
"agibotworld_beta": {
"follow_left_ee_cartesian_pos": {
"min": [0.4954, 0.0166, 0.1729],
"delta": [0.3336, 0.5123, 0.9189],
},
"follow_left_ee_rotation": {
"min": [-3.1064, -1.2629, -3.1238],
"delta": [6.2127, 2.5923, 6.2496],
},
"follow_left_gripper": {"min": [34.6222], "delta": [86.1921]},
"follow_right_ee_cartesian_pos": {
"min": [0.4615, -0.5975, 0.1638],
"delta": [0.3823, 0.5577, 0.8873],
},
"follow_right_ee_rotation": {
"min": [-3.0891, -1.0739, -2.5091],
"delta": [6.1707, 2.3074, 3.8533],
},
"follow_right_gripper": {"min": [34.6222], "delta": [85.7635]},
"height": {"min": [0.0], "delta": [0.4535]},
"head_actions": {"min": [-0.1746, 0.0523], "delta": [0.2444, 0.4713]},
},
}
"""OSS-safe compatibility constants.
The internal repository keeps canonical robot, dataset, action-key, and action
statistics tables in ``x2robot_utils.constants``. Those tables are intentionally
not bundled in the public export. This shim keeps the public VLA code importable
without exposing private dataset catalogues or default normalization stats.
"""
from __future__ import annotations
from collections.abc import Iterable
ACTION_KEY_FULL_MAPPING: dict[str, str] = {}
_ACTION_KEY_FULL_MAPPING = ACTION_KEY_FULL_MAPPING
ACTION_DATASET_NAMES: tuple[str, ...] = ()
MULTIMODAL_DATASET_NAMES: tuple[str, ...] = ()
_ACTION_DATASET_NAMES = ACTION_DATASET_NAMES
_MULTIMODAL_DATASET_NAMES = MULTIMODAL_DATASET_NAMES
VIEW_SLOT_KEYS = ["view1", "view2", "view3"]
PHYSICAL_VIEW_IDS: dict[str, int] = {}
def get_view_slot_keys(num_views: int) -> list[str]:
"""Return canonical view slot names for ``num_views`` cameras."""
return [f"view{i + 1}" for i in range(max(0, int(num_views)))]
def is_multimodal_dataset_name(dataset_name: str | None) -> bool:
"""Return whether a dataset is marked multimodal in the public shim."""
return dataset_name in MULTIMODAL_DATASET_NAMES
def is_action_dataset_name(dataset_name: str | None) -> bool:
"""Return whether a dataset row should be treated as an action sample."""
return dataset_name is not None and not is_multimodal_dataset_name(dataset_name)
def iter_action_dataset_names(dataset_names: Iterable[str]) -> list[str]:
"""Filter a sequence down to action dataset names."""
return [name for name in dataset_names if is_action_dataset_name(name)]
__all__ = [
"ACTION_KEY_FULL_MAPPING",
"_ACTION_KEY_FULL_MAPPING",
"ACTION_DATASET_NAMES",
"MULTIMODAL_DATASET_NAMES",
"_ACTION_DATASET_NAMES",
"_MULTIMODAL_DATASET_NAMES",
"VIEW_SLOT_KEYS",
"PHYSICAL_VIEW_IDS",
"get_view_slot_keys",
"is_multimodal_dataset_name",
"is_action_dataset_name",
"iter_action_dataset_names",
]
+278
View File
@@ -0,0 +1,278 @@
import os
import torch
ENABLE_CUDA_GRAPH = os.environ.get("ENABLE_CUDA_GRAPH", "True").lower() == "true"
# When using a single master buffer + multi-bucket graph, max batch must be pre-allocated;
# otherwise resize/reallocation after capture invalidates addresses recorded by the graph.
CUDA_GRAPH_MAX_BS = int(os.environ.get("CUDA_GRAPH_MAX_BS", "128"))
_SHAPE_GUARD_FASTPATH = (
os.environ.get("CUDAGRAPH_DISABLE_SHAPE_GUARD_FASTPATH", "0") != "1"
)
class CUDAGraph_Wrapper:
"""
Bucketed CUDA Graph wrapper:
- When batch size changes, select/reuse a CUDAGraph by bucket_size
- For batch_size < bucket_size: copy inputs into static buffer [:B], return outputs [:B] after replay
"""
def __init__(
self,
model,
warm_up_times: int = 3,
enable: bool | None = None,
batch_size_key: str | None = None,
):
self.model = model
self.warm_up_times = warm_up_times
self._enable = ENABLE_CUDA_GRAPH if enable is None else bool(enable)
self._batch_size_key = batch_size_key or "suffix_inputs_embeds"
# graph cache: bucket_bs -> CUDAGraph
self.graphs: dict[int, torch.cuda.CUDAGraph] = {}
# Per bucket, store views into the master buffer (graph capture depends on these views' address/shape/stride)
self.static_inputs_map: dict[int, dict[str, torch.Tensor]] = {}
self.static_output_tensor: dict[int, torch.Tensor] = {}
self.graph_pool = None
# Master buffers (shared across all buckets)
self.graph_vars: dict[str, torch.Tensor] = {}
self._max_bs: int | None = None
# Record non-batch dimension shape signature (globally consistent) to avoid silent errors
self._shape_signature: dict[str, tuple] = {}
# Batch dimension index per input (default 0; e.g. suffix_position_ids batch is at dim=1)
self._batch_dim: dict[str, int] = {}
self._shape_guard_verified: dict[int, bool] = {}
def _get_batch_size(self, **kwargs) -> int:
key = self._batch_size_key
if key in kwargs:
v = kwargs[key]
if isinstance(v, torch.Tensor) and v.dim() >= 1:
return int(v.shape[0])
raise ValueError(f"CUDAGraph_Wrapper.forward requires {key} argument")
def _select_bucket_bs(self, bs: int) -> int:
"""
Bucket selection:
- 1, 2, 4, 8 use pow2 buckets
- >8 align to 16 (16, 32, 48, ...)
"""
if bs <= 1:
return 1
if bs <= 2:
return 2
if bs <= 4:
return 4
if bs <= 8:
return 8
return int(((bs + 15) // 16) * 16)
def _batch_dim_for_key(self, key: str, tensor: torch.Tensor) -> int:
# Minimal special-case handling for known inputs only
# suffix_position_ids: shape [3, B, T], batch at dim=1
if key == "suffix_position_ids" and tensor.dim() >= 2:
return 1
return 0
def _view_for_bucket(self, tensor: torch.Tensor, bucket_bs: int, batch_dim: int):
if tensor.dim() == 0:
return tensor
if batch_dim == 0:
return tensor[:bucket_bs]
if batch_dim == 1:
return tensor[:, :bucket_bs]
raise ValueError(f"Unsupported batch_dim={batch_dim} for cudagraph wrapper")
def _copy_and_pad(
self,
master: torch.Tensor,
tensor: torch.Tensor,
bs: int,
bucket_bs: int,
batch_dim: int,
):
if tensor.dim() == 0:
master.copy_(tensor)
return
if batch_dim == 0:
master[:bs].copy_(tensor)
if bs < bucket_bs:
master[bs:bucket_bs].zero_()
return
if batch_dim == 1:
master[:, :bs].copy_(tensor)
if bs < bucket_bs:
master[:, bs:bucket_bs].zero_()
return
raise ValueError(f"Unsupported batch_dim={batch_dim} for cudagraph wrapper")
def _allocate_master_tensor(
self, key: str, tensor: torch.Tensor, max_bs: int, bs: int
) -> torch.Tensor:
"""
Allocate a master buffer for one input (single large buffer):
- Extend only the batch dimension to max_bs; other dims unchanged
- Copy the current bs valid range first; zero the rest (avoid stale values on replay)
"""
if tensor.dim() == 0:
return tensor.clone()
batch_dim = self._batch_dim_for_key(key, tensor)
shape = list(tensor.shape)
if batch_dim >= len(shape):
raise ValueError(
f"Invalid batch_dim={batch_dim} for key={key}, tensor.shape={tuple(tensor.shape)}"
)
shape[batch_dim] = max_bs
buf = torch.zeros(tuple(shape), device=tensor.device, dtype=tensor.dtype)
# Initial copy: pad to current bucket (at least cover the valid bs range)
self._copy_and_pad(buf, tensor, bs=bs, bucket_bs=bs, batch_dim=batch_dim)
return buf
def _ensure_master_buffers(self, bs: int, **kwargs):
"""
Ensure master buffers are allocated once.
- max_bs from CUDA_GRAPH_MAX_BS; if unset, use the first bucket_bs as max_bs
and disallow exceeding it later (resize would invalidate captured graphs).
"""
if self._max_bs is not None:
return
bucket_bs = self._select_bucket_bs(bs)
max_bs = CUDA_GRAPH_MAX_BS if CUDA_GRAPH_MAX_BS > 0 else bucket_bs
self._max_bs = int(max_bs)
for k, v in kwargs.items():
if not isinstance(v, torch.Tensor):
raise TypeError(
f"CUDAGraph_Wrapper only supports Tensor kwargs, got {k}={type(v)}"
)
batch_dim = self._batch_dim_for_key(k, v) if v.dim() > 0 else 0
self._batch_dim[k] = batch_dim
# Record non-batch dimensions (globally consistent)
if v.dim() == 0:
self._shape_signature[k] = tuple()
else:
sig = list(v.shape)
sig.pop(batch_dim)
self._shape_signature[k] = tuple(sig)
self.graph_vars[k] = self._allocate_master_tensor(
k, v, max_bs=self._max_bs, bs=bs
)
def _initialize_bucket(self, bucket_bs: int, bs: int, **kwargs):
assert self._max_bs is not None
self._shape_guard_verified[bucket_bs] = False
if bucket_bs > self._max_bs:
raise ValueError(
f"CUDAGraph bucket_bs({bucket_bs}) exceeds master max_bs({self._max_bs}). "
f"Set CUDA_GRAPH_MAX_BS >= {bucket_bs}."
)
print(f"[CUDAGraph] Initializing bucket_bs={bucket_bs} (current bs={bs}) ...")
# Build view map for this bucket (fixed address/shape/stride)
static_map: dict[str, torch.Tensor] = {}
for k, v in kwargs.items():
# Shape guard: non-batch dims must match
expected = self._shape_signature.get(k)
if v.dim() == 0:
got = tuple()
else:
batch_dim = self._batch_dim.get(k, self._batch_dim_for_key(k, v))
got_list = list(v.shape)
got_list.pop(batch_dim)
got = tuple(got_list)
if expected is not None and got != expected:
raise ValueError(
f"CUDAGraph bucket init: non-batch dim changed: key={k}, expected={expected}, got={got}"
)
master = self.graph_vars[k]
batch_dim = self._batch_dim.get(k, 0)
static_map[k] = self._view_for_bucket(
master, bucket_bs=bucket_bs, batch_dim=batch_dim
)
# Warmup (stabilize kernels/caches)
out = None
for _ in range(self.warm_up_times):
out = self.forward_naive(**static_map)
assert isinstance(out, torch.Tensor)
# Output master buffer (allocate once)
if "_outputs" not in self.graph_vars:
out_shape = (self._max_bs,) + tuple(out.shape[1:])
self.graph_vars["_outputs"] = torch.empty(
out_shape, device=out.device, dtype=out.dtype
)
self.static_inputs_map[bucket_bs] = static_map
self.static_output_tensor[bucket_bs] = self.graph_vars["_outputs"][:bucket_bs]
# capture
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph, pool=self.graph_pool):
self.static_output_tensor[bucket_bs].copy_(self.forward_naive(**static_map))
if self.graph_pool is None:
self.graph_pool = graph.pool()
self.graphs[bucket_bs] = graph
torch.cuda.synchronize()
print(f"[CUDAGraph] Bucket {bucket_bs} captured.")
def forward_naive(self, **kwargs):
return self.model(**kwargs)
def forward(self, **kwargs):
if not self._enable:
return self.forward_naive(**kwargs)
bs = self._get_batch_size(**kwargs)
bucket_bs = self._select_bucket_bs(bs)
# Initialize master buffers (first call)
self._ensure_master_buffers(bs=bs, **kwargs)
assert self._max_bs is not None
if bucket_bs > self._max_bs:
# Do not exceed master max_bs at runtime (resize would invalidate captured graphs)
return self.forward_naive(**kwargs)
# lazy capture for this bucket
if bucket_bs not in self.graphs:
self._initialize_bucket(bucket_bs=bucket_bs, bs=bs, **kwargs)
# Shape guard: non-batch dim change is unsafe for cudagraph; fallback
# fast-path: skip guard if already verified for this bucket (stable-state optimization)
if not (
_SHAPE_GUARD_FASTPATH and self._shape_guard_verified.get(bucket_bs, False)
):
for k, v in kwargs.items():
if not isinstance(v, torch.Tensor):
raise TypeError(
f"CUDAGraph_Wrapper only supports Tensor kwargs, got {k}={type(v)}"
)
expected = self._shape_signature.get(k)
if v.dim() == 0:
got = tuple()
else:
batch_dim = self._batch_dim.get(k, self._batch_dim_for_key(k, v))
got_list = list(v.shape)
got_list.pop(batch_dim)
got = tuple(got_list)
if expected is not None and got != expected:
return self.forward_naive(**kwargs)
# Mark verified for fast-path on subsequent calls
self._shape_guard_verified[bucket_bs] = True
# Copy inputs into master buffer (per batch_dim, zero padding)
for k, tensor in kwargs.items():
master = self.graph_vars[k]
batch_dim = self._batch_dim.get(k, 0)
self._copy_and_pad(
master, tensor, bs=bs, bucket_bs=bucket_bs, batch_dim=batch_dim
)
graph = self.graphs[bucket_bs]
graph.replay()
return self.graph_vars["_outputs"][:bs]
+95
View File
@@ -0,0 +1,95 @@
"""Rank-aware text logger for distributed training."""
from __future__ import annotations
import logging
import os
import sys
from pathlib import Path
from typing import Optional
from torch.distributed import get_rank, is_initialized
class DistributedLogger:
def __init__(
self,
name: str = "wallx",
save_path: Optional[str] = None,
level: int = logging.INFO,
):
if is_initialized():
self._rank = get_rank()
else:
self._rank = int(os.environ.get("RANK", 0))
logging.warning(
"DistributedLogger created before init_process_group; "
"falling back to RANK env var (rank=%d).",
self._rank,
)
logger = logging.getLogger(f"{name}.rank{self._rank}")
logger.setLevel(level)
logger.propagate = False
for h in logger.handlers[:]:
h.close()
logger.removeHandler(h)
fmt = logging.Formatter(
f"%(asctime)s - [rank{self._rank}] - %(levelname)s - %(message)s"
)
# All ranks write a file log (if a save_path was given).
if save_path:
log_dir = Path(save_path) / "logs"
log_dir.mkdir(parents=True, exist_ok=True)
file_handler = logging.FileHandler(log_dir / f"rank_{self._rank}.log")
file_handler.setFormatter(fmt)
logger.addHandler(file_handler)
# Only rank 0 writes to stdout.
if self._rank == 0:
stream_handler = logging.StreamHandler(sys.stdout)
stream_handler.setFormatter(fmt)
logger.addHandler(stream_handler)
self._logger = logger
# Thin pass-throughs - callers use standard logging verbs.
def info(self, msg, *args, **kwargs):
self._logger.info(msg, *args, **kwargs)
def warning(self, msg, *args, **kwargs):
self._logger.warning(msg, *args, **kwargs)
def error(self, msg, *args, **kwargs):
self._logger.error(msg, *args, **kwargs)
def debug(self, msg, *args, **kwargs):
self._logger.debug(msg, *args, **kwargs)
@property
def rank(self) -> int:
return self._rank
# --- backward-compatible shim ----------------------------------------
# Some legacy call sites still invoke `.log(msg, level=...,
# main_process_only=...)`. Forward those to the standard logger so we
# don't break them during the migration; new code should use .info /
# .warning / .error / .debug directly.
def log(
self,
message,
level: int = logging.INFO,
main_process_only: bool = False,
):
if main_process_only and self._rank != 0:
return
self._logger.log(level, message)
# Pytorch's ``accelerate``-style fallback for code that still passes an
# accelerator object to the old constructor; accept and ignore it.
# Older codepaths can be migrated incrementally.
@classmethod
def legacy(cls, name: str, level: int = logging.INFO, accelerator=None):
del accelerator # ignored
return cls(name=name, save_path=None, level=level)
+95
View File
@@ -0,0 +1,95 @@
from typing import List
import torch
def get_action_accuracy(
gt: torch.FloatTensor, # [Batch_Size, Horizon, Action_Dim]
pred: torch.FloatTensor,
thresholds: List[float] = [0.1, 0.2],
) -> torch.FloatTensor:
device = gt.device
diff = torch.abs(gt - pred).reshape(-1, gt.shape[-1])
# get the percentage of diff lower than threshold for all action dimensions
accuracies = torch.zeros(len(thresholds), device=device)
for idx, threshold in enumerate(thresholds):
accuracy = torch.mean(
(torch.mean((diff < threshold).float(), dim=1) >= 1.0).float()
)
accuracies[idx] = accuracy
return accuracies
def dtw_distance(seq1, seq2):
"""
Compute the Dynamic Time Warping distance between two sequences.
``seq1`` and ``seq2`` must have shape ``(T, D)``, where ``T`` is the
number of time steps and ``D`` is the feature dimension.
"""
n, m = seq1.shape[0], seq2.shape[0]
seq1_d, seq2_d = seq1.double(), seq2.double()
dtw_matrix = torch.full(
(n + 1, m + 1), float("inf"), dtype=torch.float64, device=seq1.device
)
dtw_matrix[0, 0] = 0
for i in range(1, n + 1):
for j in range(1, m + 1):
cost = torch.dist(seq1_d[i - 1], seq2_d[j - 1])
dtw_matrix[i, j] = cost + torch.min(
torch.stack(
[
dtw_matrix[i - 1, j],
dtw_matrix[i, j - 1],
dtw_matrix[i - 1, j - 1],
]
)
)
return dtw_matrix[n, m]
def frechet_distance(seq1, seq2):
"""
Compute the discrete Frechet distance between two sequences.
``seq1`` and ``seq2`` must have shape ``(T, D)``, where ``T`` is the
number of time steps and ``D`` is the feature dimension.
"""
n, m = seq1.shape[0], seq2.shape[0]
seq1_d, seq2_d = seq1.double(), seq2.double()
frechet_matrix = torch.full(
(n, m), float("inf"), dtype=torch.float64, device=seq1.device
)
frechet_matrix[0, 0] = torch.dist(seq1_d[0], seq2_d[0])
for j in range(1, m):
frechet_matrix[0, j] = torch.max(
frechet_matrix[0, j - 1], torch.dist(seq1_d[0], seq2_d[j])
)
for i in range(1, n):
frechet_matrix[i, 0] = torch.max(
frechet_matrix[i - 1, 0], torch.dist(seq1_d[i], seq2_d[0])
)
for i in range(1, n):
for j in range(1, m):
frechet_matrix[i, j] = torch.max(
torch.min(
torch.stack(
[
frechet_matrix[i - 1, j],
frechet_matrix[i, j - 1],
frechet_matrix[i - 1, j - 1],
]
)
),
torch.dist(seq1_d[i], seq2_d[j]),
)
return frechet_matrix[n - 1, m - 1]
+39 -178
View File
@@ -1,12 +1,15 @@
import logging
import os
import time
from torch.cuda import nvtx
from abc import ABC, abstractmethod
from contextlib import nullcontext
from functools import wraps
from typing import List
import torch
from functools import wraps
from contextlib import nullcontext
import os
from torch.cuda import nvtx
logger = logging.getLogger(__name__)
ENABLE_PERFORMANCE_TIMING = (
os.environ.get("ENABLE_PERFORMANCE_TIMING", "True").lower() == "true"
@@ -32,23 +35,14 @@ class ScopeTimerContext:
torch.cuda.synchronize()
end_time = time.perf_counter()
cost_ms = (end_time - self.start_time) * 1e3
print(f"\033[92m{self.msg} took {cost_ms:.3f} ms to execute\033[0m")
logger.info("%s took %.3f ms to execute", self.msg, cost_ms)
ScopeTimer = ScopeTimerContext if ENABLE_PERFORMANCE_TIMING else nullcontext
def timer(func, msg=None):
"""
Decorator to measure function execution time.
Args:
func: Function to be timed
Returns:
Wrapped function with timing functionality
"""
"""Decorator to measure function execution time."""
if msg is None:
msg = func.__name__
else:
@@ -58,44 +52,36 @@ def timer(func, msg=None):
def wrapper(*args, **kwargs):
with ScopeTimer(msg):
result = func(*args, **kwargs)
return result
return wrapper
# Helper functions to check for distributed environment
def _is_distributed():
"""Checks if the current environment is set up for distributed training."""
return torch.distributed.is_available() and torch.distributed.is_initialized()
def _get_world_size():
"""Safely retrieves the world size (number of processes)."""
if _is_distributed():
return torch.distributed.get_world_size()
return 1
def _get_rank():
"""Safely retrieves the rank of the current process."""
if _is_distributed():
return torch.distributed.get_rank()
return 0
def _barrier(group=None):
"""Safely executes a distributed barrier to synchronize processes."""
if _is_distributed():
torch.distributed.barrier(group=group)
# Dynamically set the all_gather function
if torch.distributed.is_available():
try:
dist_all_gather_func = torch.distributed.all_gather_into_tensor
except AttributeError:
# Fallback to standard all_gather if all_gather_into_tensor is missing
dist_all_gather_func = torch.distributed.all_gather
else:
dist_all_gather_func = None
@@ -109,51 +95,35 @@ class TimerBase(ABC):
@abstractmethod
def start(self, barrier=False):
"""Start the timer.
Args:
barrier (bool, optional): Synchronizes ranks before starting. Defaults to False.
"""
"""Start the timer, optionally syncing all ranks with a barrier first."""
pass
@abstractmethod
def stop(self, barrier=False):
"""Stop the timer.
Args:
barrier (bool, optional): Synchronizes ranks before stopping. Defaults to False.
"""
"""Stop the timer, optionally syncing all ranks with a barrier first."""
pass
@abstractmethod
def reset(self):
"""Reset timer."""
"""Reset accumulated elapsed time to zero."""
pass
@abstractmethod
def elapsed(self, reset=True, barrier=False):
"""Calculates the elapsed time and restarts timer.
Args:
reset (bool, optional): Resets timer before restarting. Defaults to True.
barrier (bool, optional): Synchronizes ranks before stopping. Defaults to False.
Returns:
float: Elapsed time.
"""
"""Return accumulated elapsed time in seconds; reset if reset=True."""
pass
class DummyTimer(TimerBase):
"""Dummy Timer."""
"""Dummy Timer - no-op placeholder used when log level exceeds threshold."""
def __init__(self):
super().__init__("dummy timer")
def start(self, barrier=False, nvtx_push=False):
def start(self, barrier=False, nvtx_push=False, sync=False, **kwargs):
return
def stop(self, barrier=False, nvtx_pop=False):
def stop(self, barrier=False, sync=False, **kwargs):
return
def reset(self):
@@ -166,9 +136,6 @@ class DummyTimer(TimerBase):
)
def active_time(self):
"""Returns the cumulative duration the timer has been active.
Note: Not supported for DummyTimer.
"""
raise Exception(
"active timer should not be used to calculate elapsed time, "
"check if timer's log_level <= self._log_level."
@@ -188,34 +155,18 @@ class Timer(TimerBase):
"""
def __init__(self, name):
"""Initialize Timer.
Args:
name (str): Name of the timer.
"""
super().__init__(name)
self._elapsed = 0.0
self._active_time = 0.0
self._started = False
# Note that None will default to the global process group
self._barrier_group = None
self._start_time = time.time()
self.nvtx = False
def set_barrier_group(self, barrier_group):
"""Sets barrier group.
Args:
barrier_group (ProcessGroup): Torch ProcessGroup for barrier.
"""
self._barrier_group = barrier_group
def start(self, barrier=False, nvtx_push=False, sync=False):
"""Start the timer.
Args:
barrier (bool, optional): Synchronizes ranks before starting. Defaults to False.
"""
assert not self._started, "timer has already been started"
if barrier:
_barrier(group=self._barrier_group)
@@ -228,11 +179,6 @@ class Timer(TimerBase):
self.nvtx = True
def stop(self, barrier=False, sync=False):
"""Stop the timer.
Args:
barrier (bool, optional): Synchronizes ranks before stopping. Defaults to False.
"""
if self.nvtx:
nvtx.range_pop()
assert self._started, "timer is not started"
@@ -246,37 +192,21 @@ class Timer(TimerBase):
self._started = False
def reset(self):
"""Reset timer."""
# Don't reset _active_time
self._elapsed = 0.0
self._started = False
def elapsed(self, reset=True, barrier=False):
"""Calculates the elapsed time and restarts timer.
Args:
reset (bool, optional): Resets timer before restarting. Defaults to True.
barrier (bool, optional): Synchronizes ranks before stopping. Defaults to False.
Returns:
float: Elapsed time.
"""
_started = self._started
# If the timing in progress, end it first.
if self._started:
self.stop(barrier=barrier)
# Get the elapsed time.
_elapsed = self._elapsed
# Reset the elapsed time
if reset:
self.reset()
# If timing was in progress, set it back.
if _started:
self.start(barrier=barrier)
return _elapsed
def active_time(self):
"""Calculates the cumulative duration for which the timer has been active"""
return self._active_time
@@ -284,13 +214,6 @@ class Timers:
"""Class for a group of Timers."""
def __init__(self, log_level, log_option):
"""Initialize group of timers.
Args:
log_level (int): Log level to control what timers are enabled.
log_option (str): Setting for logging statistics over ranks for all the timers.
Allowed: ['max', 'minmax', 'all'].
"""
self._log_level = log_level
allowed_log_options = set(["max", "minmax", "all"])
assert (
@@ -305,9 +228,6 @@ class Timers:
self._max_log_level = 2
def __call__(self, name, log_level=None):
"""Call timer with name and log level."""
# If the timer has already been set, then check if the log-level
# is provided, it matches the one that the timer was created with.
if name in self._timers:
if log_level is not None:
assert log_level == self._log_levels[name], (
@@ -317,8 +237,6 @@ class Timers:
)
)
return self._timers[name]
# If timer does not exist and no log level is provided,
# set it to the max log level which is 2.
if log_level is None:
log_level = self._max_log_level
assert (
@@ -326,38 +244,19 @@ class Timers:
), "log level {} is larger than max supported log level {}".format(
log_level, self._max_log_level
)
# Now if the input log level is larger than the one set for
# the timers class, just ignore it and return a dummy timer.
if log_level > self._log_level:
return self._dummy_timer
# Otherwise, initalize the timer and set the level.
self._timers[name] = Timer(name)
self._log_levels[name] = log_level
return self._timers[name]
def _get_elapsed_time_all_ranks(self, names, reset, barrier):
"""Returns elapsed times of timers in names.
For single-node/single-GPU cases, directly returns the time for the current rank.
For distributed cases, maintains the existing all_gather logic.
Args:
names (List[str]): list of timer names
reset (bool): reset the timer after recording the elapsed time
barrier (bool): if set, do a global barrier before time measurements
Returns:
torch.tensor: Tensor of size [world_size, len(names)] with times in float.
"""
# First make sure all the callers are in sync.
if barrier:
_barrier()
world_size = _get_world_size()
rank = _get_rank()
# Create device tensor
if torch.cuda.is_available():
device = torch.cuda.current_device()
else:
@@ -367,33 +266,26 @@ class Timers:
(world_size, len(names)), dtype=torch.float, device=device
)
# Fill timing data for the current rank
for i, name in enumerate(names):
if name in self._timers:
rank_name_to_time[rank, i] = self._timers[name].elapsed(reset=reset)
# Return directly for single-node; perform all_gather for distributed setup
if world_size > 1 and _is_distributed() and dist_all_gather_func is not None:
try:
dist_all_gather_func(
rank_name_to_time.view(-1), rank_name_to_time[rank, :].view(-1)
)
except Exception as e:
# If all_gather fails, print a warning and proceed with single rank timing
print(f"Warning: all_gather failed: {e}. Using single rank timing.")
logger.warning("all_gather failed: %s. Using single rank timing.", e)
return rank_name_to_time
def _get_global_min_max_time(self, names, reset, barrier, normalizer):
"""Report only min and max times across all ranks."""
rank_name_to_time = self._get_elapsed_time_all_ranks(names, reset, barrier)
name_to_min_max_time = {}
for i, name in enumerate(names):
rank_to_time = rank_name_to_time[:, i]
# filter out the ones we did not have any timings for
rank_to_time = rank_to_time[rank_to_time > 0.0]
# If the timer exists:
if rank_to_time.numel() > 0:
name_to_min_max_time[name] = (
rank_to_time.min().item() / normalizer,
@@ -404,7 +296,6 @@ class Timers:
def _get_global_min_max_time_string(
self, names, reset, barrier, normalizer, max_only
):
"""Report strings for max/minmax times across all ranks."""
name_to_min_max_time = self._get_global_min_max_time(
names, reset, barrier, normalizer
)
@@ -413,17 +304,13 @@ class Timers:
world_size = _get_world_size()
if world_size == 1:
# Simplified output for single-node setup
output_string = "time (ms):"
for name in name_to_min_max_time:
_, max_time = name_to_min_max_time[
name
] # min and max are identical for a single rank
_, max_time = name_to_min_max_time[name]
output_string += "\n {}: {:.2f}".format(
(name + " ").ljust(48, "."), max_time
)
else:
# Maintain original output format for multi-node setup
if max_only:
output_string = "max time across ranks (ms):"
else:
@@ -441,7 +328,6 @@ class Timers:
return output_string
def _get_all_ranks_time_string(self, names, reset, barrier, normalizer):
"""Report times across all ranks."""
rank_name_to_time = self._get_elapsed_time_all_ranks(names, reset, barrier)
world_size = _get_world_size()
@@ -474,32 +360,20 @@ class Timers:
reset: bool = True,
barrier: bool = False,
):
"""Returns the output string with logged timer values according to configured options.
"""Return a formatted timing string for the given timer names.
Args:
names (List[str]): Names of the timers to log. If None, all registered timers are
fetched. Defaults to None.
normalizer (float, optional): Normalizes the timer values by the factor.
Defaults to 1.0.
reset (bool, optional): Whether to reset timer values after logging. Defaults to True.
barrier (bool, optional): Whether to do a global barrier before time measurments.
Defaults to False.
Raises:
Exception: Raises if log option is invalid.
Returns:
str: Formatted string with the timer values.
names: Timers to include; defaults to all registered timers.
normalizer: Divide raw seconds by this value (e.g. 1000 for ms output).
reset: Reset each timer after reading its elapsed time.
barrier: Synchronize across ranks before gathering times.
"""
if names is None: # get all registered timers
if names is None:
names = list(self._timers.keys())
assert normalizer > 0.0
if self._log_option in ["max", "minmax"]:
max_only = False
if self._log_option == "max":
max_only = True
max_only = self._log_option == "max"
output_string = self._get_global_min_max_time_string(
names, reset, barrier, normalizer / 1000.0, max_only
)
@@ -519,30 +393,23 @@ class Timers:
reset: bool = True,
barrier: bool = False,
):
"""logs the timers passed in names to stdout. Example usage is to log average per step
value for timer 'foo', this function can be called with normalizer factor set to logging
interval.
"""Print timing results for the given names to stdout on one rank.
Args:
names (List[str]): Names of the timers to log.
rank (int, optional): logs the timers to a specific rank. If set to None, logs to the
last rank. Defaults to None.
normalizer (float, optional): Normalizes the timer values by the factor.
Defaults to 1.0.
reset (bool, optional): Whether to reset timer values after logging. Defaults to True.
barrier (bool, optional): Whether to do a global barrier before time measurments.
Defaults to False.
names: Timer names to log.
rank: Rank that prints; defaults to the last rank (world_size - 1).
normalizer: Divide raw seconds by this value before printing.
reset: Reset each timer after reading.
barrier: Synchronize across ranks first.
"""
output_string = self.get_all_timers_string(names, normalizer, reset, barrier)
# If no input rank is provided, log on last rank.
world_size = _get_world_size()
current_rank = _get_rank()
if rank is None:
rank = world_size - 1
if rank == current_rank and output_string is not None:
print(output_string, flush=True)
logger.info("%s", output_string)
def write(
self,
@@ -553,22 +420,16 @@ class Timers:
reset: bool = True,
barrier: bool = False,
):
"""Write timers to a tensorboard writer.
Note that we only report maximum time across ranks to tensorboard.
"""Write per-timer max times as TensorBoard scalars.
Args:
names (List[str]): Names of the timers to log.
writer (SummaryWriter): Tensorboard SummaryWriter object
iteration (int): Current iteration.
normalizer (float, optional): Normalizes the timer values by the factor.
Defaults to 1.0.
reset (bool, optional): Whether to reset timer values after logging. Defaults to True.
barrier (bool, optional): Whether to do a global barrier before time measurments.
Defaults to False.
names: Timer names to write.
writer: TensorBoard SummaryWriter instance.
iteration: Global step value for the scalar.
normalizer: Divide raw seconds by this value.
reset: Reset each timer after reading.
barrier: Synchronize across ranks first.
"""
# currently when using add_scalars,
# torch.utils.add_scalars makes each timer its own run, which
# polutes the runs list, so we just add each as a scalar
assert normalizer > 0.0
name_to_min_max_time = self._get_global_min_max_time(
names, reset, barrier, normalizer