"""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 .source import native_arrays duration, meta = _validate_physical(source) target = physical_targets(duration, self.target_steps) prepared = {} for name in MODALITIES: values, observed, intervals, quality, available = native_arrays(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 .source 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