关节角 SinCos 编码
更新时间:2026-09-14
简介
关节角 sin/cos 双通道编码算子,把角度维映射成 [sin(θ), cos(θ)],消除角度的 2π 周期性歧义(VLA 训练的标准输入形式)。移植 GR00T-N1.5 gr00t/data/transform/state_action.py 的 StateActionSinCosTransform.apply。
功能描述
- 编码结果沿最后一维拼接成
[sin(sel), cos(sel)](与 GR00T 原版的拼接顺序一致),选中 n 维输出 2n 维;不是逐维交错的[sin, cos, sin, cos...] joint_dims=None时整行参与编码,此时行为与 GR00T 原版完全一致;给定joint_dims时只编码选中列,且按joint_dims给定的顺序取值keep_original=True时输出[原始整行, sin(sel), cos(sel)],长度为原维数 + 2 × 选中维数,便于回溯原值;默认只输出编码段- 输入按弧度处理,不做任何角度归一化(sin/cos 本身对 ±2π 不敏感)
joint_dims为空或有重复在构造时报错;下标越界在处理该行时报错- 行级容错:空行 / null / 非 1D / 下标越界,该行返回空列表,不打断整批
- 纯 numpy,CPU,无模型
算子参数
输入
| 输入 | 含义 |
|---|---|
| state | 每行一个 1D 向量(list[float],单位弧度) |
输出
| 输出 | 含义 |
|---|---|
| result | list |
参数
| 参数名称 | 类型 | 默认值 | 描述 |
|---|---|---|---|
| joint_dims | list[int] 或 None | None | 参与编码的列下标,按给定顺序取值;None 表示整行全部编码。须非空且不重复 |
| keep_original | bool | False | True 时在编码段前面拼上原始整行 |
注意事项
- 输出维度会翻倍(
keep_original=True时是原维数 + 2 × 选中维数),下游 schema 与训练配置要同步改。 - 只对角度维有物理意义。夹爪开合度、末端位置这类非角度维一起编码没有意义,用
joint_dims圈出角度维。 - 编码不可逆推回具体的角度多值分支:
atan2(sin, cos)只能还原到 (-π, π],原值超出这个区间时拿不回来。
调用示例
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.sincos_joint_encode import SinCosJointEncode
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 SinCosJointEncode,
29 construct_args={
30 "joint_dims": None,
31 "keep_original": False,
32 },
33 num_cpus=1,
34 concurrency=1,
35 batch_size=256,
36 )(col("state")),
37 )
38 ds.show()
评价此篇文章
