Skip to content
Open
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
35 changes: 33 additions & 2 deletions examples/visual_gen/models/cosmos3/README.md
Original file line number Diff line number Diff line change
@@ -1,11 +1,12 @@
# Cosmos3 Text(+Image)-to-Video(+Audio) generation
# Cosmos3 Text(+Image)-to-Video(+Audio) and action generation

Cosmos3 supports four generation modes from a single checkpoint:
Cosmos3 supports the following generation modes from a single checkpoint:

- **T2V** — text-to-video (`prompts/t2v.json`).
- **T2I** — text-to-image (`prompts/t2i.json`); emits a still frame (use `--output_type image` / a non-video `--output_path`).
- **I2V / TI2V** — image-conditioned video (`prompts/i2v.json`). Condition on a reference frame via the prompt file's `vision_path` or `--image_path`. The image may be a local path, a `file://` / `http(s)://` URL, or a `data:` URI.
- **T2AV** — text-to-video with synchronized audio (`prompts/t2av.json` with `enable_audio: true`, or pass `--enable_audio`). Combine with a `vision_path` for image-conditioned audio-video (TI2AV).
- **Action** — policy / forward dynamics / inverse dynamics generation (pass `--action_mode`). Action and audio generation are mutually exclusive. Save action runs as `safetensors` or `pt` so the rollout and action tensor stay in one payload.

## Checkpoints

Expand Down Expand Up @@ -74,4 +75,34 @@ python cosmos3.py --model nvidia/Cosmos3-Nano \
python cosmos3.py --model nvidia/Cosmos3-Nano \
--prompt "A cute puppy playing with a ball in a park" \
--visual_gen_args ../configs/cosmos3-nano-1gpu.yaml

# Action — policy (first frame + instruction -> predicted action + rollout video)
python cosmos3.py --model nvidia/Cosmos3-Nano \
--prompt_file prompts/action_policy.json \
--visual_gen_args ../configs/cosmos3-nano-1gpu.yaml \
--action_mode policy \
--domain_name bridge_orig_lerobot \
--raw_action_dim 10 \
--output_path policy_rollout.safetensors \
--action_output_path policy_action.json

# Action — forward dynamics (first frame + action trajectory -> rollout video)
python cosmos3.py --model nvidia/Cosmos3-Nano \
--prompt_file prompts/action_forward_dynamics.json \
--visual_gen_args ../configs/cosmos3-nano-1gpu.yaml \
--action_mode forward_dynamics \
--domain_name av \
--action_json action_trajectory.json \
--output_path forward_dynamics.safetensors

