Cccccz's picture
Add files using upload-large-folder tool
d5e0d8f verified
Raw
History Blame Contribute Delete
13.8 kB
"""Manifest-backed offline dataset for the three Predictor-v4 transitions.
One manifest record represents one temporal chunk. This dataset expands every
usable record into the adjacent denoising pairs 0->1, 1->2 and 2->3. Chunk
zero is intentionally excluded because v4 conditions on the preceding chunk.
"""
from __future__ import annotations
import json
from functools import lru_cache
from pathlib import Path
from typing import Any, Iterable, Mapping
import torch
from safetensors import safe_open
from torch.utils.data import Dataset
SUPERVISION_PAIRS = ((0, 1), (1, 2), (2, 3))
SCHEMA_VERSION = "self_forcing_predictor_v4_bf16_v1"
FRAMES_PER_CHUNK = 3
def _resolve(root: Path, value: str | Path) -> Path:
path = Path(value)
return path if path.is_absolute() else root / path
def _load_selected(path: Path, names: Iterable[str]) -> dict[str, torch.Tensor]:
if not path.is_file():
raise FileNotFoundError(path)
with safe_open(str(path), framework="pt", device="cpu") as handle:
available = set(handle.keys())
missing = set(names).difference(available)
if missing:
raise KeyError(f"{path} is missing tensors {sorted(missing)}")
return {name: handle.get_tensor(name) for name in names}
def _load_clean_prefeature(
path: Path,
candidates: Iterable[str],
*,
expected_start_frame: int,
) -> torch.Tensor:
with safe_open(str(path), framework="pt", device="cpu") as handle:
available = set(handle.keys())
for name in candidates:
if name in available:
feature = handle.get_tensor(name)
break
else:
raise KeyError(
f"{path} has none of the expected tensors {tuple(candidates)}"
)
if "start_frame" in available:
actual_start = int(handle.get_tensor("start_frame").item())
if actual_start != expected_start_frame:
raise ValueError(
f"{path} starts at frame {actual_start}, expected "
f"{expected_start_frame}; history files are not ordered"
)
if "num_frames" in available:
actual_frames = int(handle.get_tensor("num_frames").item())
if actual_frames != FRAMES_PER_CHUNK:
raise ValueError(
f"{path} contains {actual_frames} frames, expected "
f"{FRAMES_PER_CHUNK}"
)
return feature
def _block_entry(mapping: Mapping[Any, Any], block_id: int) -> Any:
for key in (str(block_id), block_id, f"block_{block_id}", f"block_{block_id:02d}"):
if key in mapping:
return mapping[key]
raise KeyError(f"No prefeature entry for block {block_id}")
def _as_path_list(entry: Any) -> list[str]:
if isinstance(entry, (str, Path)):
return [str(entry)]
if isinstance(entry, Mapping):
# Builders may use {"files": [...]} or {"file": "..."}.
for key in ("files", "paths", "history", "file", "path"):
if key in entry:
return _as_path_list(entry[key])
if isinstance(entry, (list, tuple)):
return [str(value) for value in entry]
raise TypeError(f"Unsupported prefeature file entry: {entry!r}")
def _prefeature_names(block_id: int) -> tuple[str, ...]:
return (
"self_attn_input",
"clean_prefeature",
"prefeature",
"img_modulated",
f"block_{block_id}_self_attn_input",
f"block_{block_id:02d}_self_attn_input",
)
class PredictorV4PairDataset(Dataset):
"""Read safetensors records and expose adjacent-step supervision pairs."""
PAIRS = SUPERVISION_PAIRS
def __init__(
self,
manifest_path: str | Path,
*,
source_block_ids: tuple[int, ...] = (1, 28),
max_records: int | None = None,
require_previous_chunk: bool = True,
) -> None:
self.manifest_path = Path(manifest_path).resolve()
self.root = self.manifest_path.parent
self.source_block_ids = tuple(int(value) for value in source_block_ids)
with self.manifest_path.open("r", encoding="utf-8") as handle:
records = [json.loads(line) for line in handle if line.strip()]
if require_previous_chunk:
records = [record for record in records if int(record["chunk_id"]) > 0]
if max_records is not None:
records = records[: int(max_records)]
if not records:
raise ValueError(f"No usable records in {self.manifest_path}")
for record in records:
required = {
"step_tensor_file",
"previous_step_tensor_file",
"case_tensor_file",
"chunk_id",
}
missing = required.difference(record)
if missing:
raise ValueError(f"Manifest record lacks fields {sorted(missing)}")
if record.get("schema_version", SCHEMA_VERSION) != SCHEMA_VERSION:
raise ValueError(
f"Unsupported Predictor schema {record.get('schema_version')!r}"
)
chunk_id = int(record["chunk_id"])
expected_context_frames = chunk_id * FRAMES_PER_CHUNK
context_frames = int(
record.get("context_frames", expected_context_frames)
)
if context_frames != expected_context_frames:
raise ValueError(
f"chunk {chunk_id} context_frames={context_frames}, expected "
f"{expected_context_frames}"
)
history = record.get("history_clean_prefeature_files")
if history is None:
raise ValueError(
"Manifest record lacks history_clean_prefeature_files; "
"clean_prefeature_files contains only the current chunk"
)
for block_id in self.source_block_ids:
paths = _as_path_list(_block_entry(history, block_id))
if len(paths) != int(record["chunk_id"]):
raise ValueError(
f"chunk {record['chunk_id']} block {block_id} has "
f"{len(paths)} history files, expected {record['chunk_id']}"
)
self.records = records
def __len__(self) -> int:
return len(self.records) * len(self.PAIRS)
@lru_cache(maxsize=8)
def _load_case(self, relative_path: str) -> dict[str, torch.Tensor]:
path = _resolve(self.root, relative_path)
names: list[str] = []
with safe_open(str(path), framework="pt", device="cpu") as handle:
keys = set(handle.keys())
for block_id in self.source_block_ids:
for kind in ("k", "v"):
candidates = (
f"block_{block_id:02d}_cross_{kind}",
f"block_{block_id}_cross_{kind}",
f"block_{block_id}_text_{kind}",
f"block_{block_id:02d}_text_{kind}",
f"block_{block_id}_{kind}_txt",
f"text_{kind}_block_{block_id}",
)
found = next((name for name in candidates if name in keys), None)
if found is None:
raise KeyError(
f"{path} has no text {kind.upper()} for block {block_id}"
)
names.append(found)
return {name: handle.get_tensor(name) for name in names}
def _case_text_kv(
self,
relative_path: str,
) -> dict[int, dict[str, torch.Tensor]]:
tensors = self._load_case(relative_path)
result: dict[int, dict[str, torch.Tensor]] = {}
for block_id in self.source_block_ids:
result[block_id] = {}
for kind in ("k", "v"):
candidates = (
f"block_{block_id:02d}_cross_{kind}",
f"block_{block_id}_cross_{kind}",
f"block_{block_id}_text_{kind}",
f"block_{block_id:02d}_text_{kind}",
f"block_{block_id}_{kind}_txt",
f"text_{kind}_block_{block_id}",
)
name = next(name for name in candidates if name in tensors)
result[block_id][kind] = tensors[name]
return result
def _history_prefeature(
self,
record: Mapping[str, Any],
) -> dict[int, torch.Tensor]:
history = record["history_clean_prefeature_files"]
result = {}
for block_id in self.source_block_ids:
paths = _as_path_list(_block_entry(history, block_id))
chunks = [
_load_clean_prefeature(
_resolve(self.root, path),
_prefeature_names(block_id),
expected_start_frame=chunk_index * FRAMES_PER_CHUNK,
)
for chunk_index, path in enumerate(paths)
]
# Files are [1, chunk_tokens, dim]. Concatenate temporal chunks.
result[block_id] = torch.cat(chunks, dim=1)
return result
def __getitem__(self, index: int) -> dict[str, Any]:
record_index, pair_index = divmod(index, len(self.PAIRS))
record = self.records[record_index]
anchor_step, target_step = self.PAIRS[pair_index]
step_path = _resolve(self.root, record["step_tensor_file"])
step_names = (
f"step_{anchor_step}_final_hidden",
f"step_{target_step}_noisy_latent",
f"step_{target_step}_timestep",
f"step_{target_step}_final_hidden",
f"step_{target_step}_flow",
)
step_tensors = _load_selected(step_path, step_names)
previous_name = f"step_{target_step}_final_hidden"
previous = _load_selected(
_resolve(self.root, record["previous_step_tensor_file"]),
(previous_name,),
)
context_frames = int(
record.get(
"context_frames",
int(record["chunk_id"]) * FRAMES_PER_CHUNK,
)
)
return {
"target_latent": step_tensors[f"step_{target_step}_noisy_latent"],
"target_timestep": step_tensors[f"step_{target_step}_timestep"],
"anchor_hidden": step_tensors[f"step_{anchor_step}_final_hidden"],
"previous_chunk_hidden": previous[previous_name],
"target_hidden": step_tensors[f"step_{target_step}_final_hidden"],
"target_flow": step_tensors[f"step_{target_step}_flow"],
"clean_prefeature": self._history_prefeature(record),
"text_kv": self._case_text_kv(str(record["case_tensor_file"])),
"case_id": record.get("case_id"),
"chunk_id": int(record["chunk_id"]),
"context_frames": context_frames,
"anchor_step": anchor_step,
"target_step": target_step,
}
def predictor_v4_collate(items: list[dict[str, Any]]) -> dict[str, Any]:
"""Collate a context-length bucket into one batch."""
if not items:
raise ValueError("Cannot collate an empty batch")
context_frames = {item["context_frames"] for item in items}
if len(context_frames) != 1:
raise ValueError(
"A batch must have one history length; enable bucket_by_context"
)
tensor_keys = (
"target_latent",
"target_timestep",
"anchor_hidden",
"previous_chunk_hidden",
"target_hidden",
"target_flow",
)
batch: dict[str, Any] = {
key: torch.cat([item[key] for item in items], dim=0) for key in tensor_keys
}
block_ids = tuple(items[0]["clean_prefeature"])
batch["clean_prefeature"] = {
block_id: torch.cat(
[item["clean_prefeature"][block_id] for item in items], dim=0
)
for block_id in block_ids
}
batch["text_kv"] = {
block_id: {
kind: torch.cat(
[item["text_kv"][block_id][kind] for item in items], dim=0
)
for kind in ("k", "v")
}
for block_id in block_ids
}
for key in (
"case_id",
"chunk_id",
"context_frames",
"anchor_step",
"target_step",
):
batch[key] = [item[key] for item in items]
return batch
def _move(
tensor: torch.Tensor,
*,
device: torch.device,
dtype: torch.dtype,
) -> torch.Tensor:
target_dtype = dtype if tensor.is_floating_point() else tensor.dtype
return tensor.to(device=device, dtype=target_dtype, non_blocking=True)
def move_batch_to_device(
batch: dict[str, Any],
*,
device: torch.device,
dtype: torch.dtype,
) -> dict[str, Any]:
result = {
key: _move(batch[key], device=device, dtype=dtype)
for key in (
"target_latent",
"target_timestep",
"anchor_hidden",
"previous_chunk_hidden",
"target_hidden",
"target_flow",
)
}
result["clean_prefeature"] = {
int(block_id): _move(value, device=device, dtype=dtype)
for block_id, value in batch["clean_prefeature"].items()
}
result["text_kv"] = {
int(block_id): {
kind: _move(value, device=device, dtype=dtype)
for kind, value in values.items()
}
for block_id, values in batch["text_kv"].items()
}
for key in ("case_id", "chunk_id", "context_frames", "anchor_step", "target_step"):
result[key] = batch[key]
return result