Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -276,6 +276,8 @@ def build_webdataset(webdataset_instance, **kwargs):
# ============================================================================
from cosmos_predict2._src.predict2.action.datasets.gr00t_dreams.data.dataset import MixedLeRobotDataset
from cosmos_predict2._src.predict2.action.datasets.gr00t_dreams.groot_configs import (
JHU_DVRK_MONO_FINETUNE_TRAIN_DATASET_SPECS,
JHU_DVRK_MONO_FINETUNE_VAL_DATASET_SPECS,
MAX_ACTION_DIM,
OPEN_H_DATASET_SPECS,
)
Expand Down Expand Up @@ -310,6 +312,67 @@ def build_webdataset(webdataset_instance, **kwargs):
drop_last=True,
)

# ============================================================================
# JHU dVRK monocular reference tabletop mixture
# ============================================================================
jhu_dvrk_mono_finetune_train_dataset = L(MixedLeRobotDataset)(
dataset_specs=JHU_DVRK_MONO_FINETUNE_TRAIN_DATASET_SPECS,
num_frames=13,
data_split="train",
max_action_dim=MAX_ACTION_DIM,
downscaled_res=False,
test_split_ratio=0.02,
)
jhu_dvrk_mono_finetune_val_dataset = L(MixedLeRobotDataset)(
dataset_specs=JHU_DVRK_MONO_FINETUNE_VAL_DATASET_SPECS,
num_frames=13,
data_split="test",
max_action_dim=MAX_ACTION_DIM,
downscaled_res=False,
test_split_ratio=0.02,
)
jhu_dvrk_mono_finetune_train_dataloader = L(DataLoader)(
dataset=jhu_dvrk_mono_finetune_train_dataset,
sampler=L(get_sampler)(dataset=jhu_dvrk_mono_finetune_train_dataset),
batch_size=1,
drop_last=True,
)
jhu_dvrk_mono_finetune_val_dataloader = L(DataLoader)(
dataset=jhu_dvrk_mono_finetune_val_dataset,
sampler=L(get_sampler)(dataset=jhu_dvrk_mono_finetune_val_dataset),
batch_size=1,
drop_last=True,
)

jhu_dvrk_mono_finetune_h73_train_dataset = L(MixedLeRobotDataset)(
dataset_specs=JHU_DVRK_MONO_FINETUNE_TRAIN_DATASET_SPECS,
num_frames=73,
data_split="train",
max_action_dim=MAX_ACTION_DIM,
downscaled_res=False,
test_split_ratio=0.02,
)
jhu_dvrk_mono_finetune_h73_val_dataset = L(MixedLeRobotDataset)(
dataset_specs=JHU_DVRK_MONO_FINETUNE_VAL_DATASET_SPECS,
num_frames=73,
data_split="test",
max_action_dim=MAX_ACTION_DIM,
downscaled_res=False,
test_split_ratio=0.02,
)
jhu_dvrk_mono_finetune_h73_train_dataloader = L(DataLoader)(
dataset=jhu_dvrk_mono_finetune_h73_train_dataset,
sampler=L(get_sampler)(dataset=jhu_dvrk_mono_finetune_h73_train_dataset),
batch_size=1,
drop_last=True,
)
jhu_dvrk_mono_finetune_h73_val_dataloader = L(DataLoader)(
dataset=jhu_dvrk_mono_finetune_h73_val_dataset,
sampler=L(get_sampler)(dataset=jhu_dvrk_mono_finetune_h73_val_dataset),
batch_size=1,
drop_last=True,
)


# ============================================================================
# SutureBot Dataset Configuration
Expand Down Expand Up @@ -453,6 +516,32 @@ def register_training_and_val_data():
node=open_h_multi_val_dataloader,
)

