Self-Forcing / predictor_data /kv_validation.py
Cccccz's picture
Add files using upload-large-folder tool
510ab6b verified
Raw
History Blame Contribute Delete
3.78 kB
"""Rebuild clean Wan K/V from saved prefeatures and compare with live cache."""
from __future__ import annotations
import torch
import torch.nn.functional as F
from .schema import (
CHUNK_FRAMES,
LATENT_HEIGHT,
LATENT_WIDTH,
TOKENS_PER_FRAME,
)
@torch.no_grad()
def rebuild_clean_kv(
model: torch.nn.Module,
block_id: int,
self_attn_input: torch.Tensor,
*,
start_frame: int,
num_frames: int = CHUNK_FRAMES,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Apply the frozen Teacher's K/V projections and causal RoPE."""
from wan.modules.causal_model import causal_rope_apply
if num_frames != CHUNK_FRAMES:
raise ValueError(f"expected {CHUNK_FRAMES} clean frames, got {num_frames}")
block = model.blocks[int(block_id)]
attention = block.self_attn
device = attention.k.weight.device
dtype = attention.k.weight.dtype
value = self_attn_input.to(device=device, dtype=dtype)
batch, tokens, _ = value.shape
if tokens != num_frames * TOKENS_PER_FRAME:
raise ValueError("clean prefeature token count is not frame-aligned")
grid_sizes = torch.tensor(
[[num_frames, LATENT_HEIGHT // 2, LATENT_WIDTH // 2]],
device=device,
dtype=torch.long,
).expand(batch, -1)
freqs = model.freqs.to(device=device)
key = attention.norm_k(attention.k(value)).view(
batch, tokens, attention.num_heads, attention.head_dim
)
rebuilt_k = causal_rope_apply(
key,
grid_sizes,
freqs,
start_frame=int(start_frame),
).type_as(value)
rebuilt_v = attention.v(value).view(
batch, tokens, attention.num_heads, attention.head_dim
)
return rebuilt_k, rebuilt_v
def comparison_metrics(actual: torch.Tensor, expected: torch.Tensor) -> dict[str, float]:
actual_flat = actual.float().flatten()
expected_flat = expected.float().flatten()
difference = actual_flat - expected_flat
return {
"relative_l2": float(
difference.norm() / expected_flat.norm().clamp_min(1e-12)
),
"cosine": float(F.cosine_similarity(actual_flat, expected_flat, dim=0)),
"max_abs": float(difference.abs().max()),
}
@torch.no_grad()
def validate_against_live_cache(
model: torch.nn.Module,
kv_cache: list[dict[str, torch.Tensor]],
clean_features: dict[int, torch.Tensor],
*,
start_frame: int,
num_frames: int = CHUNK_FRAMES,
relative_l2_limit: float = 5e-3,
cosine_limit: float = 0.9999,
) -> dict[str, float]:
"""Validate reconstructed K/V against the clean pass cache slice."""
start_token = int(start_frame) * TOKENS_PER_FRAME
end_token = start_token + int(num_frames) * TOKENS_PER_FRAME
metrics: dict[str, float] = {}
for block_id, feature in clean_features.items():
rebuilt_k, rebuilt_v = rebuild_clean_kv(
model,
block_id,
feature,
start_frame=start_frame,
num_frames=num_frames,
)
for field, rebuilt in (("k", rebuilt_k), ("v", rebuilt_v)):
expected = kv_cache[block_id][field][:, start_token:end_token]
values = comparison_metrics(rebuilt, expected)
for metric_name, value in values.items():
metrics[f"block_{block_id:02d}_{field}_{metric_name}"] = value
if (
values["relative_l2"] > relative_l2_limit
or values["cosine"] < cosine_limit
):
raise RuntimeError(
"clean KV rebuild failed "
f"block={block_id} field={field}: "
f"relative_l2={values['relative_l2']:.6g}, "
f"cosine={values['cosine']:.8f}"
)
return metrics