ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

ACT 训练阶段完整 Mermaid

ACT 训练阶段完整 Mermaid 可以。下面把前面讨论的Backbone、CVAE、ACT Transformer、Action Chunk、Loss 和梯度路径全部重新合成一张训练阶段总图。为了保证论文 / 代码真实性(Paper / Code Fidelity),图中的网络拓扑以论文方法定义为主,具体实现细节以官方仓库为准:官方实现使用单个共享的 ResNet18 对不同相机逐个编码,视觉特征沿宽度维拼接;CVAE Encoder 只接收当前qtq_tqt​与专家动作块,不接收图像;主 ACT Transformer Encoder 的输入则是[z token, qpos token, visual tokens]。turn840745view3标准 ACT 配置中,状态/动作维度为 14,latent dimension 为 32,hidden dimension 为 512,CVAE Transformer Encoder 与主 ACT Transformer Encoder 均为 4 层,ACT Transformer Decoder 为 7 层,attention heads 为 8;官方示例使用k=100k=100k=100、β=10\beta=10β=10。turn936103view0turn189676view1ACT 训练阶段完整 Mermaidflowchart TD subgraph S0["阶段 S0:原始训练数据(Raw Training Data)"] direction LR D_IMG_RAW[/"数据模块:四视角 RGB 图像(Multi-view RGB Images I_t)br/Shape: [B, 4, 480, 640, 3]"/] D_QPOS_RAW[/"数据模块:当前双臂关节位置(Current Joint Positions q_t)br/Shape: [B, 14]"/] D_ACT_RAW[/"数据模块:专家目标动作序列(Expert Target Actions)br/Shape: [B, T, 14]"/] D_PAD_RAW[/"数据模块:动作填充掩码(Action Padding Mask)br/Shape: [B, T]"/] end subgraph S1["阶段 S1:数据预处理(Data Preprocessing)"] direction LR O_IMG_PRE(["操作模块:图像预处理(Image Preprocessing)br/参数:Channel Last→First,/255,ImageNet Normalization"]) D_IMG[/"数据模块:标准化多视角图像(Normalized Multi-view Images)br/Shape: [B, 4, 3, 480, 640]"/] O_QPOS_NORM(["操作模块:关节状态标准化(QPos Normalization)br/参数:Dataset Mean / Std"]) D_QPOS[/"数据模块:标准化关节位置(Normalized q_t)br/Shape: [B, 14]"/] O_ACT_CHUNK(["操作模块:动作标准化与分块(Action Normalization and Chunking)br/参数:截取前 k 步,Dataset Mean / Std"]) D_ACT[/"数据模块:标准化专家动作块(Normalized Target Action Chunk A_t)br/Shape: [B, k, 14]"/] O_PAD_CHUNK(["操作模块:填充掩码截取(Padding Mask Truncation)br/参数:保留前 k 步"]) D_PAD[/"数据模块:动作块填充掩码(Action Chunk Padding Mask)br/Shape: [B, k]"/] D_IMG_RAW -- O_IMG_PRE -- D_IMG D_QPOS_RAW -- O_QPOS_NORM -- D_QPOS D_ACT_RAW -- O_ACT_CHUNK -- D_ACT D_PAD_RAW -- O_PAD_CHUNK -- D_PAD end subgraph S2["阶段 S2:CVAE 后验分支(CVAE Posterior Branch)"] direction TB subgraph S2A["CVAE 输入嵌入(CVAE Input Embedding)"] direction LR M_CVAE_EMB["模型模块:🔥 CVAE 输入嵌入(CVAE Input Embedding)br/结构:Learned [CLS] Embedding + q_t Linear 14→512 + Action Linear 14→512"] D_CLS[/"数据模块:学习式分类 Token(Learned [CLS] Token)br/Shape: [B, 1, 512]"/] D_QPOS_CVAE[/"数据模块:CVAE 关节 Token(CVAE Joint Token)br/Shape: [B, 1, 512]"/] D_ACT_EMB[/"数据模块:CVAE 动作 Token(CVAE Action Tokens)br/Shape: [B, k, 512]"/] end D_QPOS -- M_CVAE_EMB D_ACT -- M_CVAE_EMB M_CVAE_EMB -- D_CLS M_CVAE_EMB -- D_QPOS_CVAE M_CVAE_EMB -- D_ACT_EMB O_CVAE_CAT(["操作模块:CVAE Token 拼接(CVAE Token Concatenation)br/顺序:[CLS] → q_t → Action Tokens"]) D_CVAE_SEQ_BF[/"数据模块:CVAE Token 序列(CVAE Token Sequence)br/Shape: [B, k+2, 512]"/] D_CLS -- O_CVAE_CAT D_QPOS_CVAE -- O_CVAE_CAT D_ACT_EMB -- O_CVAE_CAT O_CVAE_CAT -- D_CVAE_SEQ_BF O_CVAE_PERMUTE(["操作模块:维度置换(Sequence-First Permutation)"]) D_CVAE_SEQ[/"数据模块:CVAE Encoder 输入(CVAE Encoder Input)br/Shape: [k+2, B, 512]"/] D_CVAE_SEQ_BF -- O_CVAE_PERMUTE -- D_CVAE_SEQ D_CVAE_POS[/"数据模块:固定正弦位置编码(Fixed Sinusoidal Positional Embedding)br/Shape: [k+2, 1, 512]"/] O_PAD_EXT(["操作模块:扩展填充掩码(Extend Padding Mask)br/参数:[CLS] 与 q_t 对应位置前置 False"]) D_CVAE_PAD[/"数据模块:CVAE Encoder 填充掩码(CVAE Encoder Padding Mask)br/Shape: [B, k+2]"/] D_PAD -- O_PAD_EXT -- D_CVAE_PAD M_CVAE_ENC["模型模块:🔥 CVAE Transformer 编码器(CVAE Transformer Encoder q_φ)br/结构:Transformer Encoder × 4,d_model = 512,8 Heads,FFN = 3200,Dropout = 0.1"] D_CVAE_H[/"数据模块:CVAE 编码器隐藏序列(CVAE Encoder Hidden Sequence)br/Shape: [k+2, B, 512]"/] D_CVAE_SEQ -- M_CVAE_ENC D_CVAE_POS -. "PosEmb" .- M_CVAE_ENC D_CVAE_PAD -. "Padding Mask" .- M_CVAE_ENC M_CVAE_ENC -- D_CVAE_H O_CLS_SELECT(["操作模块:提取 [CLS] 表示(Select [CLS] Representation)"]) D_CLS_H[/"数据模块:[CLS] 隐藏表示([CLS] Hidden Representation)br/Shape: [B, 512]"/] D_CVAE_H -- O_CLS_SELECT -- D_CLS_H M_LATENT_HEAD["模型模块:🔥 潜变量参数头(Latent Parameter Head)br/结构:Linear 512→64"] D_LATENT_INFO[/"数据模块:潜变量分布参数(Latent Distribution Parameters)br/Shape: [B, 64]"/] D_CLS_H -- M_LATENT_HEAD -- D_LATENT_INFO O_LATENT_SPLIT(["操作模块:潜变量参数拆分(Latent Parameter Split)"]) D_MU[/"数据模块:潜变量均值(Latent Mean μ)br/Shape: [B, 32]"/] D_LOGVAR[/"数据模块:潜变量对数方差(Latent Log-Variance logvar)br/Shape: [B, 32]"/] D_LATENT_INFO -- O_LATENT_SPLIT O_LATENT_SPLIT -- D_MU O_LATENT_SPLIT -- D_LOGVAR R_EPS{ {"随机模块:标准高斯噪声(Gaussian Noise ε)br/Shape: [B, 32]"}} O_REPARAM(["操作模块:重参数化(Reparameterization Trick)br/公式:z = μ + exp(0.5·logvar) ⊙ ε"]) D_Z[/"数据模块:风格潜变量(Style Variable z)br/Shape: [B, 32]"/] D_MU -- O_REPARAM D_LOGVAR -- O_REPARAM R_EPS -- O_REPARAM O_REPARAM -- D_Z end subgraph S3["阶段 S3:多视角视觉 Backbone(Multi-view Visual Backbone)"] direction TB M_BACKBONE["模型模块:🔥 共享视觉主干(Shared ResNet18 Backbone)br/结构:ImageNet-pretrained ResNet18;❄️ FrozenBatchNorm2d + 🔥 Convolution Weights"] D_RESNET[/"数据模块:逐相机 ResNet18 特征(Per-camera ResNet Features)br/Shape: [B, 4, 512, 15, 20]"/] D_IMG --|"Shared Weights"| M_BACKBONE M_BACKBONE -- D_RESNET M_VIS_PROJ["模型模块:🔥 视觉特征投影(Visual Input Projection)br/结构:1×1 Conv 512→512"] D_CAM_FEATURE[/"数据模块:逐相机投影视觉特征(Projected Camera Features)br/Shape: [B, 4, 512, 15, 20]"/] D_RESNET -- M_VIS_PROJ -- D_CAM_FEATURE O_VIS_POS(["操作模块:二维正弦位置编码(2D Sinusoidal Position Encoding)"]) D_CAM_POS[/"数据模块:逐相机视觉位置编码(Per-camera Visual Positional Embedding)br/Shape: [B, 4, 512, 15, 20]"/] D_RESNET -- O_VIS_POS -- D_CAM_POS O_CAM_CAT(["操作模块:相机特征宽度拼接(Camera Feature Width Concatenation)"]) D_VIS_MAP[/"数据模块:拼接视觉特征图(Concatenated Visual Feature Map)br/Shape: [B, 512, 15, 80]"/] D_CAM_FEATURE -- O_CAM_CAT -- D_VIS_MAP O_POS_CAT(["操作模块:相机位置编码宽度拼接(Camera Position Width Concatenation)"]) D_VIS_POS_MAP[/"数据模块:拼接视觉位置编码(Concatenated Visual Position Map)br/Shape: [B, 512, 15, 80]"/] D_CAM_POS -- O_POS_CAT -- D_VIS_POS_MAP end subgraph S4["阶段 S4:ACT 条件 Token 与视觉 Token 构造(ACT Condition and Visual Token Construction)"] direction TB M_COND["模型模块:🔥 ACT 条件投影(ACT Condition Projections)br/结构:Latent Linear 32→512 + q_t Linear 14→512 + Learned 2-Token Positional Embedding"] D_Z_TOKEN[/"数据模块:潜变量条件 Token(Latent Condition Token)br/Shape: [B, 512]"/] D_Q_TOKEN[/"数据模块:本体感知条件 Token(Proprioceptive Condition Token)br/Shape: [B, 512]"/] D_COND_POS[/"数据模块:条件位置编码(Condition Positional Embedding)br/Shape: [2, B, 512]"/] D_Z -- M_COND D_QPOS -- M_COND M_COND -- D_Z_TOKEN M_COND -- D_Q_TOKEN M_COND -- D_COND_POS O_VIS_FLAT(["操作模块:视觉特征空间展平(Visual Spatial Flatten)"]) D_VIS_TOKEN[/"数据模块:视觉 Token 序列(Visual Token Sequence)br/Shape: [1200, B, 512]"/] D_VIS_MAP -- O_VIS_FLAT -- D_VIS_TOKEN O_VIS_POS_FLAT(["操作模块:视觉位置编码空间展平(Visual Position Flatten)"]) D_VIS_POS[/"数据模块:视觉 Token 位置编码(Visual Token Positional Embedding)br/Shape: [1200, B, 512]"/] D_VIS_POS_MAP -- O_VIS_POS_FLAT -- D_VIS_POS O_COND_STACK(["操作模块:条件 Token 堆叠(Condition Token Stacking)br/顺序:z → q_t"]) D_COND_TOKEN[/"数据模块:ACT 条件 Token(ACT Condition Tokens)br/Shape: [2, B, 512]"/] D_Z_TOKEN -- O_COND_STACK D_Q_TOKEN -- O_COND_STACK O_COND_STACK -- D_COND_TOKEN O_SRC_CAT(["操作模块:ACT Encoder 输入拼接(ACT Encoder Input Concatenation)br/顺序:z Token → q_t Token → Visual Tokens"]) D_ACT_SRC[/"数据模块:ACT Transformer Encoder 输入(ACT Transformer Encoder Input)br/Shape: [1202, B, 512]"/] D_COND_TOKEN -- O_SRC_CAT D_VIS_TOKEN -- O_SRC_CAT O_SRC_CAT -- D_ACT_SRC O_MAIN_POS_CAT(["操作模块:ACT 位置编码拼接(ACT Positional Encoding Concatenation)"]) D_ACT_POS[/"数据模块:ACT Transformer 位置编码(ACT Transformer Positional Embedding)br/Shape: [1202, B, 512]"/] D_COND_POS -- O_MAIN_POS_CAT D_VIS_POS -- O_MAIN_POS_CAT O_MAIN_POS_CAT -- D_ACT_POS end subgraph S5["阶段 S5:ACT Transformer 动作块生成(ACT Transformer Action Chunk Generation)"] direction TB M_ACT_ENC["模型模块:🔥 ACT Transformer 编码器(ACT Transformer Encoder)br/结构:Self-Attention + FFN × 4,d_model = 512,8 Heads,FFN = 3200"] D_MEMORY[/"数据模块:ACT Encoder Memory(Encoder Memory)br/Shape: [1202, B, 512]"/] D_ACT_SRC -- M_ACT_ENC D_ACT_POS -. "PosEmb" .- M_ACT_ENC M_ACT_ENC -- D_MEMORY M_QUERY["模型模块:🔥 动作查询嵌入(Action Query Embedding)br/结构:Learned nn.Embedding(k, 512)"] D_QUERY[/"数据模块:动作查询(Action Queries)br/Shape: [k, B, 512]"/] M_QUERY -- D_QUERY O_ZERO_TGT(["操作模块:Decoder Target 零初始化(Zero Decoder Target Initialization)"]) D_TGT[/"数据模块:Decoder 初始目标(Initial Decoder Target)br/Shape: [k, B, 512]"/] D_QUERY -- O_ZERO_TGT -- D_TGT M_ACT_DEC["模型模块:🔥 ACT Transformer 解码器(ACT Transformer Decoder)br/结构:Self-Attention + Cross-Attention + FFN × 7"] D_DEC_H[/"数据模块:动作查询隐藏表示(Action Query Hidden States)br/Shape: [B, k, 512]"/] D_TGT -- M_ACT_DEC D_QUERY -. "Query Pos" .- M_ACT_DEC D_MEMORY --|"Memory"| M_ACT_DEC D_ACT_POS -. "Memory Pos" .- M_ACT_DEC M_ACT_DEC -- D_DEC_H M_ACTION_HEAD["模型模块:🔥 动作输出头(Action Head)br/结构:官方实现 Linear 512→14"] D_PRED[/"数据模块:预测动作块(Predicted Action Chunk Â_t)br/Shape: [B, k, 14]"/] D_DEC_H -- M_ACTION_HEAD -- D_PRED end subgraph S6["阶段 S6:联合损失与优化目标(Joint Training Objective)"] direction TB subgraph SL["子损失(Sub Losses)"] direction LR L_REC["损失模块:掩码 L1 动作重建损失(Masked L1 Reconstruction Loss / L_reconst)br/Shape: Scalar"] L_KL["损失模块:KL 正则损失(KL Regularization Loss / L_KL)br/Shape: Scalar"] end L_TOTAL["损失模块:总损失(Total Loss / L_total)br/Shape: Scalar"] L_REC -- L_TOTAL L_KL --|"β = 10"| L_TOTAL end D_PRED --|"Prediction"| L_REC D_ACT --|"Target"| L_REC D_PAD -. "Mask" .- L_REC D_MU -- L_KL D_LOGVAR -- L_KL L_TOTAL -. "Gradient" .- M_ACTION_HEAD L_TOTAL -. "Gradient" .- M_ACT_DEC L_TOTAL -. "Gradient" .- M_QUERY L_TOTAL -. "Gradient" .- M_ACT_ENC L_TOTAL -. "Gradient" .- M_COND L_TOTAL -. "Gradient" .- M_VIS_PROJ L_TOTAL -. "Gradient: Conv Weights" .- M_BACKBONE L_TOTAL -. "Gradient via z" .- M_LATENT_HEAD L_TOTAL -. "Gradient via z" .- M_CVAE_ENC L_TOTAL -. "Gradient via z" .- M_CVAE_EMB classDef data fill:#F3F4F6,stroke:#6B7280,stroke-width:1.5px,color:#111827; classDef model fill:#ECFDF3,stroke:#16A34A,stroke-width:2.5px,color:#111827; classDef frozen fill:#F0FDF4,stroke:#65A30D,stroke-width:2px,color:#111827; classDef operation fill:#F5F3FF,stroke:#7C3AED,stroke-width:1.8px,color:#111827; classDef random fill:#FFF7E6,stroke:#D97706,stroke-width:1.8px,color:#111827; classDef loss fill:#FFF1F2,stroke:#DC2626,stroke-width:2.5px,color:#111827; class D_IMG_RAW,D_QPOS_RAW,D_ACT_RAW,D_PAD_RAW,D_IMG,D_QPOS,D_ACT,D_PAD,D_CLS,D_QPOS_CVAE,D_ACT_EMB,D_CVAE_SEQ_BF,D_CVAE_SEQ,D_CVAE_POS,D_CVAE_PAD,D_CVAE_H,D_CLS_H,D_LATENT_INFO,D_MU,D_LOGVAR,D_Z,D_RESNET,D_CAM_FEATURE,D_CAM_POS,D_VIS_MAP,D_VIS_POS_MAP,D_Z_TOKEN,D_Q_TOKEN,D_COND_POS,D_VIS_TOKEN,D_VIS_POS,D_COND_TOKEN,D_ACT_SRC,D_ACT_POS,D_MEMORY,D_QUERY,D_TGT,D_DEC_H,D_PRED data; class M_CVAE_EMB,M_CVAE_ENC,M_LATENT_HEAD,M_BACKBONE,M_VIS_PROJ,M_COND,M_ACT_ENC,M_QUERY,M_ACT_DEC,M_ACTION_HEAD model; class O_IMG_PRE,O_QPOS_NORM,O_ACT_CHUNK,O_PAD_CHUNK,O_CVAE_CAT,O_CVAE_PERMUTE,O_PAD_EXT,O_CLS_SELECT,O_LATENT_SPLIT,O_REPARAM,O_VIS_POS,O_CAM_CAT,O_POS_CAT,O_VIS_FLAT,O_VIS_POS_FLAT,O_COND_STACK,O_SRC_CAT,O_MAIN_POS_CAT,O_ZERO_TGT operation; class R_EPS random; class L_REC,L_KL,L_TOTAL loss;这张图里有几个特别重要、也最容易画错的地方。1.Backbone(主干网络)是 ResNet18,不是 CVAE官方实现创建一个 ResNet18 backbone,并在循环里对所有 camera 调用同一个self.backbones[0],所以这些摄像头共享同一套 Backbone 权重。ResNet18 的最后一层特征通道数为 512;对论文的480×640480\times640480×640图像,对应空间分辨率为15×2015\times2015×20。不同相机的特征不是按 token 序列直接相加,而是在代码中沿 width dimension 拼接。turn840745view0因此四相机情况下:[B,4,512,15,20]→[B,512,15,80][B,4,512,15,20]\rightarrow[B,512,15,80][B,4,512,15,20]→[B,512,15,80]再 flatten:[B,512,15,80]→[1200,B,512][B,512,15,80]\rightarrow[1200,B,512][B,512,15,80]→[1200,B,512]这比把图简单画成“4 cameras → 1200 tokens”更忠实于官方实现。2. Backbone 整体不是完全 Frozen这里也值得纠正一个常见误解。官方backbone.py使用了FrozenBatchNorm2d,但是把冻结整个 ResNet 参数的代码注释掉了;优化器还专门把"backbone"参数放到一个独立 parameter group,并使用:lrbackbone=10−5\mathrm{lr}_{\mathrm{backbone}}=10^{-5}lrbackbone​=10−5因此更准确的表示是:部分训练模型(Partially Trainable Model):BatchNorm statistics / parameters 采用 frozen 实现,而 ResNet convolutional parameters 参与训练。官方优化器是 AdamW,并给 backbone 单独设置学习率。turn382943view0turn936103view0所以图里写成:🔥 Shared ResNet18 Backbone同时在结构行中明确:❄️ FrozenBatchNorm2d + 🔥 Convolution Weights而不是把整个 ResNet18 画成 ❄️。3. CVAE 和视觉 Backbone 是两条完全不同的分支CVAE posterior branch不读取图像。它只读取:qtq_tqt​和:At=at:t+kA_t=a_{t:t+k}At​=at:t+k​然后构造:[CLS],qt,at,…,at+k−1[\mathrm{CLS}],q_t,a_t,\ldots,a_{t+k-1}[CLS],qt​,at​,…,at+k−1​
返回列表