# JHU dVRK monocular tabletop reference recipe (short and long horizons).
cs.store(
group="data_train",
package="dataloader_train",
name="jhu_dvrk_mono_finetune_train",
node=jhu_dvrk_mono_finetune_train_dataloader,
)
cs.store(
group="data_val",
package="dataloader_val",
name="jhu_dvrk_mono_finetune_val",
node=jhu_dvrk_mono_finetune_val_dataloader,
)
cs.store(
group="data_train",
package="dataloader_train",
name="jhu_dvrk_mono_finetune_h73_train",
node=jhu_dvrk_mono_finetune_h73_train_dataloader,
)
cs.store(
group="data_val",
package="dataloader_val",
name="jhu_dvrk_mono_finetune_h73_val",
node=jhu_dvrk_mono_finetune_h73_val_dataloader,
)

# ============================================================================
# SutureBot dataset (20D actions zero-padded to 44D)
# ============================================================================
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -16,9 +16,11 @@
# Configs for resuming from stage3 training

import functools
import os

from hydra.core.config_store import ConfigStore

from cosmos_predict2._src.imaginaire.functional.lr_scheduler import LambdaWarmUpCosineScheduler
from cosmos_predict2._src.imaginaire.lazy_config import LazyCall as L
from cosmos_predict2._src.imaginaire.lazy_config import LazyDict
from cosmos_predict2._src.imaginaire.utils.checkpoint_db import get_checkpoint_path
Expand All @@ -37,6 +39,19 @@

DEFAULT_CHECKPOINT = MODEL_CHECKPOINTS[ModelKey()] # This uses post_trained=True by default

_TABLETOP_OUTPUT_ROOT = os.environ.get("IMAGINAIRE_OUTPUT_ROOT", "imaginaire/output")
_TABLETOP_CHSS_CHECKPOINT = os.environ.get(
"CHSS_CHECKPOINT_DIR",
"checkpoints/cosmos-h-surgical-simulator",
)


def _tabletop_teacher_checkpoint(run_name: str, iteration: int) -> str:
return (
f"{_TABLETOP_OUTPUT_ROOT}/cosmos_predict2_action_conditioned/"
f"official_runs_vid2vid/{run_name}/checkpoints/iter_{iteration:09d}"
)

_TRAINER_DEBUG_CONFIG = dict(
max_iter=1000,
logging_iter=50,
Expand Down Expand Up @@ -951,6 +966,111 @@ def build_debug_runs(job):
flags={"allow_objects": True},
)

# =============================================================================
# JHU dVRK monocular reference tabletop recipe
# =============================================================================
AC_CHUNK_SINGLE_VIEW_2B_JHU_DVRK_MONO_FINETUNE_13FRAME_8NODES_OSS = LazyDict(
dict(
defaults=[
"/experiment/2b_bridge_action_conditioned_oss",
{"override /net": "cosmos_v1_2B_action_chunk_conditioned"},
{"override /data_train": "jhu_dvrk_mono_finetune_train"},
{"override /data_val": "jhu_dvrk_mono_finetune_val"},
"_self_",
],
job=dict(
group="official_runs_vid2vid",
name="cosmos_predict2p5_2B_action_conditioned_jhu_dvrk_mono_finetune_13frame_8nodes_release_oss",
project="cosmos_predict2_action_conditioned",
),
checkpoint=dict(
load_path=_TABLETOP_CHSS_CHECKPOINT,
load_training_state=False,
strict_resume=False,
),
model=dict(
config=dict(
state_t=1 + 12 // 4,
net=dict(action_dim=44),
),
),
dataloader_train=dict(batch_size=16),
optimizer=dict(lr=1.6e-4, weight_decay=0.1),
trainer=dict(max_iter=16000),
),
flags={"allow_objects": True},
)

_JHU_H13_TEACHER_RUN = (
"cosmos_predict2p5_2B_action_conditioned_jhu_dvrk_mono_"
"finetune_13frame_8nodes_release_oss"
)
AC_CHUNK_SINGLE_VIEW_2B_JHU_DVRK_MONO_FINETUNE_13FRAME_8NODES_OSS_FINE_ANNEAL_4K = LazyDict(
dict(
defaults=[
f"/experiment/{_JHU_H13_TEACHER_RUN}",
"_self_",
],
job=dict(
group="official_runs_vid2vid",
name=f"{_JHU_H13_TEACHER_RUN}_fine_anneal_4k",
project="cosmos_predict2_action_conditioned",
),
checkpoint=dict(
load_path=_tabletop_teacher_checkpoint(_JHU_H13_TEACHER_RUN, 16000),
load_training_state=False,
strict_resume=False,
),
scheduler=L(LambdaWarmUpCosineScheduler)(
warm_up_steps=[100],
f_start=[0.10],
f_max=[1.00],
f_min=[0.05],
cycle_lengths=[4000],
),
trainer=dict(max_iter=4000),
),
flags={"allow_objects": True},
)

