状态动作归一化
更新时间:2026-09-14
简介
state / action 归一化算子,按数据集统计量做 mean_std / min_max / q99 / scale / binary 五种归一化,并支持反归一化。移植 GR00T-N1.5 gr00t/data/transform/state_action.py 的 Normalizer,统计量可以直接吃 lerobot 数据集的 meta/stats.json。
功能描述
- 五种模式的 forward 与 GR00T 逐分支一致:
- ◦
mean_std:(x - mean) / std,std == 0的维透传原值 - ◦
min_max:2*(x - min)/(max - min) - 1,min == max的维置 0,不做截断 - ◦
q99:2*(x - q01)/(q99 - q01) - 1再截断到[-1, 1],q01 == q99的维透传后同样受截断 - ◦
scale:x / max(|min|, |max|),abs_max == 0的维置 0 - ◦
binary:(x > 0.5)→ {0, 1},不需要统计量 inverse=True走反变换;scale的逆在原版里缺失,此处按x * abs_max补齐,其余分支与原版语义一致- 统计量三种来源:构造参数
statistics(dict)优先,其次读statistics_path指向的stats.json,binary模式无需统计量。dict 既可以是内层单特征形式{"mean": [...], "std": [...]},也可以是外层映射{"observation.state": {...}, "action": {...}}——外层映射有多个特征时必须给feature_key,只有一个特征时自动选中 - 构造时校验本模式需要的统计键(
mean_std→mean/std,min_max/scale→min/max,q99→q01/q99),缺键立即报错;处理每行时校验向量维度与统计量维度一致 - 输出 struct 回填本模式 forward/inverse 实际用到的统计量,未用到的字段为 null,下游不必重读
stats.json就能反演 - 行级容错:空行 / null / 维度不符时
normalized为空列表且统计字段全 null(mode仍回填),不打断整批 - 纯 numpy,CPU,无模型
算子参数
输入
| 输入 | 含义 |
|---|---|
| state | 每行一个 1D 向量(list[float]),维度须与统计量一致(binary 模式无此约束) |
输出
| 输出 | 含义 |
|---|---|
| result.normalized | list |
| result.mode | string,本次使用的模式,取值 mean_std / min_max / q99 / scale / binary |
| result.mean | list |
| result.std | list |
| result.min | list |
| result.max | list |
| result.q01 | list |
| result.q99 | list |
参数
| 参数名称 | 类型 | 默认值 | 描述 |
|---|---|---|---|
| mode | str | "mean_std" | 归一化模式,取值 mean_std / min_max / q99 / scale / binary |
| statistics | dict 或 None | None | 直接给统计量,内层单特征 dict 或外层 feature 映射;优先于 statistics_path |
| statistics_path | str 或 None | None | stats.json 路径(lerobot / 数据集统计算子写出的格式),文件不存在报错 |
| feature_key | str 或 None | None | 外层映射里取哪个特征,如 observation.state;多特征时必填 |
| inverse | bool | False | True 走反归一化 |
注意事项
- 退化维不可逆:
mean_std的std == 0维 forward 透传原值,inverse 算的是x * 0 + mean,还原不回原值;min_max的min == max维 forward 置 0,inverse 得到min。做往返校验时要把退化维排除。 q99的截断也不可逆:落在[q01, q99]之外的值被夹到 ±1,反归一化只能回到分位边界。min_max与scale不做截断(GR00T 原语义),源值超出统计区间时输出会超出[-1, 1]。binary模式的 inverse 与 forward 相同(都是x > 0.5),只是幂等,不能还原原值。- 输入维度必须与统计量维度严格相等,否则该行返回空列表;不会按前缀对齐。
- 一个算子实例只带一个特征的统计量。
observation.state和action要各起一个实例(feature_key不同),反归一化再起一个inverse=True的实例。
调用示例
Python
1from __future__ import annotations
2
3import os
4
5import daft
6import pyarrow.parquet as pq
7from daft import col
8
9from daft.aihc.common.udf import aihc_udf
10from daft.aihc.functions.embodied.lerobot.normalize_state_action import NormalizeStateAction
11
12ROOT = "/path/to/lerobot_dataset"
13
14if __name__ == "__main__":
15 if os.getenv("DAFT_RUNNER", "native") == "ray":
16 import ray
17 ray.init(dashboard_host="0.0.0.0", ignore_reinit_error=True)
18 daft.set_runner_ray()
19 daft.set_execution_config(actor_udf_ready_timeout=6000, min_cpu_per_task=0)
20
21table = pq.read_table(
22 f"{ROOT}/data/chunk-000/episode_000000.parquet", columns=["observation.state"]
23 )
24 ds = daft.from_pydict({"state": table.column("observation.state").to_pylist()})
25 ds = ds.with_column(
26 "result",
27 aihc_udf(
28 NormalizeStateAction,
29 construct_args={
30 "mode": "mean_std",
31 "statistics_path": f"{ROOT}/meta/stats.json",
32 "feature_key": "observation.state",
33 "inverse": False,
34 },
35 num_cpus=1,
36 concurrency=1,
37 batch_size=256,
38 )(col("state")),
39 )
40 ds.show()
评价此篇文章
