#!/usr/bin/env python3 """ data.py (stage3 / Flow Matching — Pi0.5 风格) ============================================== 多数据集 LeRobot 格式数据加载器 (兼容 v2.0 和 v3.0),支持: - 多数据集加权混合采样 (BridgeData / DROID / Libero / 自有数据) - 机器人状态 (proprioception) 读取与归一化 - Action chunking (K=16) - 多相机输入 (base_rgb + wrist_rgb) - 域随机化 (Domain Randomization) 缓解背景过拟合 - 标准 LeRobot 命名 与 HuggingFace LeRobot 下载的 "data__chunk__file" 命名兼容 - 缺视频时自动回退到占位图,避免训练中断 v2.0 vs v3.0: v2.0: 1 parquet = 1 episode, 视频 episode_*.mp4, task_index + tasks.jsonl v3.0: 1 parquet = N episodes (is_first/is_last 分界), 视频 file-*.mp4, 多 camera keys, language_instruction 列直接读指令 接口契约 ======== build_dataloader(config) -> (train_loader, val_loader) 每个 batch 必须包含: { "image": List[PIL.Image] 或 List[List[PIL.Image]], "instruction": List[str], "state": torch.Tensor, # (B, state_dim) 或 None "action": torch.Tensor, # (B, action_dim) "action_chunk": torch.Tensor, # (B, K, action_dim) "raw_action": torch.Tensor, # (B, action_dim) "frame_index": List[int], "episode_index": List[int], "dataset_name": List[str], # 可选 } config 示例见 config_flow.yaml;data.py 新增字段: data: datasets: - name: bridge dir: /mnt/workspace/Dataset/stage_frozen01/full weight: 0.4 - name: droid dir: /mnt/workspace/Dataset/droid/full weight: 0.2 camera_keys: [observation.images.exterior_1_left, ...] val_ratio: 0.05 action_chunk_size: 16 action_dim: 7 state_dim: 7 domain_randomization: true """ from __future__ import annotations import json import glob import os import random import threading from pathlib import Path from typing import Any, Callable, Dict, List, Optional, Tuple, Union import numpy as np import pandas as pd import torch from PIL import Image from torch.utils.data import ( ConcatDataset, DataLoader, Dataset, WeightedRandomSampler, ) try: from .utils.chunking import TemporalEnsembler HAS_CHUNKING = True except ImportError: HAS_CHUNKING = False try: from torchvision import transforms except ModuleNotFoundError: transforms = None # type: ignore # 视频解码 (优先 decord,否则 imageio) try: import decord decord.bridge.set_bridge("torch") HAS_DECORD = True except Exception: HAS_DECORD = False try: import imageio except Exception: imageio = None logger = __import__("logging").getLogger("LingArm.stage3") # --------------------------------------------------------------------------- # 视频帧缓存 # --------------------------------------------------------------------------- class VideoCache: """按 episode 缓存解码后的视频帧 (LRU)。""" def __init__(self, max_cached_episodes: int = 200): self.cache: Dict[str, np.ndarray] = {} self.access_order: List[str] = [] self.max_cached = max_cached_episodes self.lock = threading.Lock() def __getstate__(self): state = self.__dict__.copy() state.pop("lock", None) state["cache"] = {} # 不跨进程传输大缓存 state["access_order"] = [] return state def __setstate__(self, state): self.__dict__.update(state) self.lock = threading.Lock() def get_frame( self, video_path: str, frame_idx: int, ) -> np.ndarray: with self.lock: if video_path not in self.cache: self._load_video(video_path) else: self.access_order.remove(video_path) self.access_order.append(video_path) frames = self.cache[video_path] if frame_idx >= len(frames): frame_idx = len(frames) - 1 return frames[frame_idx] def _load_video(self, video_path: str): while len(self.cache) >= self.max_cached and self.access_order: old = self.access_order.pop(0) self.cache.pop(old, None) if not os.path.exists(video_path): raise FileNotFoundError(f"Video not found: {video_path}") if HAS_DECORD: vr = decord.VideoReader(video_path) frames = vr.get_batch(list(range(len(vr)))) # decord 0.6.0+ 返回 torch.Tensor, 旧版返回 NDArray with .asnumpy() if hasattr(frames, 'numpy'): self.cache[video_path] = frames.numpy() else: self.cache[video_path] = frames.asnumpy() elif imageio is not None: reader = imageio.get_reader(video_path) self.cache[video_path] = np.array([ reader.get_data(i) for i in range(reader.count_frames()) ]) reader.close() else: raise RuntimeError( "Need decord or imageio to read videos. " "Install: pip install decord" ) self.access_order.append(video_path) # --------------------------------------------------------------------------- # 缺视频时的占位图 # --------------------------------------------------------------------------- class DummyImageCache: """为没有视频的样本生成一张随机纹理占位图。""" def __init__(self, size: int = 224): self.size = size self._cache: Dict[Tuple[int, int], Image.Image] = {} def get(self, seed: int = 0) -> Image.Image: key = (self.size, seed) if key not in self._cache: rng = np.random.default_rng(seed) arr = rng.integers(0, 255, (self.size, self.size, 3), dtype=np.uint8) self._cache[key] = Image.fromarray(arr).convert("RGB") return self._cache[key] DUMMY_IMAGE = DummyImageCache() # --------------------------------------------------------------------------- # LeRobot 数据集 # --------------------------------------------------------------------------- class LeRobotDataset(Dataset): """ LeRobot 格式数据集加载器 (stage3), 兼容 v2.0 和 v3.0。 v2.0: 1 parquet = 1 episode, 视频 episode_*.mp4, task_index → tasks.jsonl v3.0: 1 parquet = N episodes (is_first/is_last 分界), 视频 file-*.mp4, videos/{cam}/chunk-*/ 路径, language_instruction 列直接读取任务指令 Args: data_dir: 数据集根目录 (含 meta/info.json) dataset_name: 数据集名称 (用于日志/加权) action_dim: 目标动作维度 (默认 7) state_dim: 目标状态维度 (默认 7) action_chunk_size: K (默认 16) image_size: 图像目标尺寸 (PIL resize;最终由 processor 决定) action_normalize: 是否 z-score 标准化 action state_normalize: 是否 z-score 标准化 state 后 clip 到 [-1,1] state_low/state_high: state clip 范围 frame_sampling: "all" | "uniform" max_frames_per_episode: uniform 采样时最大帧数 use_processor: True 返回 PIL.Image,False 返回 tensor camera_keys: 多相机 key 列表 dummy_image_on_missing_video: 缺视频是否用占位图 augment: 是否启用域随机化 action_adapter: 可选动作转换函数 state_adapter: 可选状态转换函数 val_episode_ids: 若提供,仅使用这些 episode (验证集) """ def __init__( self, data_dir: str, dataset_name: str = "lerobot", action_dim: int = 7, state_dim: int = 7, action_chunk_size: int = 16, image_size: int = 224, action_normalize: bool = True, state_normalize: bool = True, state_low: float = -1.0, state_high: float = 1.0, frame_sampling: str = "all", max_frames_per_episode: Optional[int] = None, use_processor: bool = True, max_cache_episodes: int = 200, camera_keys: Optional[List[str]] = None, dummy_image_on_missing_video: bool = True, augment: bool = True, action_adapter: Optional[Callable[[np.ndarray, str], np.ndarray]] = None, state_adapter: Optional[Callable[[np.ndarray, str], np.ndarray]] = None, val_episode_ids: Optional[set] = None, skip_stats_compute: Optional[bool] = None, ): self.data_dir = Path(data_dir) self.dataset_name = dataset_name self.action_dim = action_dim self.state_dim = state_dim self.action_chunk_size = action_chunk_size self.image_size = image_size self.action_normalize = action_normalize self.state_normalize = state_normalize self.state_low = state_low self.state_high = state_high self.frame_sampling = frame_sampling self.max_frames_per_episode = max_frames_per_episode self.use_processor = use_processor self.camera_keys = camera_keys or ["observation.images.image_0"] self.dummy_image_on_missing_video = dummy_image_on_missing_video self.augment = augment self.action_adapter = action_adapter self.state_adapter = state_adapter self.val_episode_ids = val_episode_ids # 检测 LeRobot 格式版本 (v2.0 / v3.0) self.codebase_version = self._detect_version() logger.info("[Dataset %s] codebase_version=%s", dataset_name, self.codebase_version) # 读取 meta / 任务指令 self.tasks = self._load_tasks() # 构建 episode 索引 self.episodes = self._build_index() if not self.episodes: raise RuntimeError(f"No valid episodes found in {data_dir}") # 计算统计量 (仅训练集) if skip_stats_compute is None: skip_stats_compute = (val_episode_ids is not None) if skip_stats_compute: self.action_mean = self.action_std = None self.state_mean = self.state_std = None else: self.action_mean, self.action_std = self._compute_action_stats() self.state_mean, self.state_std = self._compute_state_stats() # 视频缓存 self.video_cache = VideoCache(max_cached_episodes=max_cache_episodes) # 图像预处理 if not use_processor: if transforms is None: raise ImportError("torchvision is required when use_processor=False") self.image_transform = transforms.Compose([ transforms.Resize((image_size, image_size)), transforms.ToTensor(), transforms.Normalize( mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225], ), ]) else: self.image_transform = None # 域随机化 (仅对 PIL 做轻量变换) self.aug_transform = None if augment and transforms is not None: self.aug_transform = transforms.Compose([ transforms.ColorJitter(brightness=0.15, contrast=0.15, saturation=0.15, hue=0.05), transforms.RandomGrayscale(p=0.05), ]) # 镜像翻转增广 (保持关闭, 镜像几何变换需要逐相机标定) self.hflip_aug = False self.hflip_prob = 0.5 # 随机透视抖动 (模拟相机位置变化) self.perspective_aug = augment self.perspective_distortion = 0.15 total_frames = sum(len(ep["frame_indices"]) for ep in self.episodes) logger.info( "[Dataset %s] episodes=%d frames=%d action_dim=%d state_dim=%d cameras=%d", dataset_name, len(self.episodes), total_frames, self.action_dim, self.state_dim, len(self.camera_keys), ) # ------------------------------------------------------------------ # 版本检测 # ------------------------------------------------------------------ def _detect_version(self) -> str: """从 meta/info.json 读取 codebase_version,回退到 v2.0。""" for candidate in ["meta/info.json", "info.json"]: info_path = self.data_dir / candidate if info_path.exists(): try: with open(info_path, "r", encoding="utf-8") as f: info = json.load(f) return info.get("codebase_version", "v2.0") except Exception: pass return "v2.0" # ------------------------------------------------------------------ # Meta / 任务指令 # ------------------------------------------------------------------ def _load_tasks(self) -> Dict[int, str]: """加载任务指令表。v3.0 返回空 dict (指令从 parquet language_instruction 列读取)。""" if self.codebase_version == "v3.0": return {} tasks: Dict[int, str] = {} # 1) meta/tasks.jsonl jsonl_path = self.data_dir / "meta" / "tasks.jsonl" if jsonl_path.exists(): with open(jsonl_path, "r", encoding="utf-8") as f: for line in f: line = line.strip() if not line: continue obj = json.loads(line) tasks[obj["task_index"]] = obj.get("task", "") return tasks # 2) meta__tasks.parquet (HF LeRobot 下载格式) parquet_paths = list(self.data_dir.glob("meta__tasks*.parquet")) if not parquet_paths: parquet_paths = list(self.data_dir.glob("*/meta__tasks*.parquet")) for p in parquet_paths: try: df = pd.read_parquet(p) if "task_index" not in df.columns: continue has_task_col = "task" in df.columns for idx, row in df.iterrows(): ti = int(row["task_index"]) task_text = str(row["task"]) if has_task_col else ( str(idx) if isinstance(idx, str) else "") tasks[ti] = task_text return tasks except Exception as e: logger.warning("[Dataset %s] failed to read task parquet %s: %s", self.dataset_name, p, e) return tasks # ------------------------------------------------------------------ # 索引构建 # ------------------------------------------------------------------ def _find_parquet_files(self) -> List[Path]: patterns = [ "data/chunk-*/episode_*.parquet", # 标准 v2.0 "data/chunk-*/file-*.parquet", # 标准 v3.0 (DROID) "data__chunk-*__file-*.parquet", # HF 下载格式 "*/data__chunk-*__file-*.parquet", # 嵌套 HF 下载 ] files: List[Path] = [] for pat in patterns: files.extend(self.data_dir.glob(pat)) return sorted(set(files)) def _build_index(self) -> List[Dict[str, Any]]: episodes = [] for pf in self._find_parquet_files(): try: df = pd.read_parquet(pf) except Exception as e: logger.warning("[Dataset %s] skip parquet %s: %s", self.dataset_name, pf, e) continue if "action" not in df.columns: logger.warning("[Dataset %s] skip %s: no action column", self.dataset_name, pf) continue if len(df) == 0: continue if self.codebase_version == "v3.0": episodes.extend(self._build_index_v3(pf, df)) else: episodes.extend(self._build_index_v2(pf, df)) return episodes def _build_index_v2(self, pf: Path, df: pd.DataFrame) -> List[Dict[str, Any]]: """v2.0: 1 parquet = 1 episode。""" n_frames = len(df) if "episode_index" in df.columns: ep_idx = int(df["episode_index"].iloc[0]) else: ep_idx = 0 if self.val_episode_ids is not None and ep_idx not in self.val_episode_ids: return [] chunk_name, video_stem = self._video_location(pf) video_paths_per_cam, has_any_video = self._resolve_video_paths_v2(chunk_name, video_stem) task_index = int(df["task_index"].iloc[0]) if "task_index" in df.columns else 0 instruction = self.tasks.get(task_index, "") frame_indices = self._sample_frame_indices(n_frames, range(n_frames)) raw_frame_indices = ( df["frame_index"].tolist() if "frame_index" in df.columns else list(range(n_frames)) ) actions_raw = self._extract_actions(df) states_raw = self._extract_states(df) return [{ "parquet": pf, "video_paths": video_paths_per_cam, "has_video": has_any_video, "n_frames": n_frames, "frame_indices": frame_indices, "instruction": instruction, "task_index": task_index, "episode_index": ep_idx, "actions_raw": actions_raw, "states_raw": states_raw, "raw_frame_indices": raw_frame_indices, "frame_offset": 0, # v2.0: 无偏移 (1 episode = 1 mp4) }] def _build_index_v3(self, pf: Path, df: pd.DataFrame) -> List[Dict[str, Any]]: """v3.0: 1 parquet = N episodes, 按 is_first/is_last 拆分。""" if "is_first" not in df.columns or "is_last" not in df.columns: logger.warning( "[Dataset %s] v3.0 parquet lacks is_first/is_last: %s, " "treating as single episode", self.dataset_name, pf ) return self._build_index_v2(pf, df) n_total = len(df) first_rows = df.index[df["is_first"] == True].tolist() last_rows = df.index[df["is_last"] == True].tolist() if not first_rows: first_rows = [0] if not last_rows: last_rows = [n_total - 1] if len(last_rows) < len(first_rows): last_rows.append(n_total - 1) chunk_name, video_stem = self._video_location(pf) video_paths_per_cam, has_any_video = self._resolve_video_paths_v3(chunk_name, video_stem) episodes = [] for ep_i, (start_row, end_row) in enumerate(zip(first_rows, last_rows)): if end_row < start_row: continue ep_df = df.loc[start_row:end_row] n_frames = len(ep_df) if n_frames == 0: continue # episode_index if "episode_index" in ep_df.columns: ep_idx = int(ep_df["episode_index"].iloc[0]) else: ep_idx = ep_i if self.val_episode_ids is not None and ep_idx not in self.val_episode_ids: continue # v3.0: 语言指令从 language_instruction 列读取 if "language_instruction" in ep_df.columns: instruction = str(ep_df["language_instruction"].iloc[0]) else: instruction = "" # 帧采样 (0-based within episode) frame_indices = self._sample_frame_indices(n_frames, range(n_frames)) # 原始 frame_index (parquet 中的绝对帧号 → mp4 内的帧位置) if "frame_index" in ep_df.columns: raw_frame_indices = ep_df["frame_index"].tolist() else: raw_frame_indices = list(range(n_frames)) frame_offset = int(raw_frame_indices[0]) if raw_frame_indices else 0 actions_raw = self._extract_actions(ep_df) states_raw = self._extract_states(ep_df) episodes.append({ "parquet": pf, "video_paths": video_paths_per_cam, "has_video": has_any_video, "n_frames": n_frames, "frame_indices": frame_indices, "instruction": instruction, "task_index": 0, "episode_index": ep_idx, "actions_raw": actions_raw, "states_raw": states_raw, "raw_frame_indices": raw_frame_indices, "frame_offset": frame_offset, # mp4 内起始帧号 }) return episodes # ------------------------------------------------------------------ # 视频路径解析 # ------------------------------------------------------------------ def _video_location(self, parquet_path: Path) -> Tuple[str, str]: """从 parquet 路径推断 chunk 名和视频 stem。""" stem = parquet_path.stem # e.g. episode_000000 or file-000 or data__chunk-000__file-000 if stem.startswith("data__"): parts = stem.split("__") return parts[1], parts[2] # chunk-000, file-000 else: chunk_name = parquet_path.parent.name # chunk-xxx return chunk_name, stem def _resolve_video_paths_v2( self, chunk_name: str, video_stem: str ) -> Tuple[Dict[str, Optional[str]], bool]: """v2.0 视频路径: videos///.mp4""" video_paths: Dict[str, Optional[str]] = {} has_any = False for cam_key in self.camera_keys: cam_dir = cam_key.replace(".", "_") if "." in cam_key else cam_key candidates = [ self.data_dir / "videos" / chunk_name / cam_key / f"{video_stem}.mp4", self.data_dir / "videos" / chunk_name / cam_dir / f"{video_stem}.mp4", ] found = None for vp in candidates: if vp.exists(): found = str(vp) break video_paths[cam_key] = found if found: has_any = True return video_paths, has_any def _resolve_video_paths_v3( self, chunk_name: str, video_stem: str ) -> Tuple[Dict[str, Optional[str]], bool]: """v3.0 视频路径: videos///.mp4 DROID v3.0 使用 videos/{video_key}/chunk-*/file-*.mp4 格式。 尝试多种路径以兼容不同下载布局。 """ video_paths: Dict[str, Optional[str]] = {} has_any = False for cam_key in self.camera_keys: candidates = [ # (1) v3 标准: videos///.mp4 self.data_dir / "videos" / cam_key / chunk_name / f"{video_stem}.mp4", # (2) v2 风格回退: videos///.mp4 self.data_dir / "videos" / chunk_name / cam_key / f"{video_stem}.mp4", # (3) 下划线版本 self.data_dir / "videos" / chunk_name / cam_key.replace(".", "_") / f"{video_stem}.mp4", ] found = None for vp in candidates: if vp.exists(): found = str(vp) break video_paths[cam_key] = found if found: has_any = True return video_paths, has_any # ------------------------------------------------------------------ # 辅助: 帧采样 / action 提取 / state 提取 # ------------------------------------------------------------------ def _sample_frame_indices( self, n_frames: int, src_indices: Any ) -> List[int]: src = list(src_indices) if self.frame_sampling == "uniform" and self.max_frames_per_episode: if n_frames <= self.max_frames_per_episode: return src else: return np.linspace(0, n_frames - 1, self.max_frames_per_episode, dtype=int).tolist() return src def _extract_actions(self, df: pd.DataFrame) -> np.ndarray: arr = np.array(df["action"].tolist(), dtype=np.float32) if self.action_adapter is not None: arr = np.stack([self.action_adapter(a, self.dataset_name) for a in arr]) else: arr = self._adapt_dim_batch(arr, self.action_dim, "action") return arr def _extract_states(self, df: pd.DataFrame) -> Optional[np.ndarray]: if "observation.state" not in df.columns: return None arr = np.array(df["observation.state"].tolist(), dtype=np.float32) if self.state_adapter is not None: arr = np.stack([self.state_adapter(s, self.dataset_name) for s in arr]) else: arr = self._adapt_dim_batch(arr, self.state_dim, "state") return arr # ------------------------------------------------------------------ # 统计量 # ------------------------------------------------------------------ def _compute_action_stats(self) -> Tuple[np.ndarray, np.ndarray]: """批量 Welford 在线算法。""" dim = self.action_dim count = 0 mean = np.zeros(dim, dtype=np.float64) m2 = np.zeros(dim, dtype=np.float64) for ep in self.episodes: acts = ep["actions_raw"].astype(np.float64) n = len(acts) if n == 0: continue if count == 0: mean = acts.mean(axis=0) m2 = acts.var(axis=0) * n count = n else: batch_mean = acts.mean(axis=0) batch_var = acts.var(axis=0) delta = batch_mean - mean new_count = count + n new_mean = mean + delta * n / new_count m2 = m2 + batch_var * n + (delta ** 2) * count * n / new_count mean = new_mean count = new_count if count < 2: return np.zeros(dim, dtype=np.float32), np.ones(dim, dtype=np.float32) mean = mean.astype(np.float32) std = np.sqrt(m2 / count).astype(np.float32) std = np.where(std < 1e-6, 1.0, std) return mean, std def _compute_state_stats(self) -> Tuple[Optional[np.ndarray], Optional[np.ndarray]]: """批量 Welford 在线算法计算 state 统计。""" dim = self.state_dim count = 0 mean = np.zeros(dim, dtype=np.float64) m2 = np.zeros(dim, dtype=np.float64) for ep in self.episodes: if ep["states_raw"] is None: continue states = ep["states_raw"].astype(np.float64) n = len(states) if n == 0: continue if count == 0: mean = states.mean(axis=0) m2 = states.var(axis=0) * n count = n else: batch_mean = states.mean(axis=0) batch_var = states.var(axis=0) delta = batch_mean - mean new_count = count + n new_mean = mean + delta * n / new_count m2 = m2 + batch_var * n + (delta ** 2) * count * n / new_count mean = new_mean count = new_count if count < 2: return None, None mean = mean.astype(np.float32) std = np.sqrt(m2 / count).astype(np.float32) std = np.where(std < 1e-6, 1.0, std) return mean, std def set_stats( self, action_mean: Optional[np.ndarray], action_std: Optional[np.ndarray], state_mean: Optional[np.ndarray], state_std: Optional[np.ndarray], ): """验证集复用训练集统计量。""" self.action_mean = action_mean self.action_std = action_std self.state_mean = state_mean self.state_std = state_std # ------------------------------------------------------------------ # 适配到目标维度 # ------------------------------------------------------------------ def _adapt_action(self, raw: np.ndarray) -> np.ndarray: if self.action_adapter is not None: return self.action_adapter(raw, self.dataset_name) return self._pad_or_truncate(raw, self.action_dim, "action") def _adapt_state(self, raw: np.ndarray) -> np.ndarray: if self.state_adapter is not None: return self.state_adapter(raw, self.dataset_name) return self._pad_or_truncate(raw, self.state_dim, "state") def _pad_or_truncate(self, vec: np.ndarray, target_dim: int, name: str) -> np.ndarray: vec = np.asarray(vec, dtype=np.float32) d = vec.shape[0] if d == target_dim: return vec if d > target_dim: logger.warning( "[Dataset %s] %s dim %d > target %d, truncating first %d dims", self.dataset_name, name, d, target_dim, target_dim, ) return vec[:target_dim] logger.warning( "[Dataset %s] %s dim %d < target %d, zero-padding", self.dataset_name, name, d, target_dim, ) out = np.zeros(target_dim, dtype=np.float32) out[:d] = vec return out def _adapt_dim_batch(self, arr: np.ndarray, target_dim: int, name: str) -> np.ndarray: """批量 pad/truncate 到目标维度。""" d = arr.shape[1] if d == target_dim: return arr if d > target_dim: logger.warning( "[Dataset %s] %s dim %d > target %d, truncating (batch)", self.dataset_name, name, d, target_dim, ) return arr[:, :target_dim] logger.warning( "[Dataset %s] %s dim %d < target %d, zero-padding (batch)", self.dataset_name, name, d, target_dim, ) pad = np.zeros((arr.shape[0], target_dim - d), dtype=arr.dtype) return np.concatenate([arr, pad], axis=1) # ------------------------------------------------------------------ # 图像读取 # ------------------------------------------------------------------ def _read_images(self, ep: Dict[str, Any], frame_idx: int) -> List[Image.Image]: images: List[Image.Image] = [] seed = int(ep["episode_index"] * 100000 + frame_idx) # v3.0: 多 episode 共用 mp4, 实际帧号 = episode 内偏移 + mp4 起始偏移 actual_frame = frame_idx + ep.get("frame_offset", 0) for cam_key in self.camera_keys: vp = ep["video_paths"].get(cam_key) if vp and os.path.exists(vp): frame = self.video_cache.get_frame(vp, actual_frame) img = Image.fromarray(frame).convert("RGB") elif self.dummy_image_on_missing_video: img = DUMMY_IMAGE.get(seed).copy() seed += 1 else: raise FileNotFoundError(f"Missing video for camera {cam_key}: {vp}") if img.size[0] != self.image_size or img.size[1] != self.image_size: img = img.resize((self.image_size, self.image_size), Image.BILINEAR) if self.augment and self.aug_transform is not None and random.random() < 0.5: img = self.aug_transform(img) if self.perspective_aug and random.random() < 0.3: import torchvision.transforms.functional as TF w, h = img.size max_off = int(w * self.perspective_distortion) startpoints = [[0, 0], [w, 0], [w, h], [0, h]] endpoints = [ [random.randint(-max_off, max_off), random.randint(-max_off, max_off)], [w + random.randint(-max_off, max_off), random.randint(-max_off, max_off)], [w + random.randint(-max_off, max_off), h + random.randint(-max_off, max_off)], [random.randint(-max_off, max_off), h + random.randint(-max_off, max_off)], ] img = TF.perspective(img, startpoints, endpoints, interpolation=Image.BILINEAR) images.append(img) return images # ------------------------------------------------------------------ # Dataset 接口 # ------------------------------------------------------------------ def __len__(self) -> int: return sum(len(ep["frame_indices"]) for ep in self.episodes) def __getitem__(self, idx: int) -> Dict[str, Any]: cumulative = 0 for ep in self.episodes: n = len(ep["frame_indices"]) if idx < cumulative + n: frame_idx = ep["frame_indices"][idx - cumulative] break cumulative += n else: raise IndexError(f"Index {idx} out of range") images = self._read_images(ep, frame_idx) instruction = ep["instruction"] do_hflip = False if self.hflip_aug and random.random() < self.hflip_prob: do_hflip = True images = [img.transpose(Image.FLIP_LEFT_RIGHT) for img in images] instruction = instruction.replace("left", "\x00L\x00").replace("right", "left").replace("\x00L\x00", "right") raw_action = ep["actions_raw"][frame_idx].copy() if do_hflip and raw_action.shape[0] >= 6: raw_action[0] = -raw_action[0] raw_action[1] = -raw_action[1] raw_action[3] = -raw_action[3] raw_action[4] = -raw_action[4] norm_action = raw_action.copy() if self.action_normalize and self.action_mean is not None: norm_action = (raw_action - self.action_mean) / self.action_std if ep["states_raw"] is not None: raw_state = ep["states_raw"][frame_idx].copy() if do_hflip and raw_state.shape[0] >= 6: raw_state[0] = -raw_state[0] raw_state[1] = -raw_state[1] if raw_state.shape[0] >= 6: raw_state[3] = -raw_state[3] raw_state[4] = -raw_state[4] norm_state = raw_state.copy() if self.state_normalize and self.state_mean is not None: norm_state = (raw_state - self.state_mean) / self.state_std norm_state = np.clip(norm_state, self.state_low, self.state_high) else: norm_state = np.zeros(self.state_dim, dtype=np.float32) ep_actions = ep["actions_raw"] if self.action_normalize and self.action_mean is not None: ep_actions = (ep_actions - self.action_mean) / self.action_std if do_hflip and ep_actions.shape[1] >= 6: ep_actions = ep_actions.copy() ep_actions[:, 0] = -ep_actions[:, 0] ep_actions[:, 1] = -ep_actions[:, 1] ep_actions[:, 3] = -ep_actions[:, 3] ep_actions[:, 4] = -ep_actions[:, 4] T, D = ep_actions.shape K = self.action_chunk_size if T < K: ep_actions = np.concatenate([ep_actions, np.tile(ep_actions[-1:], (K - T, 1))], axis=0) T = K padded = np.concatenate([ep_actions, np.tile(ep_actions[-1:], (K - 1, 1))], axis=0) chunks = np.lib.stride_tricks.sliding_window_view(padded, (K, D))[:, 0, :, :] action_chunk = chunks[frame_idx].copy() if not self.use_processor: images = torch.stack([self.image_transform(img) for img in images]) return { "image": images, "instruction": instruction, "state": torch.from_numpy(norm_state).float(), "action": torch.from_numpy(norm_action).float(), "action_chunk": torch.from_numpy(action_chunk).float(), "raw_action": torch.from_numpy(raw_action).float(), "frame_index": ep["raw_frame_indices"][frame_idx], "episode_index": ep["episode_index"], "dataset_name": self.dataset_name, } # --------------------------------------------------------------------------- # Collate 函数 # --------------------------------------------------------------------------- def collate_fn(batch: List[Dict[str, Any]]) -> Dict[str, Any]: images = [b["image"] for b in batch] if isinstance(images[0], list): max_cam = max(len(imgs) for imgs in images) if max_cam > 1: for i in range(len(images)): if len(images[i]) < max_cam: images[i] = images[i] + [images[i][0]] * (max_cam - len(images[i])) elif isinstance(images[0], torch.Tensor): max_cam = max(imgs.shape[0] for imgs in images) if max_cam > 1: for i in range(len(images)): if images[i].shape[0] < max_cam: rep = images[i][0:1].repeat(max_cam - images[i].shape[0], 1, 1, 1) images[i] = torch.cat([images[i], rep], dim=0) images = torch.stack(images) B, N, C, H, W = images.shape images = images.view(B * N, C, H, W) states = torch.stack([b["state"] for b in batch]) actions = torch.stack([b["action"] for b in batch]) action_chunks = torch.stack([b["action_chunk"] for b in batch]) raw_actions = torch.stack([b["raw_action"] for b in batch]) return { "image": images, "instruction": [b["instruction"] for b in batch], "state": states, "action": actions, "action_chunk": action_chunks, "raw_action": raw_actions, "frame_index": [b["frame_index"] for b in batch], "episode_index": [b["episode_index"] for b in batch], "dataset_name": [b["dataset_name"] for b in batch], } # --------------------------------------------------------------------------- # DataLoader 构建 # --------------------------------------------------------------------------- def build_dataloader(config: dict) -> Tuple[DataLoader, Optional[DataLoader]]: """从 config 构建 (train_loader, val_loader)。""" data_cfg = config.get("data", {}) model_cfg = config.get("model", {}) train_cfg = config.get("training", {}) action_dim = data_cfg.get("action_dim", model_cfg.get("action_dim", 7)) state_dim = data_cfg.get("state_dim", model_cfg.get("state_dim", 7)) action_chunk_size = data_cfg.get("action_chunk_size", model_cfg.get("action_chunk_size", 16)) image_size = data_cfg.get("image_size", 224) use_processor = model_cfg.get("use_processor", True) val_ratio = data_cfg.get("val_ratio", 0.0) augment = data_cfg.get("domain_randomization", True) num_workers = data_cfg.get("num_workers", 1) pin_memory = data_cfg.get("pin_memory", True) dataset_configs = data_cfg.get("datasets", []) if not dataset_configs: dataset_configs = [{ "name": "single", "dir": config.get("data_dir", ""), "weight": 1.0, }] action_adapter = data_cfg.get("action_adapter", None) state_adapter = data_cfg.get("state_adapter", None) train_dataset_list: List[LeRobotDataset] = [] val_dataset_list: List[LeRobotDataset] = [] for ds_cfg in dataset_configs: ds_dir = ds_cfg["dir"] ds_name = ds_cfg.get("name", Path(ds_dir).name) if not ds_dir or not Path(ds_dir).exists(): logger.warning("[build_dataloader] dataset dir not found: %s", ds_dir) continue all_episodes = _collect_episode_indices(ds_dir) if not all_episodes: logger.warning("[build_dataloader] no episodes in %s", ds_dir) continue if val_ratio > 0 and len(all_episodes) > 1: n_val = max(1, int(len(all_episodes) * val_ratio)) split_rng = random.Random(config.get("seed", 42)) val_ids = set(split_rng.sample(sorted(all_episodes), n_val)) train_ids = all_episodes - val_ids else: val_ids = None train_ids = all_episodes common_kwargs = { "data_dir": ds_dir, "dataset_name": ds_name, "action_dim": action_dim, "state_dim": state_dim, "action_chunk_size": action_chunk_size, "image_size": image_size, "action_normalize": data_cfg.get("action_normalize", True), "state_normalize": data_cfg.get("state_normalize", True), "frame_sampling": data_cfg.get("frame_sampling", "all"), "max_frames_per_episode": data_cfg.get("max_frames_per_episode", None), "use_processor": use_processor, "max_cache_episodes": data_cfg.get("max_cache_episodes", 200), "camera_keys": ds_cfg.get("camera_keys", data_cfg.get("camera_keys", ["observation.images.image_0"])), "dummy_image_on_missing_video": data_cfg.get("dummy_image_on_missing_video", True), "augment": augment, "action_adapter": action_adapter, "state_adapter": state_adapter, } train_ds = LeRobotDataset(val_episode_ids=train_ids, skip_stats_compute=False, **common_kwargs) train_dataset_list.append(train_ds) if val_ids: val_ds = LeRobotDataset(val_episode_ids=val_ids, skip_stats_compute=True, **{**common_kwargs, "augment": False}) val_ds.set_stats( train_ds.action_mean, train_ds.action_std, train_ds.state_mean, train_ds.state_std, ) val_dataset_list.append(val_ds) if not train_dataset_list: raise RuntimeError("No training datasets could be loaded. Check your config.") train_concat = ConcatDataset(train_dataset_list) dataset_weights = [ds_cfg.get("weight", 1.0) for ds_cfg in dataset_configs] sample_weights = [] for ds, w in zip(train_dataset_list, dataset_weights): if len(ds) == 0: continue sample_weights.extend([w / max(1, len(ds))] * len(ds)) sample_weights = torch.tensor(sample_weights, dtype=torch.double) total_train_samples = len(sample_weights) samples_per_epoch = data_cfg.get("samples_per_epoch", total_train_samples) num_samples = min(samples_per_epoch, total_train_samples) logger.info("[build_dataloader] samples_per_epoch=%d (total=%d, capped=%s)", samples_per_epoch, total_train_samples, "yes" if num_samples < total_train_samples else "no") sampler = WeightedRandomSampler( weights=sample_weights, num_samples=num_samples, replacement=True, ) train_loader = DataLoader( train_concat, batch_size=train_cfg.get("batch_size", 16), sampler=sampler, num_workers=num_workers, pin_memory=pin_memory, collate_fn=collate_fn, drop_last=True, persistent_workers=num_workers > 0, ) val_loader = None if val_dataset_list: val_concat = ConcatDataset(val_dataset_list) val_loader = DataLoader( val_concat, batch_size=train_cfg.get("batch_size", 16), shuffle=False, num_workers=num_workers, pin_memory=pin_memory, collate_fn=collate_fn, drop_last=False, persistent_workers=num_workers > 0, ) logger.info("[build_dataloader] train datasets:") for ds in train_dataset_list: logger.info(" %s: %d samples", ds.dataset_name, len(ds)) if val_loader: logger.info("[build_dataloader] val datasets:") for ds in val_dataset_list: logger.info(" %s: %d samples", ds.dataset_name, len(ds)) _set_global_action_stats(train_concat, train_dataset_list, logger) return train_loader, val_loader def _set_global_action_stats( train_concat: ConcatDataset, dataset_list: List["LeRobotDataset"], logger, ): """per-dataset 归一化 + 虚拟全局统计 fallback。""" total_samples = 0 action_dim = dataset_list[0].action_dim weighted_mean = np.zeros(action_dim, dtype=np.float64) weighted_var = np.zeros_like(weighted_mean) for ds in dataset_list: n = len(ds) total_samples += n am = getattr(ds, "action_mean", None) as_ = getattr(ds, "action_std", None) if am is not None: weighted_mean += n * am.astype(np.float64) if as_ is not None: weighted_var += n * (as_.astype(np.float64) ** 2) if total_samples > 0: weighted_mean = (weighted_mean / total_samples).astype(np.float32) weighted_std = np.sqrt(weighted_var / total_samples).astype(np.float32) train_concat.action_mean = weighted_mean train_concat.action_std = weighted_std train_concat.state_mean = dataset_list[0].state_mean train_concat.state_std = dataset_list[0].state_std train_concat.per_dataset_stats = [ { "name": ds.dataset_name, "action_mean": getattr(ds, "action_mean", None), "action_std": getattr(ds, "action_std", None), "state_mean": getattr(ds, "state_mean", None), "state_std": getattr(ds, "state_std", None), } for ds in dataset_list ] logger.info("[action_stats] per-dataset normalization (no global mixing):") for st in train_concat.per_dataset_stats: logger.info(" %s: mean=%s std=%s", st["name"], st["action_mean"], st["action_std"]) logger.info("[action_stats] fallback global mean=%s std=%s (for model registration only)", weighted_mean, weighted_std) else: logger.warning("[action_stats] no action stats available, using default (0, 1)") def _collect_episode_indices(data_dir: str) -> set: """收集数据集里所有 episode_index (v2.0 + v3.0)。""" indices = set() base = Path(data_dir) patterns = [ "data/chunk-*/episode_*.parquet", # v2.0 "data/chunk-*/file-*.parquet", # v3.0 "data__chunk-*__file-*.parquet", "*/data__chunk-*__file-*.parquet", ] files: List[Path] = [] for pat in patterns: files.extend(base.glob(pat)) for pf in set(files): try: df = pd.read_parquet(pf) if "episode_index" in df.columns: indices.update(df["episode_index"].unique().tolist()) else: indices.add(0) except Exception: continue return indices # --------------------------------------------------------------------------- # 便捷测试 # --------------------------------------------------------------------------- if __name__ == "__main__": import logging logging.basicConfig(level=logging.INFO) test_cfg = { "data_dir": "G:/LingArm/datasets/pretrain_pool/droid/sample_download", "model": {"use_processor": True}, "training": {"batch_size": 2}, "data": { "num_workers": 0, "image_size": 224, "action_chunk_size": 16, "action_dim": 7, "state_dim": 7, "action_normalize": True, "state_normalize": True, "domain_randomization": True, "datasets": [ {"name": "droid", "dir": "G:/LingArm/datasets/pretrain_pool/droid/sample_download", "weight": 1.0}, ], }, } train_loader, val_loader = build_dataloader(test_cfg) batch = next(iter(train_loader)) print("Batch keys:", batch.keys()) print("image type:", type(batch["image"])) print("state shape:", batch["state"].shape) print("action shape:", batch["action"].shape) print("action_chunk shape:", batch["action_chunk"].shape) print("raw_action shape:", batch["raw_action"].shape) print("Instructions:", batch["instruction"][:2]) print("Datasets:", batch["dataset_name"][:5])