Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
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
@@ -0,0 +1,61 @@
_target_: egomimic.pl_utils.pl_data_utils.MultiDataModuleWrapper

# The source root stays physically complete at 720 generated episodes. This
# recipe selects 719 by excluding the audited 3,118-frame idle-heavy episode.
train_datasets:
pushshapes_sim_chain_gripper:
_target_: egomimic.rldb.zarr.zarr_dataset_multi.MultiDataset._from_resolver
resolver:
_target_: egomimic.rldb.zarr.zarr_dataset_multi.LocalEpisodeResolverManyWithEmbodimentOverride
folder_paths: &chain_roots
- /coc/flash7/paphiwetsa3/datasets/Tsim_v2/chain_gripper_3000_v2
- /coc/flash7/paphiwetsa3/datasets/Tsim_v2/chain_gripper_gen
embodiment_override: pushshapes_sim_chain_gripper
key_map:
_target_: egomimic.rldb.embodiment.pushshapes.get_keymap_hpt
action_horizon: 16
action_zarr_key: actions
transform_list:
_target_: egomimic.rldb.embodiment.pushshapes.get_chain_gripper_point_transform_list
world_size: 512.0
filters: &chain_episode_filter
_target_: egomimic.rldb.filters.DatasetFilter
filter_lambdas:
- "lambda row: row.get('episode_hash') != 'episode_T_chain_gripper_obs7_000050'"
mode: train
valid_ratio: 0.0
bounds_check: false

valid_datasets:
pushshapes_sim_chain_gripper:
_target_: egomimic.rldb.zarr.zarr_dataset_multi.MultiDataset._from_resolver
resolver:
_target_: egomimic.rldb.zarr.zarr_dataset_multi.LocalEpisodeResolverManyWithEmbodimentOverride
folder_paths: *chain_roots
embodiment_override: pushshapes_sim_chain_gripper
key_map:
_target_: egomimic.rldb.embodiment.pushshapes.get_keymap_hpt
action_horizon: 16
action_zarr_key: actions
transform_list:
_target_: egomimic.rldb.embodiment.pushshapes.get_chain_gripper_point_transform_list
world_size: 512.0
filters: *chain_episode_filter
mode: valid
valid_ratio: 0.02
bounds_check: false

# One A40 per BC run; batch 64 is the global batch for embodiment 20.
train_dataloader_params:
pushshapes_sim_chain_gripper:
batch_size: 64
num_workers: 4
pin_memory: true
persistent_workers: true
prefetch_factor: 2

valid_dataloader_params:
pushshapes_sim_chain_gripper:
batch_size: 16
num_workers: 4
pin_memory: true
Original file line number Diff line number Diff line change
@@ -0,0 +1,62 @@
# @package _global_
defaults:
- override /model: bf/bf_pipeline_diffusion_chain_gripper_points
- override /data: pusht/chain_gripper_newdata_h16_points_full
- override /evaluator: eval_pipeline_action_mse
- override /logger: wandb

# ChainGripper-only BC over the frozen base-plus-generated corpus.
name: flow_transfer_bc_chain_newdata_dp_h16
description: standard_dp_epsilon_h16_chain_3000v2_plus_gen
model:
enable_grad_norm: false
# Scale the complete original 1e-4 -> 1e-5 schedule by 0.3.
optimizer:
lr: 3.0e-5
scheduler:
eta_min: 3.0e-6
train_metrics_on_step: true
train_metrics_on_epoch: true
launch_params:
gpus_per_node: 1
nodes: 1
trainer:
max_steps: 240000
max_epochs: -1
min_epochs: null
limit_train_batches: 1.0
val_check_interval: 10000
limit_val_batches: 0
check_val_every_n_epoch: 1
num_sanity_val_steps: 0
accumulate_grad_batches: 1
log_every_n_steps: 1
logger:
wandb: {project: pushshapes-flow-transfer}
norm_stats:
norm_mode: minmax
reduce_all_but_last: true
callbacks:
model_checkpoint:
every_n_epochs: null
every_n_train_steps: null
train_time_interval: {_target_: datetime.timedelta, hours: 1}
save_top_k: 1
save_last: true
save_on_train_epoch_end: true
terminal_checkpoint:
_target_: lightning.pytorch.callbacks.ModelCheckpoint
dirpath: ${paths.output_dir}/checkpoints/final
filename: "step-{step}"
monitor: null
verbose: true
every_n_epochs: null
every_n_train_steps: ${trainer.max_steps}
train_time_interval: null
save_top_k: 1
save_last: false
save_on_train_epoch_end: false
save_on_exception: false
save_weights_only: false
auto_insert_metric_name: false
enable_version_counter: false
Original file line number Diff line number Diff line change
@@ -0,0 +1,62 @@
# @package _global_
defaults:
- override /model: bf/bf_pipeline_sampler_chain_gripper_points_dense_medium_h16
- override /data: pusht/chain_gripper_newdata_h16_points_full
- override /evaluator: eval_pipeline_action_mse
- override /logger: wandb

