SCAIL-2 character animation port (Phases 1-4) #1
+1
-1
@@ -40,7 +40,7 @@
|
|||||||
| 3.4 | LoRA 模式解凍 patchify_proj | ✅ | `trainer._unfreeze_patchify_proj`:mask_channels>0 時把 video patchify_proj 設 trainable,讓新欄位隨 LoRA 一起訓 |
|
| 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.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.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.8 | validation runner 接 driving/mask | ⬜ | 驗證期取樣尚未接 SCAIL 條件(config 內 validation 先停用),與 Phase 4 一起 |
|
||||||
| 3.9 | 實機訓練跑通 | ⬜ | 需 Linux + GPU + checkpoint,本機無法 |
|
| 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
|
`latents/`, `conditions/`, and `driving_latents/` come from the existing
|
||||||
`packages/ltx-trainer/scripts/process_dataset.py` (run it once per video set).
|
`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
|
`char_masks/` is produced by `packages/ltx-trainer/scripts/process_char_masks.py`
|
||||||
masks — that preprocessing is not implemented yet (see `docs/tasks.md` 3.7).
|
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
|
## 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