# Action — inverse dynamics (video -> predicted action)
python cosmos3.py --model nvidia/Cosmos3-Nano \
--prompt_file prompts/action_inverse_dynamics.json \
--video_path /path/to/frames/ \
--visual_gen_args ../configs/cosmos3-nano-1gpu.yaml \
--action_mode inverse_dynamics \
--domain_name bridge_orig_lerobot \
--raw_action_dim 10 \
--output_path inverse_video.safetensors \
--action_output_path inverse_action.json
```
188 changes: 186 additions & 2 deletions examples/visual_gen/models/cosmos3/cosmos3.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,8 @@
from tensorrt_llm import VisualGen, VisualGenArgs

_SCRIPT_DIR = Path(__file__).resolve().parent
_ACTION_MODES = ("policy", "forward_dynamics", "inverse_dynamics")
_TENSOR_OUTPUT_SUFFIXES = {".pt", ".safetensors"}


def _resolve_path(path: str) -> str:
Expand Down Expand Up @@ -80,6 +82,91 @@ def resolve_prompt_and_options(
return resolved_prompt, resolved_image, resolved_enable_audio, resolved_output_type


def _validate_action_args(
args: argparse.Namespace, resolved_image_path: Optional[str] = None
) -> None:
if args.action_mode is None:
return

# The first frame may come from --image_path or a prompt file's vision_path.
has_first_frame = resolved_image_path is not None or args.video_path is not None

mode = args.action_mode.strip().lower()
if mode not in _ACTION_MODES:
raise SystemExit(
f"Invalid --action_mode {args.action_mode!r}; expected one of {list(_ACTION_MODES)}."
)
args.action_mode = mode
if args.enable_audio:
raise SystemExit("Cosmos3 does not support joint action and audio generation.")
if args.output_type != "video":
raise SystemExit("Action generation requires --output_type video.")

if mode == "forward_dynamics":
if args.action_json is None:
raise SystemExit(f"{mode} requires --action_json.")
if not has_first_frame:
raise SystemExit(
f"{mode} requires --image_path, a prompt-file vision_path, or --video_path "
"for the first frame."
)
elif mode == "policy":
if not has_first_frame:
raise SystemExit(
f"{mode} requires --image_path, a prompt-file vision_path, or --video_path "
"for the first frame."
)
if args.raw_action_dim is None and args.domain_name is None and args.domain_id is None:
raise SystemExit(f"{mode} requires --raw_action_dim, --domain_name, or --domain_id.")
elif mode == "inverse_dynamics":
if args.video_path is None:
raise SystemExit(
f"{mode} requires --video_path (frame directory, .mp4/.avi, or image)."
)
if args.raw_action_dim is None and args.domain_name is None and args.domain_id is None:
raise SystemExit(f"{mode} requires --raw_action_dim, --domain_name, or --domain_id.")


def _resolved_output_path(path: str, action_mode: Optional[str]) -> str:
if action_mode is None:
return path
output_path = Path(path)
if output_path.suffix.lower() in _TENSOR_OUTPUT_SUFFIXES:
return str(output_path)
return str(output_path.with_suffix(".safetensors"))


def _default_action_output_path(output_path: str) -> str:
stem = Path(output_path)
return str(stem.with_suffix(".action.json"))


def _save_action_output(output, path: str) -> None:
if output.action is None:
return

action = output.action
if action.ndim == 3 and action.shape[0] == 1:
action_data = action[0].tolist()
shape = list(action.shape[1:])
else:
action_data = action.tolist()
shape = list(action.shape)

payload = {
"action_mode": output.action_mode,
"domain_id": output.domain_id,
"raw_action_dim": output.raw_action_dim,
"shape": shape,
"dtype": str(action.dtype).replace("torch.", ""),
"data": action_data,
}
out_path = Path(path)
out_path.parent.mkdir(parents=True, exist_ok=True)
with out_path.open("w", encoding="utf-8") as f:
json.dump(payload, f, indent=2)


def main():
parser = argparse.ArgumentParser(description="Cosmos3 Text(+Image)-to-Video(+Audio) example")
parser.add_argument(
Expand Down Expand Up @@ -139,6 +226,68 @@ def main():
"--use_system_prompt", action="store_true", help="Use system prompt in prompt"
)
parser.add_argument("--enable_audio", action="store_true", help="Enable audio generation")
parser.add_argument(
"--action_mode",
type=str,
default=None,
choices=list(_ACTION_MODES),
help="Action mode: policy, forward_dynamics, or inverse_dynamics",
)
parser.add_argument(
"--domain_name",
type=str,
default=None,
help="Embodiment domain name (e.g. bridge_orig_lerobot, av, droid_lerobot)",
)
parser.add_argument(
"--domain_id",
type=int,
default=None,
help="Embodiment domain id (alternative to --domain_name)",
)
parser.add_argument(
"--raw_action_dim",
type=int,
default=None,
help="Raw action DOF for policy/inverse_dynamics",
)
parser.add_argument(
"--action_chunk_size",
type=int,
default=None,
help="Action tokens to generate. Defaults to the domain preset or model default.",
)
parser.add_argument(
"--action_json",
type=str,
default=None,
help="JSON file with action trajectory [T, D] for forward_dynamics",
)
parser.add_argument(
"--video_path",
type=str,
default=None,
help="Frame directory, .mp4/.avi video, or image path for inverse_dynamics",
)
parser.add_argument(
"--action_resolution",
type=int,
default=None,
choices=[256, 480, 704, 720],
help=("Resolution bucket for action image sizing. Defaults to the domain preset or 480."),
)
parser.add_argument(
"--action_fps",
type=float,
default=None,
help="Action-token temporal rate for mRoPE (Hz). Defaults to frame_rate.",
)
parser.add_argument(
"--action_output_path",
type=str,
default=None,
help="Path to save predicted action JSON (default: <output_stem>_action.json)",
)
parser.add_argument(
"--output_type", type=str, default="video", help="Output type (video, image)"
)
Expand All @@ -156,6 +305,7 @@ def main():
enable_audio=args.enable_audio,
output_type=args.output_type,
)
_validate_action_args(args, resolved_image_path=image_path)

# Engine config from shared YAML (optional); model-specific defaults apply otherwise.
extra_args = VisualGenArgs.from_yaml(args.visual_gen_args) if args.visual_gen_args else None
Expand All @@ -166,6 +316,8 @@ def main():
params = visual_gen.default_params
if image_path is not None:
params.image = image_path
if args.action_mode is not None and args.video_path is not None:
params.image = args.video_path

negative_prompt_path = _resolve_path(args.negative_prompt)
if args.negative_prompt is not None:
Expand All @@ -186,6 +338,24 @@ def main():
params.extra_params["use_guardrails"] = not args.disable_guardrails
params.extra_params["output_type"] = output_type

if args.action_mode is not None:
params.extra_params["action_mode"] = args.action_mode
if args.domain_name is not None:
params.extra_params["domain_name"] = args.domain_name
if args.domain_id is not None:
params.extra_params["domain_id"] = args.domain_id
if args.raw_action_dim is not None:
params.extra_params["raw_action_dim"] = args.raw_action_dim
if args.action_chunk_size is not None:
params.extra_params["action_chunk_size"] = args.action_chunk_size
if args.action_resolution is not None:
params.extra_params["action_resolution"] = args.action_resolution
if args.action_fps is not None:
params.extra_params["action_fps"] = args.action_fps
if args.action_json is not None:
with open(args.action_json, encoding="utf-8") as f:
params.extra_params["action"] = json.load(f)

if negative_prompt is None:
params.negative_prompt = None
elif isinstance(negative_prompt, str):
Expand All @@ -198,8 +368,22 @@ def main():
params=params,
)

output.save(args.output_path)
print(f"Saved: {args.output_path}")
output_path = _resolved_output_path(args.output_path, args.action_mode)
output.save(output_path)
print(f"Saved: {output_path}")

if args.action_mode is not None:
action_path = args.action_output_path or _default_action_output_path(output_path)
_save_action_output(output, action_path)
if output.action is not None:
print(f"Saved action: {action_path}")
print(
f"Action shape: {tuple(output.action.shape)}, "
f"raw_action_dim={output.raw_action_dim}, domain_id={output.domain_id}"
)
else:
print("Warning: action_mode was set but the output carried no action tensor.")

print(output.metrics)


Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
{
"model_mode": "image2video",
"prompt": "Robot manipulation rollout: the right arm extends to the fruit display, picks up a pear, and places it into the bag in the shopping cart.",
"vision_path": "https://github.com/nvidia-cosmos/cosmos-dependencies/raw/refs/heads/assets/cosmos3/inputs/vision/robot_153.jpg"
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,4 @@
{
"model_mode": "image2video",
"prompt": "Recover the robot action trajectory from this clip of the arm picking up a pear and placing it into the bag."
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
{
"model_mode": "image2video",
"prompt": "Pick up the pear from the fruit display and place it into the plastic bag in the shopping cart.",
"vision_path": "https://github.com/nvidia-cosmos/cosmos-dependencies/raw/refs/heads/assets/cosmos3/inputs/vision/robot_153.jpg"
}
Loading
Loading