# ChainGripper-only BC over the frozen base-plus-generated corpus.
name: flow_transfer_bc_chain_newdata_latent_dense_medium_h16
description: decoder_only_latent96_dense16_chain_3000v2_plus_gen
model:
enable_grad_norm: false
# Scale the complete original 1e-4 -> 1e-5 schedule by 0.3.
optimizer:
lr: 3.0e-5
scheduler:
eta_min: 3.0e-6
train_metrics_on_step: true
train_metrics_on_epoch: true
launch_params:
gpus_per_node: 1
nodes: 1
trainer:
max_steps: 240000
max_epochs: -1
min_epochs: null
limit_train_batches: 1.0
val_check_interval: 10000
limit_val_batches: 0
check_val_every_n_epoch: 1
num_sanity_val_steps: 0
accumulate_grad_batches: 1
log_every_n_steps: 1
logger:
wandb: {project: pushshapes-flow-transfer}
norm_stats:
norm_mode: minmax
reduce_all_but_last: true
callbacks:
model_checkpoint:
every_n_epochs: null
every_n_train_steps: null
train_time_interval: {_target_: datetime.timedelta, hours: 1}
save_top_k: 1
save_last: true
save_on_train_epoch_end: true
terminal_checkpoint:
_target_: lightning.pytorch.callbacks.ModelCheckpoint
dirpath: ${paths.output_dir}/checkpoints/final
filename: "step-{step}"
monitor: null
verbose: true
every_n_epochs: null
every_n_train_steps: ${trainer.max_steps}
train_time_interval: null
save_top_k: 1
save_last: false
save_on_train_epoch_end: false
save_on_exception: false
save_weights_only: false
auto_insert_metric_name: false
enable_version_counter: false
Original file line number Diff line number Diff line change
@@ -0,0 +1,81 @@
# Flow Transfer ChainGripper Medium direct-dense H16 latent denoiser. Native4
# commands become ordered points6 in the loader; the action path is decoder-only.
_target_: egomimic.pl_utils.pl_model.ModelWrapper

robomimic_model:
_target_: egomimic.pipeline.algo.PipelineAlgo
action_horizon: 16
domains: [pushshapes_sim_chain_gripper]
ac_keys: {pushshapes_sim_chain_gripper: actions}
rollout_adapters:
pushshapes_sim_chain_gripper:
_target_: egomimic.pipeline.pushshapes.ChainGripperPointRolloutAdapter
action_horizon: 16
stages:
- _target_: egomimic.pipeline.stages_sampler.FusedObsEncoder
n_obs_steps: 1
encoder:
_target_: egomimic.pipeline.stages_sampler.DPStyleObsEncoder
obs_specs:
state_agent_obj: {input_dim: 3, input_slice: [0, 3]}
img_encoders:
front_img_1:
_target_: egomimic.models.stems.visual_core.VisualCore
in_channels: 3
image_size: 96
num_kp: 32
feature_dimension: 64
pretrained: false
crop_aug: true
crop_height: 84
crop_width: 84
crop_eval_mode: center
crop_sample_mode: v02
crop_scope: frame
norm_layer: group
pool_type: spatial_softmax
- _target_: egomimic.pipeline.stages_sampler.GaussianLatentNoise
action_horizon: 16
latent_dim: 96
- _target_: egomimic.pipeline.stages_sampler.MultiJActionSampler
condition_input_dim: 67
condition_dim: 384
gradient_accumulation_steps: 1
schedule_anchor_domain: pushshapes_sim_chain_gripper
action_horizon: 16
action_dims: {pushshapes_sim_chain_gripper: 6}
latent_dim: 96
decoder_hidden_dim: 512
denoiser_hidden_dim: 384
num_inference_steps: 16
sampling_schedule:
1: {1: 0.50, 2: 0.50}
2001: {2: 0.80, 4: 0.15, 8: 0.05}
gradient_checkpointing: true
denoising_module:
_target_: egomimic.models.denoising_nets.CrossTransformer
nblocks: 16
cond_dim: 384
hidden_dim: 384
act_dim: 96
act_seq: 16
n_heads: 8
dropout: 0.1
mlp_layers: 4
mlp_ratio: 4
time_conditioning: additive
- _target_: egomimic.pipeline.stages_sampler.NativeActionMSELoss

enable_grad_norm: false
optimizer:
_target_: torch.optim.AdamW
_partial_: true
lr: 1.0e-4
weight_decay: 1.0e-4
scheduler:
_target_: egomimic.utils.schedulers.warmup_cosine_scheduler
_partial_: true
max_steps: 240000
warmup_steps: 3000
warmup_start_factor: 0.1
eta_min: 1.0e-5
11 changes: 10 additions & 1 deletion egomimic/scripts/verify_training_smoke.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,8 @@
from wandb.proto import wandb_internal_pb2
from wandb.sdk.internal.datastore import DataStore

import egomimic.utils.hydra_resolvers # noqa: F401


def _sha256(path: Path) -> str:
digest = hashlib.sha256()
Expand Down Expand Up @@ -157,6 +159,13 @@ def _has_required_metrics(
return True


def _load_training_config(config_path: Path):
"""Load a Hydra snapshot with the same resolvers used by trainHydra."""
if not OmegaConf.has_resolver("eval"):
OmegaConf.register_new_resolver("eval", eval)
return OmegaConf.load(config_path)


def verify_training_smoke(
output_dir: Path,
required_embodiments: list[int],
Expand All @@ -166,7 +175,7 @@ def verify_training_smoke(
output_dir = output_dir.resolve()
config_path = output_dir / ".hydra" / "config.yaml"
assert config_path.is_file(), config_path
config = OmegaConf.load(config_path)
config = _load_training_config(config_path)

assert int(config.trainer.max_steps) == 2
assert int(config.trainer.limit_train_batches) == 2
Expand Down
Loading