AC_CHUNK_SINGLE_VIEW_2B_JHU_DVRK_MONO_TABLETOP_H73_8NODES_OSS = LazyDict(
dict(
defaults=[
f"/experiment/{_JHU_H13_TEACHER_RUN}",
{"override /data_train": "jhu_dvrk_mono_finetune_h73_train"},
{"override /data_val": "jhu_dvrk_mono_finetune_h73_val"},
"_self_",
],
job=dict(
group="official_runs_vid2vid",
name=f"{_JHU_H13_TEACHER_RUN}_h73_tabletop",
project="cosmos_predict2_action_conditioned",
),
checkpoint=dict(
load_path=_tabletop_teacher_checkpoint(f"{_JHU_H13_TEACHER_RUN}_fine_anneal_4k", 4000),
load_training_state=False,
strict_resume=False,
),
model=dict(
config=dict(
state_t=1 + 72 // 4,
net=dict(action_dim=44),
),
),
dataloader_train=dict(batch_size=4),
optimizer=dict(lr=4e-5, weight_decay=0.1),
scheduler=L(LambdaWarmUpCosineScheduler)(
warm_up_steps=[1000],
f_start=[0.10],
f_max=[1.00],
f_min=[0.05],
cycle_lengths=[5000],
),
trainer=dict(max_iter=5000),
),
flags={"allow_objects": True},
)


cs = ConfigStore.instance()

