Add char_masks preprocessing script for SCAIL-2 training (task 3.7)
scripts/process_char_masks.py turns per-sample character label-map videos/images
(integer pixel labels: 0 = environment, 1..K = characters -> binding slots) into
the pixel-space semantic-mask tensors the SCAIL training/inference path consumes:
{"mask": [K+1, F_pix, H_pix, W_pix]} (ch0 = environment switch, ch1..K = slots).
- Aligns to the target video's latent grid read from the saved latent metadata
(F_pix=(F-1)*8+1, H*32, W*32), so char_masks/ lines up file-for-file with
latents/ / driving_latents/ for PrecomputedDataset.
- Nearest-neighbour resize so integer labels are never blended; labels > K are
dropped with a warning; ch0 filled uniformly with --environment-switch.
- Reuses process_videos.py helpers (naming, atomic save, VAE factors) and matches
its typer CLI conventions.
Verified on CPU: a synthetic 2-character label map (plus an out-of-range id)
produces mask (7,17,128,128) with ch0 uniform, slots placed correctly, id>K
dropped, and feeds encode_mask_channels to the 8*(K+1)=56 channels. README +
docs/tasks.md 3.7 updated (upstream label-map generation via SAM/tracking is
dataset-specific and still out of scope).
Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
+1
-1
@@ -40,7 +40,7 @@
|
||||
| 3.4 | LoRA 模式解凍 patchify_proj | ✅ | `trainer._unfreeze_patchify_proj`:mask_channels>0 時把 video patchify_proj 設 trainable,讓新欄位隨 LoRA 一起訓 |
|
||||
| 3.5 | 範例 config + docs | ✅ | `configs/scail_animation_lora.yaml`;`configs/README.md` 與 `docs/training-modes.md` 表格 row |
|
||||
| 3.6 | CPU 單元驗證 | ✅ | `verify_phase3_trainer.py`:config round-trip、prepare_training_inputs 建 cond_channels[B,T,56]、driving 前置/mask placement、widened model forward + compute_loss finite、widen helper zero-init 等價 |
|
||||
| 3.7 | dataset 前處理(產 driving latents + 語意 mask) | ⬜ | 需 process_dataset 產出 `driving_latents/`(同 target 形狀)與 `char_masks/`(mask=[K+1,F_pix,H,W]);語意 mask 需分割模型,屬資料工程,未做 |
|
||||
| 3.7 | dataset 前處理(產 driving latents + 語意 mask) | 🟡 | `char_masks/` 前處理已做:`scripts/process_char_masks.py`(label-map 影片/圖 → `[K+1,F_pix,H,W]`,對齊 target latent,nearest 保留整數 label,ch0 環境開關)。`driving_latents/` 沿用既有 `process_videos.py`(driving 影片走 video latent 路徑,同 target 形狀)。**仍缺**:從原始影片產生 label-map 的分割/追蹤步驟(SAM 等,資料集特定,未含) |
|
||||
| 3.8 | validation runner 接 driving/mask | ⬜ | 驗證期取樣尚未接 SCAIL 條件(config 內 validation 先停用),與 Phase 4 一起 |
|
||||
| 3.9 | 實機訓練跑通 | ⬜ | 需 Linux + GPU + checkpoint,本機無法 |
|
||||
|
||||
|
||||
+12
-2
@@ -86,8 +86,18 @@ char_masks/ # per-sample "mask" = [K+1, F_pix, H, W] (ch0 env switch, 1.
|
||||
|
||||
`latents/`, `conditions/`, and `driving_latents/` come from the existing
|
||||
`packages/ltx-trainer/scripts/process_dataset.py` (run it once per video set).
|
||||
`char_masks/` still needs a segmentation step (e.g. SAM) to produce the semantic
|
||||
masks — that preprocessing is not implemented yet (see `docs/tasks.md` 3.7).
|
||||
`char_masks/` is produced by `packages/ltx-trainer/scripts/process_char_masks.py`
|
||||
from per-sample **label-map** videos/images (integer pixel labels: `0` =
|
||||
environment, `1..K` = characters → binding slots):
|
||||
|
||||
```bash
|
||||
python packages/ltx-trainer/scripts/process_char_masks.py dataset.csv \
|
||||
--mask-column char_labels --latents-dir ./latents \
|
||||
--output-dir ./char_masks --num-slots 6 --main-media-column media_path
|
||||
```
|
||||
|
||||
You still need a segmentation/tracking model (e.g. SAM) to *produce* those label
|
||||
maps from raw video — that upstream step is dataset-specific and not included.
|
||||
|
||||
## Cost notes
|
||||
|
||||
|
||||
@@ -0,0 +1,217 @@
|
||||
#!/usr/bin/env python3
|
||||
|
||||
"""Preprocess SCAIL-2 character binding-slot masks into per-sample ``.pt`` tensors.
|
||||
|
||||
Each sample's input is a **label-map** video or image whose integer pixel values index
|
||||
characters: ``0`` = environment/background, ``k`` in ``1..K`` = character *k* (assigned to
|
||||
binding slot *k*). This script aligns the label map to the target video's latent grid (read
|
||||
from the saved video-latent metadata), and emits the pixel-space semantic-mask tensor that
|
||||
``ltx_core.conditioning.encode_mask_channels`` (and the trainer's ``mask_channels`` condition)
|
||||
consumes:
|
||||
|
||||
{"mask": tensor[K+1, F_pix, H_pix, W_pix]} # ch0 = environment switch, ch1..K = binding slots
|
||||
|
||||
where ``F_pix = (latent_frames - 1) * 8 + 1``, ``H_pix = latent_h * 32``, ``W_pix = latent_w * 32``
|
||||
(the SCAIL encoder then spatially downsamples + temporally stacks these to the 8*(K+1)=56 channels).
|
||||
|
||||
Label maps are resized with **nearest-neighbour** interpolation so integer labels are never
|
||||
blended. The environment-switch channel (ch0) is filled uniformly with ``--environment-switch``
|
||||
(``0.0`` = derive the environment from the reference image, ``1.0`` = from the driving video),
|
||||
matching the paper's single-bit environment signal.
|
||||
|
||||
Output naming mirrors the target latents (same relative path), so ``char_masks/`` lines up
|
||||
file-for-file with ``latents/`` / ``driving_latents/`` for the trainer's ``PrecomputedDataset``.
|
||||
|
||||
Standalone usage::
|
||||
|
||||
python scripts/process_char_masks.py dataset.csv \\
|
||||
--mask-column char_labels --latents-dir ./latents --output-dir ./char_masks --num-slots 6
|
||||
"""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import typer
|
||||
from PIL import Image
|
||||
|
||||
# Sibling scripts (resolved via scripts/ on sys.path), reused so naming + alignment match the other stages.
|
||||
from process_videos import (
|
||||
IMAGE_FILE_EXTENSIONS,
|
||||
VAE_SPATIAL_FACTOR,
|
||||
VAE_TEMPORAL_FACTOR,
|
||||
_atomic_save,
|
||||
_load_paths_from_dataset,
|
||||
_output_relative,
|
||||
)
|
||||
from rich.console import Console
|
||||
from rich.progress import BarColumn, MofNCompleteColumn, Progress, SpinnerColumn, TextColumn, TimeElapsedColumn
|
||||
|
||||
from ltx_trainer import logger
|
||||
from ltx_trainer.video_utils import read_video
|
||||
|
||||
app = typer.Typer(
|
||||
pretty_exceptions_enable=False,
|
||||
no_args_is_help=True,
|
||||
help="Preprocess SCAIL-2 character label maps into [K+1, F_pix, H_pix, W_pix] mask tensors.",
|
||||
)
|
||||
|
||||
|
||||
def _load_label_frames(mask_file: Path, pixel_f: int) -> torch.Tensor:
|
||||
"""Load a label-map video/image as integer labels ``[F_pix, H, W]`` (float-typed integers).
|
||||
|
||||
Images are tiled across ``pixel_f`` frames. Videos are read via the shared ``read_video`` helper
|
||||
(which returns ``[F, C, H, W]`` in ``[0, 1]``); labels are recovered as ``round(value * 255)``,
|
||||
so label maps must be stored as small-integer grayscale (label ``k`` -> pixel value ``k``).
|
||||
"""
|
||||
if mask_file.suffix.lower() in IMAGE_FILE_EXTENSIONS:
|
||||
arr = np.array(Image.open(mask_file).convert("L")) # exact integer labels [H, W] (writable copy)
|
||||
labels = torch.from_numpy(arr).float()
|
||||
return labels.unsqueeze(0).expand(pixel_f, -1, -1).contiguous()
|
||||
|
||||
frames, _ = read_video(str(mask_file), max_frames=pixel_f) # [F, C, H, W] in [0, 1]
|
||||
labels = frames[:, 0].mul(255.0).round() # channel 0 -> integer labels [F, H, W]
|
||||
if labels.shape[0] < pixel_f:
|
||||
# Pad by repeating the last frame so every latent frame has a label map.
|
||||
pad = labels[-1:].expand(pixel_f - labels.shape[0], -1, -1)
|
||||
labels = torch.cat([labels, pad], dim=0)
|
||||
return labels[:pixel_f]
|
||||
|
||||
|
||||
def _labels_to_slot_masks(labels: torch.Tensor, num_slots: int, environment_switch: float) -> torch.Tensor:
|
||||
"""Convert integer label frames ``[F, H, W]`` to ``[K+1, F, H, W]`` (ch0 env switch, ch1..K slots)."""
|
||||
f_pix, h, w = labels.shape
|
||||
out = torch.zeros(num_slots + 1, f_pix, h, w, dtype=torch.float32)
|
||||
out[0] = environment_switch # uniform environment-switch channel
|
||||
for k in range(1, num_slots + 1):
|
||||
out[k] = (labels == k).float()
|
||||
if bool(((labels > num_slots) & (labels > 0)).any().item()):
|
||||
logger.warning(
|
||||
f"Label map contains ids > num_slots ({num_slots}); those pixels are dropped (treated as environment)."
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
def _resize_labels(labels: torch.Tensor, pixel_h: int, pixel_w: int) -> torch.Tensor:
|
||||
"""Nearest-neighbour resize of integer label frames ``[F, H, W]`` to ``[F, pixel_h, pixel_w]``."""
|
||||
if labels.shape[1:] == (pixel_h, pixel_w):
|
||||
return labels
|
||||
return torch.nn.functional.interpolate(
|
||||
labels.unsqueeze(1), size=(pixel_h, pixel_w), mode="nearest"
|
||||
).squeeze(1)
|
||||
|
||||
|
||||
def compute_char_masks(
|
||||
dataset_file: str | Path,
|
||||
mask_column: str,
|
||||
latents_dir: str,
|
||||
output_dir: str,
|
||||
num_slots: int = 6,
|
||||
environment_switch: float = 0.0,
|
||||
main_media_column: str | None = None,
|
||||
overwrite: bool = False,
|
||||
) -> None:
|
||||
"""Preprocess character label maps into ``[K+1, F_pix, H_pix, W_pix]`` mask tensors.
|
||||
|
||||
Args:
|
||||
dataset_file: Metadata file (CSV/JSON/JSONL) with a column of label-map paths.
|
||||
mask_column: Column containing the per-sample label-map video/image paths.
|
||||
latents_dir: Directory of target video latents (read for spatial/temporal alignment).
|
||||
output_dir: Directory to write ``char_masks`` ``.pt`` files.
|
||||
num_slots: Number of binding slots K (output has ``K+1`` channels).
|
||||
environment_switch: Uniform value for ch0 (0.0 = env from reference, 1.0 = from driving video).
|
||||
main_media_column: Column used for output naming (defaults to ``mask_column``); set it to the
|
||||
target-video column so masks align with ``latents/`` when label maps live elsewhere.
|
||||
overwrite: Recompute even if the output already exists.
|
||||
"""
|
||||
dataset_path = Path(dataset_file)
|
||||
data_root = dataset_path.parent
|
||||
latents_path = Path(latents_dir)
|
||||
output_path = Path(output_dir)
|
||||
output_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
naming_column = main_media_column or mask_column
|
||||
mask_paths = _load_paths_from_dataset(dataset_path, mask_column)
|
||||
naming_paths = _load_paths_from_dataset(dataset_path, naming_column) if naming_column != mask_column else mask_paths
|
||||
|
||||
console = Console()
|
||||
success = 0
|
||||
with Progress(
|
||||
SpinnerColumn(),
|
||||
TextColumn("[progress.description]{task.description}"),
|
||||
BarColumn(),
|
||||
MofNCompleteColumn(),
|
||||
TimeElapsedColumn(),
|
||||
console=console,
|
||||
) as progress:
|
||||
task = progress.add_task("Processing char masks", total=len(mask_paths))
|
||||
for mask_file, naming_file in zip(mask_paths, naming_paths, strict=True):
|
||||
progress.advance(task)
|
||||
rel_path = _output_relative(naming_file, data_root)
|
||||
latent_file = latents_path / rel_path.with_suffix(".pt")
|
||||
out_file = output_path / rel_path.with_suffix(".pt")
|
||||
|
||||
if not latent_file.exists():
|
||||
logger.warning(f"No target latent at {latent_file}, skipping mask {mask_file}")
|
||||
continue
|
||||
if not overwrite and out_file.is_file():
|
||||
continue
|
||||
|
||||
meta = torch.load(latent_file, map_location="cpu", weights_only=True)
|
||||
pixel_h = meta["height"] * VAE_SPATIAL_FACTOR
|
||||
pixel_w = meta["width"] * VAE_SPATIAL_FACTOR
|
||||
pixel_f = (meta["num_frames"] - 1) * VAE_TEMPORAL_FACTOR + 1
|
||||
|
||||
labels = _load_label_frames(mask_file, pixel_f)
|
||||
labels = _resize_labels(labels, pixel_h, pixel_w)
|
||||
mask = _labels_to_slot_masks(labels, num_slots, environment_switch)
|
||||
|
||||
out_file.parent.mkdir(parents=True, exist_ok=True)
|
||||
_atomic_save({"mask": mask.contiguous()}, out_file)
|
||||
success += 1
|
||||
|
||||
logger.info(f"Char-mask preprocessing complete: {success} masks saved to {output_path}")
|
||||
|
||||
|
||||
@app.command()
|
||||
def main(
|
||||
dataset_file: str = typer.Argument(..., help="Metadata file (CSV/JSON/JSONL) with a label-map column"),
|
||||
mask_column: str = typer.Option(..., help="Column of per-sample label-map video/image paths"),
|
||||
latents_dir: str = typer.Option(..., help="Directory of target video latents (for alignment)"),
|
||||
output_dir: str = typer.Option(..., help="Output directory for char_masks .pt files"),
|
||||
num_slots: int = typer.Option(6, help="Number of character binding slots K (channels = K+1)"),
|
||||
environment_switch: float = typer.Option(
|
||||
0.0, help="Uniform ch0 value: 0.0 = environment from reference, 1.0 = from driving video"
|
||||
),
|
||||
main_media_column: str | None = typer.Option(
|
||||
None, help="Column for output naming (defaults to --mask-column; set to the target-video column to align)"
|
||||
),
|
||||
overwrite: bool = typer.Option(False, help="Recompute even if the output already exists"),
|
||||
) -> None:
|
||||
"""Preprocess SCAIL-2 character label maps into ``[K+1, F_pix, H_pix, W_pix]`` mask tensors.
|
||||
|
||||
Example::
|
||||
|
||||
python scripts/process_char_masks.py dataset.csv \\
|
||||
--mask-column char_labels --latents-dir ./latents \\
|
||||
--output-dir ./char_masks --num-slots 6 --main-media-column media_path
|
||||
"""
|
||||
if not Path(dataset_file).is_file():
|
||||
raise typer.BadParameter(f"Dataset file not found: {dataset_file}")
|
||||
if num_slots < 1:
|
||||
raise typer.BadParameter("--num-slots must be >= 1")
|
||||
|
||||
compute_char_masks(
|
||||
dataset_file=dataset_file,
|
||||
mask_column=mask_column,
|
||||
latents_dir=latents_dir,
|
||||
output_dir=output_dir,
|
||||
num_slots=num_slots,
|
||||
environment_switch=environment_switch,
|
||||
main_media_column=main_media_column,
|
||||
overwrite=overwrite,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
app()
|
||||
Reference in New Issue
Block a user