Files

239 lines
11 KiB
Python

"""Q1's common sample contract and evidence-based coordinate dispatch."""
from __future__ import annotations
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
import numpy as np
from scipy import sparse
from .coordinates import physical_targets, relative_cells, relative_targets
from .projection import project_intervals
MODALITIES = ("text", "audio", "vision")
DIMS = {"text": 768, "audio": 74, "vision": 35}
STEPS = {"text": 50, "audio": 500, "vision": 500}
K = 50
VERSION = "q1-unified-1"
class AlignmentError(ValueError):
"""The source cannot be assigned a defensible alignment coordinate."""
@dataclass
class ModalityProvenance:
source_weights: sparse.csr_matrix
source_count: np.ndarray
first_source: np.ndarray
last_source: np.ndarray
source_span: int
original_source_length: int
reported_length: int | None
length_conflict: bool = False
tail_ambiguous: bool = False
observed_dimensions: np.ndarray | None = None
coverage_dimensions: np.ndarray | None = None
quality_available: np.ndarray | None = None
@dataclass
class AlignedMultimodalSample:
features: dict[str, np.ndarray]
observed: dict[str, np.ndarray]
coverage: dict[str, np.ndarray]
provenance: dict[str, ModalityProvenance]
metadata: dict[str, Any]
target_intervals: np.ndarray
quality_mean: dict[str, np.ndarray] = field(default_factory=dict)
quality_available_fraction: dict[str, np.ndarray] = field(default_factory=dict)
def q2_arrays(self) -> tuple[dict[str, np.ndarray], np.ndarray]:
"""The existing Q2 model input shapes, without changing that model."""
return self.features, np.stack([self.observed[m] for m in MODALITIES], axis=-1)
@dataclass
class _Prepared:
values: np.ndarray
intervals: np.ndarray
observed: np.ndarray
quality: np.ndarray
reported_length: int | None
length_conflict: bool = False
tail_ambiguous: bool = False
quality_available: np.ndarray | None = None
def _relative_modality(record: dict[str, Any], name: str) -> _Prepared:
raw = np.asarray(record[name])
if raw.shape != (STEPS[name], DIMS[name]) or not np.isfinite(raw).all():
raise AlignmentError(f"{name}: expected finite {(STEPS[name], DIMS[name])}, got {raw.shape}")
if name == "text":
attention = np.asarray(record["attention_mask"], bool)
if attention.shape != (50,) or not np.array_equal(attention, np.arange(50) < int(attention.sum())):
raise AlignmentError("text attention mask must be a 50-position prefix")
length = int(attention.sum())
if length < 1:
raise AlignmentError("empty text attention mask")
span = length
observed = attention[:span] & np.any(raw[:span] != 0, axis=1)
conflict = ambiguous = False
else:
length = int(record[f"{name}_length"])
if length < 1 or length > STEPS[name]:
raise AlignmentError(f"{name}: invalid official length {length}")
nonzero = np.any(raw != 0, axis=1)
last = int(np.flatnonzero(nonzero)[-1]) + 1 if nonzero.any() else 0
conflict = last > length
if name == "audio" and conflict:
raise AlignmentError("audio contains observed positions beyond audio_lengths")
span = max(length, last)
observed = nonzero[:span]
ambiguous = bool(conflict)
return _Prepared(raw[:span].astype(np.float32, copy=False), relative_cells(span),
observed, np.ones(span, np.float32), length, bool(conflict), bool(ambiguous),
np.zeros(span, bool))
def _validate_physical(source: dict[str, Any]) -> tuple[float, dict[str, Any]]:
meta = source.get("_meta")
if not isinstance(meta, dict):
raise AlignmentError("physical mode requires stored Q1 metadata")
duration = float(meta.get("duration_s", float("nan")))
if not np.isfinite(duration) or duration <= 0:
raise AlignmentError("physical mode requires a finite positive duration")
if meta.get("media", {}).get("status") != "ok" or not meta.get("source_video_sha256"):
raise AlignmentError("physical mode requires verified media status and source hash")
for name in MODALITIES:
if f"native_{name}_intervals" not in source:
raise AlignmentError(f"physical mode lacks {name} timestamps")
return duration, meta
class Q1AlignmentAdapter:
"""Align either verified Q1 physical sources or official ordered sequences.
`auto` uses evidence in the input contract only; tensor shape never decides
whether time is physical. A malformed physical source is an error rather
than a silent relative fallback.
"""
def __init__(self, target_steps: int = K):
if target_steps < 1:
raise ValueError("target_steps must be positive")
self.target_steps = target_steps
def align(self, source: dict[str, Any], mode: str = "auto") -> AlignedMultimodalSample:
if mode not in {"auto", "physical", "relative"}:
raise AlignmentError(f"unsupported coordinate mode: {mode}")
if mode == "auto":
if "_meta" in source or any(k.startswith("native_") for k in source):
mode = "physical"
elif source.get("sequence_order_verified") is True:
mode = "relative"
else:
raise AlignmentError("auto mode requires Q1 physical evidence or verified sequence order")
if mode == "physical":
from q1_io import _native
duration, meta = _validate_physical(source)
target = physical_targets(duration, self.target_steps)
prepared = {}
for name in MODALITIES:
values, observed, intervals, quality, available = _native(source, name)
if np.asarray(intervals).shape != (len(values), 2) or np.any(np.asarray(intervals) < -1e-5) or np.any(np.asarray(intervals) > duration + 1e-5):
raise AlignmentError(f"{name}: physical timestamps outside media duration")
prepared[name] = _Prepared(values, intervals, observed, quality, len(values),
quality_available=available)
metadata = {"sample_id": meta.get("sample_id", f"{meta.get('video_id')}/{meta.get('clip_id')}"),
"coordinate_mode": "physical", "coordinate_unit": "seconds",
"physical_time_alignment": True, "duration_s": duration,
"source_video_sha256": meta["source_video_sha256"],
"dense_view": "views_sec_* (stored 0.1 s Q1 artifact)",
"quality_fields_available": {m: bool(np.asarray(prepared[m].quality_available).any()) for m in MODALITIES}}
else:
if source.get("sequence_order_verified") is not True:
raise AlignmentError("relative mode requires verified source order")
target = relative_targets(self.target_steps)
prepared = {name: _relative_modality(source, name) for name in MODALITIES}
metadata = {"sample_id": str(source.get("id", "")), "coordinate_mode": "relative",
"coordinate_unit": "normalized_progress", "physical_time_alignment": False,
"word_or_frame_timestamps_available": False,
"quality_fields_available": {m: False for m in MODALITIES}}
features = {}
observed = {}
coverage = {}
quality_mean = {}
quality_available_fraction = {}
provenance = {}
for name in MODALITIES:
item = prepared[name]
result = project_intervals(item.values, item.intervals, target, item.observed, item.quality,
item.quality_available)
features[name] = result.x
observed[name] = result.observed
coverage[name] = result.coverage
quality_mean[name] = result.quality_mean
quality_available_fraction[name] = result.quality_available_fraction
provenance[name] = ModalityProvenance(result.source_weights, result.source_count,
result.first_source, result.last_source, len(item.values),
len(source[name]) if mode == "relative" else len(item.values), item.reported_length,
item.length_conflict, item.tail_ambiguous, result.observed_dimensions,
result.coverage_dimensions, item.quality_available)
metadata.update({"target_steps": self.target_steps, "adapter_version": VERSION})
return AlignedMultimodalSample(features, observed, coverage, provenance, metadata, target,
quality_mean, quality_available_fraction)
def from_q1_sample(self, sample_id: str, feature_dir: Path | None = None) -> AlignedMultimodalSample:
from q1_io import FEATURE_DIR, load_sample
return self.align(load_sample(sample_id, FEATURE_DIR if feature_dir is None else feature_dir), "auto")
def from_unaligned_record(self, split: dict[str, Any], index: int) -> AlignedMultimodalSample:
"""Build verified ordered input from one official unaligned pickle row."""
attention = np.asarray(split["text_bert"][index, 1], bool)
raw_id = split["id"][index]
if isinstance(raw_id, bytes):
raw_id = raw_id.decode("utf-8", errors="replace")
record = {"id": str(raw_id), "sequence_order_verified": True,
"attention_mask": attention,
"text": split["text"][index], "audio": split["audio"][index],
"vision": split["vision"][index],
"audio_length": int(split["audio_lengths"][index]),
"vision_length": int(split["vision_lengths"][index])}
return self.align(record, "auto")
def adapt_official_split(split: dict[str, Any]) -> tuple[dict[str, np.ndarray], np.ndarray, dict[str, Any]]:
"""Q2's batch bridge; all rows are produced through Q1AlignmentAdapter."""
n = len(split["id"])
output = {m: np.zeros((n, K, DIMS[m]), np.float32) for m in MODALITIES}
masks = np.zeros((n, K, len(MODALITIES)), bool)
conflicts = ambiguous = padding = 0
coverage_sum = {m: 0.0 for m in MODALITIES}
observed_rows = {m: 0 for m in MODALITIES}
adapter = Q1AlignmentAdapter()
for i in range(n):
attention = np.asarray(split["text_bert"][i, 1], bool)
padding += int(np.count_nonzero(np.any(split["text"][i] != 0, axis=1) & ~attention))
sample = adapter.from_unaligned_record(split, i)
for j, name in enumerate(MODALITIES):
output[name][i] = sample.features[name]
masks[i, :, j] = sample.observed[name]
coverage_sum[name] += float(sample.coverage[name].sum())
observed_rows[name] += int(sample.observed[name].sum())
conflicts += int(sample.provenance["vision"].length_conflict)
ambiguous += int(sample.provenance["vision"].tail_ambiguous)
audit = {"method": "shared_interval_overlap_on_normalized_progress",
"coordinate_mode": "relative", "physical_time_alignment": False,
"samples": n, "vision_length_conflict_samples": conflicts,
"vision_tail_ambiguous_samples": ambiguous,
"nonzero_text_rows_outside_attention": padding,
"observed_target_rows": observed_rows,
"mean_target_coverage": {m: coverage_sum[m] / (n * K) for m in MODALITIES},
"quality_fields_available": False,
"word_or_frame_timestamps_available": False}
return output, masks, audit