239 lines
11 KiB
Python
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
|