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:
2026-07-13 09:49:25 +08:00
parent b69eedbd54
commit 5594d49c76
3 changed files with 230 additions and 3 deletions
+1 -1
View File
@@ -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 latentnearest 保留整數 labelch0 環境開關)。`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
View File
@@ -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()