Update Wall-X to 1.1.0 (#104)
This commit is contained in:
+58
-362
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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]
|
||||
@@ -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)
|
||||
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user