Expand Down Expand Up @@ -1005,6 +1125,18 @@ def build_debug_runs(job):
AC_CHUNK_SINGLE_VIEW_2B_SUTUREBOT_13FRAME_NODES_OSS,
*build_debug_runs(AC_CHUNK_SINGLE_VIEW_2B_SUTUREBOT_13FRAME_NODES_OSS),
],
[
AC_CHUNK_SINGLE_VIEW_2B_JHU_DVRK_MONO_FINETUNE_13FRAME_8NODES_OSS,
*build_debug_runs(AC_CHUNK_SINGLE_VIEW_2B_JHU_DVRK_MONO_FINETUNE_13FRAME_8NODES_OSS),
],
[
AC_CHUNK_SINGLE_VIEW_2B_JHU_DVRK_MONO_FINETUNE_13FRAME_8NODES_OSS_FINE_ANNEAL_4K,
*build_debug_runs(AC_CHUNK_SINGLE_VIEW_2B_JHU_DVRK_MONO_FINETUNE_13FRAME_8NODES_OSS_FINE_ANNEAL_4K),
],
[
AC_CHUNK_SINGLE_VIEW_2B_JHU_DVRK_MONO_TABLETOP_H73_8NODES_OSS,
*build_debug_runs(AC_CHUNK_SINGLE_VIEW_2B_JHU_DVRK_MONO_TABLETOP_H73_8NODES_OSS),
],
]:
cs.store(group="experiment", package="_global_", name=f"{_item['job']['name']}", node=_item)
if _item_wo_resume is not None:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1457,13 +1457,22 @@ def __init__(
self,
*args,
data_split="full",
test_split_ratio: float = 0.05,
modality_filename: str | None = None,
exclude_splits: list[str] | None = None,
**kwargs,
):
"""Wrap ``LeRobotSingleDataset`` with a deterministic train/test split.

The split is taken from the trailing end of ``_all_steps``.
``data_split="full"`` skips partitioning entirely.
"""
if not 0.0 < test_split_ratio < 1.0:
raise ValueError(f"test_split_ratio must be in (0, 1), got {test_split_ratio}")
# Store data_split BEFORE calling super().__init__() because
# _get_all_steps_cmr_filtered needs it for cache path generation
self.data_split = data_split
self.test_split_ratio = test_split_ratio
super().__init__(
*args,
modality_filename=modality_filename,
Expand All @@ -1474,11 +1483,16 @@ def __init__(
if data_split == "full":
pass
elif data_split == "train":
self._all_steps = self._all_steps[: -len(self) // 20]
n_test = max(1, int(len(self) * test_split_ratio))
self._all_steps = self._all_steps[:-n_test]
elif data_split == "test":
self._all_steps = self._all_steps[-len(self) // 20 :]
n_test = max(1, int(len(self) * test_split_ratio))
self._all_steps = self._all_steps[-n_test:]

print(f"Dataset is split into {data_split} data, with {len(self._all_steps)} steps.")
print(
f"Dataset is split into {data_split} data (test_split_ratio={test_split_ratio:.4f}), "
f"with {len(self._all_steps)} steps."
)

def _get_trajectories(self) -> tuple[np.ndarray, np.ndarray]:
"""Get the trajectories in the dataset."""
Expand Down Expand Up @@ -1730,11 +1744,19 @@ class MixedLeRobotDataset(torch.utils.data.Dataset):
- ``embodiment`` (str): Embodiment tag string (must be in
EMBODIMENT_REGISTRY or be one of the built-in embodiments).
- ``mix_ratio`` (float, optional): Relative sampling weight. Default 1.0.
- ``data_split_override`` (str, optional): Per-spec override of the
global ``data_split``.
- ``test_split_ratio_override`` (float, optional): Per-spec override
of the global ``test_split_ratio``.
- ``exclude_splits`` (list[str], optional): Episode-level split names
from ``meta/info.json`` to exclude.
num_frames: Number of video frames per sample (e.g. 13 = 1 context + 12 pred).
data_split: One of ``"train"``, ``"test"``, ``"full"``.
data_split: One of ``"train"``, ``"test"``, ``"full"``. Applied to every
spec unless overridden by ``data_split_override``.
max_action_dim: All action tensors are zero-padded to this dimension.
Default 44 (CMR Versius conditioning dimension).
downscaled_res: If True, use 256x256 resolution for all videos.
test_split_ratio: Default held-out fraction for each sub-dataset.

Example::

Expand All @@ -1752,6 +1774,7 @@ def __init__(
data_split: str = "train",
max_action_dim: int = 44,
downscaled_res: bool = False,
test_split_ratio: float = 0.05,
):
from cosmos_predict2._src.predict2.action.datasets.gr00t_dreams.groot_configs import (
construct_modality_config_and_transforms,
Expand All @@ -1775,7 +1798,13 @@ def __init__(
embodiment = raw_embodiment.value if isinstance(raw_embodiment, EmbodimentTag) else raw_embodiment
mix_ratio = spec.get("mix_ratio", 1.0)

print(f"\n[{i}] Loading: embodiment={embodiment}, mix_ratio={mix_ratio}")
spec_data_split = spec.get("data_split_override", data_split)
spec_test_split_ratio = spec.get("test_split_ratio_override", test_split_ratio)

print(
f"\n[{i}] Loading: embodiment={embodiment}, mix_ratio={mix_ratio}, "
f"data_split={spec_data_split}, test_split_ratio={spec_test_split_ratio}"
)
print(f" path={path}")

config, train_transform, test_transform = construct_modality_config_and_transforms(
Expand All @@ -1789,7 +1818,7 @@ def __init__(
if isinstance(config, dict) and "modality_filename" in config:
modality_filename = config.pop("modality_filename")

transform = train_transform if data_split in ("train", "full") else test_transform
transform = train_transform if spec_data_split in ("train", "full") else test_transform

# Per-dataset episode filtering (e.g., exclude "fail", "bad_frames" splits)
exclude_splits = spec.get("exclude_splits", None)
Expand All @@ -1799,7 +1828,8 @@ def __init__(
modality_configs=config,
transforms=transform,
embodiment_tag=embodiment,
data_split=data_split,
data_split=spec_data_split,
test_split_ratio=spec_test_split_ratio,
modality_filename=modality_filename,
exclude_splits=exclude_splits,
)
Expand Down
Loading
Loading