Prepare minimum submission bundle

This commit is contained in:
2026-09-26 16:36:33 +08:00
parent 9cdd604117
commit 411f0f97e5
172 changed files with 12565 additions and 0 deletions
+82
View File
@@ -0,0 +1,82 @@
# E题最小提交材料
本目录只保留原题“四、结果与提交说明”中要求随附件提交的材料。特征文件、代码、模型参数与预测 CSV 合计 **47.19 MB**,低于 50 MB。官方原始附件由赛题另行提供,不重复打包;竞赛论文中的方法、实验表格和分析按题目要求写入论文正文。
## 提交内容
| 题目 | 文件 | 用途 |
|---|---|---|
| Q1 | `final/output/q1/features_v2/` | 100 条自生成多模态时序特征 `.npz`;附 `manifest_q1.jsonl`、`feature_manifest.json`、100 行样本汇总和 300 行模态汇总,用于追溯样本、维度与有效时长 |
| Q2 | `final/q2/math/`、`final/model/`、`final/adapter/` | C0–C7 数学方案的训练、数据处理、附件 3 推理和统一未对齐输入接口代码;随包参数对应验证集选定的 C6 |
| Q2 | `final/experiments/q2/unaligned_math_all_b128/` | C6 选定模型权重、结构化填补器、训练集预处理统计、验证集校准与运行配置 |
| Q2 | `final/output/q2/attachment3_predictions.csv` | 附件 3 的 30 条无标签预测 |
| Q3 | `final/q3/ati_ho/`、`final/q3/run_experiments.py`、`final/q2/deep_learning/q2/`、相关 `final/model/` 与 `final/adapter/` | ATI–HO 训练、预测、解释与局部证据计算代码;含运行说明 |
| Q3 | `final/experiments/q3/ati_ho/`、`final/experiments/q2/unaligned_deep_two_b128/unaligned_50_robust_stats.npz` | 验证集选定的 ATI–HO A0 三种子权重、配置和推理所需缩放器 |
| Q3 | `final/output/q3/ati_ho/attachment4_predictions.csv`、`attachment4_explanations.csv`、`attachment4_local_evidence.csv` | 附件 4 的全量预测、模态作用解释和局部证据位置(20 条样本、600 条局部证据记录) |
未放入提交目录的训练日志、消融表、Bootstrap 表、错误归因表和历史方案输出不属于第四小节要求的附件。模型比较、验证结果与分析应呈现在论文正文中。
## 运行环境与数据位置
使用 Python 3.11 或更新版本。安装依赖:
```bash
python -m venv .venv
source .venv/bin/activate # Windows PowerShell: .venv\Scripts\Activate.ps1
python -m pip install -r requirements.txt
```
GPU 运行时,请安装与本机 CUDA 驱动匹配的 PyTorch;没有 GPU 时可使用 CPU,但训练和局部解释会更慢。附件 3 的文本重建会调用 `google-bert/bert-base-uncased`;首次运行需能下载该模型,或提前放入 Hugging Face 缓存。
将官方数据根目录放到任意位置,并设置 `FINAL_DATA_DIR` 指向该根目录。目录中需有:
- `附件2-数据集特征文件/unaligned_50.pkl`
- `附件3-模态缺失特征样本/未对齐版本/`
- `附件4-可解释专项视频样本与特征文件/附件4-可解释专项视频样本与特征文件/未对齐版本/`
Windows PowerShell 示例:
```powershell
$env:FINAL_DATA_DIR = "D:/E题数据"
```
Linux/WSL 示例:
```bash
export FINAL_DATA_DIR="/data/E题数据"
```
## 重新生成专项预测文件
从 `submit/` 目录运行。Q2 使用随包提供的 C6 权重、预处理统计和校准参数:
```bash
python -m final.q2.math.predict_attachment3 --input-version unaligned_50 --device auto
```
Q3 使用随包提供的 A0 三种子权重,生成预测、模态解释和局部 Owen 证据 CSV:
```bash
python -m final.q3.ati_ho.predict_attachment4 --device auto
```
输出分别写入 `final/output/q2/` 和 `final/output/q3/ati_ho/`。附件 3、附件 4 是无标签专项测试集,输出文件不包含测试指标。
## 从头训练
Q2 的 C0–C7 训练与比较代码在 `final/q2/math/train.py`。按本次 C6 运行配置执行:
```bash
python -m final.q2.math.train --input-version unaligned_50 --epochs 12 --imputer-epochs 8 --batch-size 128 --patience 3 --seed 20260924 --device cuda --output-dir final/experiments/q2/unaligned_math_rerun
```
无 CUDA 时将 `--device cuda` 改为 `--device cpu`。训练完成后可用 `--results-dir final/experiments/q2/unaligned_math_rerun` 让 `predict_attachment3.py` 使用新权重。
Q3 完整训练会生成 ATI 消融与 EarlyConcat、MoFE 基线权重;之后运行完整验证和解释评估:
```bash
python -m final.q3.ati_ho.train --phase all --device cuda --force
python -m final.q3.ati_ho.evaluate --device auto
```
完整训练需要附件 2 训练集以及随包提供的训练集缩放器。训练产生的比较模型和审计结果写入 `final/experiments/`,不需要作为当前最小附件提交。
+1
View File
@@ -0,0 +1 @@
"""Consolidated Q1 alignment and Q2 models."""
+17
View File
@@ -0,0 +1,17 @@
"""Unified Q1 alignment interface for physical time and relative progress."""
from .core import (
AlignedMultimodalSample,
AlignmentError,
ModalityProvenance,
Q1AlignmentAdapter,
adapt_official_split,
)
__all__ = [
"AlignedMultimodalSample",
"AlignmentError",
"ModalityProvenance",
"Q1AlignmentAdapter",
"adapt_official_split",
]
+26
View File
@@ -0,0 +1,26 @@
"""Coordinate constructors; no feature aggregation happens here."""
from __future__ import annotations
import numpy as np
def physical_targets(duration_s: float, steps: int) -> np.ndarray:
"""K fixed bins spanning verified real media duration, in seconds."""
duration = float(duration_s)
if not np.isfinite(duration) or duration <= 0 or steps < 1:
raise ValueError("physical targets require positive finite duration and steps")
edges = np.linspace(0.0, duration, steps + 1, dtype=np.float64)
return np.column_stack((edges[:-1], edges[1:]))
def relative_cells(length: int) -> np.ndarray:
"""Ordered source cells on a unit progress axis, with no time claim."""
if length < 1:
raise ValueError("relative source length must be positive")
left = np.arange(length, dtype=np.float64) / length
return np.column_stack((left, left + 1.0 / length))
def relative_targets(steps: int) -> np.ndarray:
"""K fixed cells on the same unit progress axis."""
return relative_cells(steps)
+238
View File
@@ -0,0 +1,238 @@
"""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
+108
View File
@@ -0,0 +1,108 @@
"""The sole interval overlap projection kernel used by both coordinate modes."""
from __future__ import annotations
from dataclasses import dataclass
import numpy as np
from scipy import sparse
@dataclass
class Projection:
x: np.ndarray
observed: np.ndarray
coverage: np.ndarray
source_count: np.ndarray
first_source: np.ndarray
last_source: np.ndarray
source_weights: sparse.csr_matrix
observed_dimensions: np.ndarray
coverage_dimensions: np.ndarray
quality_mean: np.ndarray
quality_available_fraction: np.ndarray
def project_intervals(
values: np.ndarray,
source_intervals: np.ndarray,
target_intervals: np.ndarray,
observed: np.ndarray,
quality: np.ndarray | None = None,
quality_available: np.ndarray | None = None,
) -> Projection:
"""Project source cells using overlap * quality * observed validity.
Row provenance uses any valid dimension. Feature values use validity per
dimension, so partially observed physical features remain partially missing.
"""
source = np.asarray(values, dtype=np.float32)
src = np.asarray(source_intervals, dtype=np.float64)
dst = np.asarray(target_intervals, dtype=np.float64)
if source.ndim != 2 or src.shape != (len(source), 2) or dst.ndim != 2 or dst.shape[1] != 2:
raise ValueError("inconsistent source features or interval dimensions")
if not np.isfinite(source).all() or not np.isfinite(src).all() or not np.isfinite(dst).all():
raise ValueError("non-finite source features or intervals")
if np.any(src[:, 1] < src[:, 0]) or np.any(dst[:, 1] <= dst[:, 0]):
raise ValueError("source widths must be nonnegative and target widths positive")
obs = np.asarray(observed, bool)
if obs.shape == (len(source),):
obs_dim = np.broadcast_to(obs[:, None], source.shape)
elif obs.shape == source.shape:
obs_dim = obs
obs = obs.any(axis=1)
else:
raise ValueError("observed must have source-row or source-feature shape")
q = np.ones(len(source), dtype=np.float64) if quality is None else np.asarray(quality, dtype=np.float64)
if q.shape != (len(source),) or not np.isfinite(q).all() or np.any(q < 0):
raise ValueError("quality must be finite and nonnegative per source row")
overlap = np.maximum(0.0, np.minimum(dst[:, None, 1], src[None, :, 1])
- np.maximum(dst[:, None, 0], src[None, :, 0]))
available = np.zeros(len(source), bool) if quality_available is None else np.asarray(quality_available, bool)
if available.shape != (len(source),):
raise ValueError("quality availability must be per source row")
# Keep the same multiplication and accumulation order as the original
# official-unaligned projection when quality is uniformly one.
physical = overlap.copy()
physical *= obs[None, :]
row_weight = physical.copy()
row_weight *= q[None, :]
mass = row_weight.sum(axis=1)
row_valid = mass > 0
normalized = np.zeros_like(row_weight, dtype=np.float32)
normalized[row_valid] = (row_weight[row_valid] / mass[row_valid, None]).astype(np.float32)
support = row_weight > 0
count = support.sum(axis=1).astype(np.uint16)
first = np.full(len(dst), -1, dtype=np.int32)
last = np.full(len(dst), -1, dtype=np.int32)
if row_valid.any():
first[row_valid] = support[row_valid].argmax(axis=1)
last[row_valid] = len(source) - 1 - support[row_valid, ::-1].argmax(axis=1)
width = dst[:, 1] - dst[:, 0]
physical_mass = physical.sum(axis=1)
coverage = np.clip(physical_mass / width, 0.0, 1.0).astype(np.float32)
qmean = np.ones(len(dst), np.float32)
qavailable = np.zeros(len(dst), np.float32)
physical_valid = physical_mass > 0
qmean[physical_valid] = (mass[physical_valid] / physical_mass[physical_valid]).astype(np.float32)
qavailable[physical_valid] = ((physical[physical_valid] @ available.astype(np.float64))
/ physical_mass[physical_valid]).astype(np.float32)
# The common full-dimension case follows the original matrix product
# exactly; this is also much faster for 500 x 768 input.
if np.array_equal(obs_dim, np.broadcast_to(obs[:, None], source.shape)):
x = np.zeros((len(dst), source.shape[1]), np.float32)
x[row_valid] = ((row_weight[row_valid] @ source) / mass[row_valid, None]).astype(np.float32)
observed_dimensions = np.broadcast_to(row_valid[:, None], x.shape).copy()
coverage_dimensions = np.broadcast_to(coverage[:, None], x.shape).copy()
else:
dim_physical = overlap[:, :, None] * obs_dim[None, :, :]
dim_weight = dim_physical * q[None, :, None]
dim_mass = dim_weight.sum(axis=1)
observed_dimensions = dim_mass > 0
x = np.zeros((len(dst), source.shape[1]), np.float32)
numerator = np.einsum("ksd,sd->kd", dim_weight, source, optimize=True)
x[observed_dimensions] = (numerator[observed_dimensions] / dim_mass[observed_dimensions]).astype(np.float32)
coverage_dimensions = np.clip(dim_physical.sum(axis=1) / width[:, None], 0, 1).astype(np.float32)
return Projection(x, row_valid, coverage, count, first, last,
sparse.csr_matrix(normalized), observed_dimensions, coverage_dimensions,
qmean, qavailable)
+51
View File
@@ -0,0 +1,51 @@
"""Read the native Q1 artifacts needed by the unified alignment adapter."""
from __future__ import annotations
import json
from pathlib import Path
from typing import Any
import numpy as np
FEATURE_DIR = Path(__file__).resolve().parents[1] / "output" / "q1" / "features_v2"
def load_sample(sample_id: str, feature_dir: Path = FEATURE_DIR) -> dict[str, Any]:
"""Load a Q1 sample by its video/clip ID from a feature directory."""
feature_dir = Path(feature_dir)
manifest = feature_dir / "manifest_q1.jsonl"
for line in manifest.read_text(encoding="utf-8").splitlines():
row = json.loads(line)
if row["sample_id"] != sample_id:
continue
path = feature_dir / Path(row["feature_path"]).name
with np.load(path, allow_pickle=False) as archive:
result = {name: archive[name] for name in archive.files}
result["_path"] = path
result["_manifest"] = row
result["_meta"] = json.loads(str(result["meta_json"]))
return result
raise KeyError(f"sample_id not found: {sample_id}")
def native_arrays(sample: dict[str, Any], modality: str):
"""Return values, observed dimensions, time intervals and quality evidence."""
if modality == "text":
values = sample["native_text_features"].astype(np.float32)
return (
values,
np.broadcast_to(sample["native_text_observed"][:, None], values.shape),
sample["native_text_intervals"].astype(np.float32),
sample["native_text_quality_effective"].astype(np.float32),
sample["native_text_quality_available"].astype(bool),
)
if modality in ("audio", "vision"):
prefix = f"native_{modality}_"
return (
sample[prefix + "features"].astype(np.float32),
sample[prefix + "mask"].astype(bool),
sample[prefix + "intervals"].astype(np.float32),
sample[prefix + "quality"].astype(np.float32),
sample[prefix + "quality_available"].astype(bool),
)
raise ValueError(f"unknown modality: {modality}")
+13
View File
@@ -0,0 +1,13 @@
"""Shared input-data locations for the standalone deliverable."""
from __future__ import annotations
import os
from pathlib import Path
PROJECT_ROOT = Path(__file__).resolve().parent
DATA_ROOT = Path(os.environ.get("FINAL_DATA_DIR", PROJECT_ROOT / "data")).expanduser().resolve()
ATTACHMENT1 = DATA_ROOT / "附件1-数据集原始多模态样本" / "MOSEI数据集部分原始视频-100条"
ATTACHMENT2 = DATA_ROOT / "附件2-数据集特征文件"
ATTACHMENT3 = DATA_ROOT / "附件3-模态缺失特征样本"
ATTACHMENT4 = DATA_ROOT / "附件4-可解释专项视频样本与特征文件"
@@ -0,0 +1,422 @@
{
"seed": 20260924,
"text_encoder": "official precomputed text field; encoder revision not supplied",
"training_configuration": {
"student_epoch_limit": 12,
"imputer_epochs": 8,
"batch_size": 128,
"early_stopping_patience": 3,
"device": "cuda",
"device_name": "NVIDIA GeForce RTX 5070 Ti",
"optimizer": "AdamW",
"student_learning_rate": 0.0003,
"student_weight_decay": 0.001,
"imputer_learning_rate": 0.0003,
"imputer_weight_decay": 0.0001,
"early_stopping_metric": "mean untempered selection_nll over fixed group-disjoint internal training scenarios",
"inner_selection_scenarios": [
"0.0/natural",
"0.3/single",
"0.3/sync",
"0.5/async"
],
"inner_selection_source_video_groups": 76
},
"training_input": "E题数据/附件2-数据集特征文件/unaligned_50.pkl",
"input_version": "unaligned_50",
"q1_alignment_adapter": {
"train": {
"method": "shared_interval_overlap_on_normalized_progress",
"coordinate_mode": "relative",
"physical_time_alignment": false,
"samples": 3395,
"vision_length_conflict_samples": 618,
"vision_tail_ambiguous_samples": 618,
"nonzero_text_rows_outside_attention": 86078,
"observed_target_rows": {
"text": 169750,
"audio": 169750,
"vision": 163302
},
"mean_target_coverage": {
"text": 1.0,
"audio": 1.0,
"vision": 0.9540337701314328
},
"quality_fields_available": false,
"word_or_frame_timestamps_available": false
},
"valid": {
"method": "shared_interval_overlap_on_normalized_progress",
"coordinate_mode": "relative",
"physical_time_alignment": false,
"samples": 728,
"vision_length_conflict_samples": 141,
"vision_tail_ambiguous_samples": 141,
"nonzero_text_rows_outside_attention": 17772,
"observed_target_rows": {
"text": 36400,
"audio": 36400,
"vision": 35315
},
"mean_target_coverage": {
"text": 1.0,
"audio": 1.0,
"vision": 0.96090314748523
},
"quality_fields_available": false,
"word_or_frame_timestamps_available": false
},
"test": {
"method": "shared_interval_overlap_on_normalized_progress",
"coordinate_mode": "relative",
"physical_time_alignment": false,
"samples": 727,
"vision_length_conflict_samples": 131,
"vision_tail_ambiguous_samples": 131,
"nonzero_text_rows_outside_attention": 18041,
"observed_target_rows": {
"text": 36350,
"audio": 36350,
"vision": 35034
},
"mean_target_coverage": {
"text": 1.0,
"audio": 1.0,
"vision": 0.9549938172651288
},
"quality_fields_available": false,
"word_or_frame_timestamps_available": false
}
},
"training_sha256": "77eda14a06be9749a96c52ae45470c7cffcfa7219011eae391d231a0664c3762",
"official_group_overlap": {
"train_valid": 0,
"train_test": 0,
"valid_test": 0
},
"official_splits": {
"train": {
"n": 3395,
"source_video_groups": 1528
},
"valid": {
"n": 728,
"source_video_groups": 239
},
"test": {
"n": 727,
"source_video_groups": 381
}
},
"internal_train_holdouts": {
"fit": {
"n": 3030,
"video_groups": 1375
},
"reliability_selection": {
"n": 190,
"video_groups": 76
},
"temperature_calibration": {
"n": 175,
"video_groups": 77
},
"all_group_disjoint": true
},
"feature_standardization": "fit-only observed rows, per-dimension; fixed for valid/test/attachment3",
"missing_mask": "official text attention and source lengths plus row observation; normalized-progress overlap preserves empty bins; q*=1 only where visible, J_Q=0",
"observation_quality": {
"quality_score_fields_present": false,
"quality_available_flag_present": false,
"fallback": "q*=1 and J_Q=0 for visible rows; R_eff=R",
"quality_noise_mapping_ablation": "not identifiable on unaligned_50 because no row quality score varies"
},
"imputer": {
"type": "structured linear Gaussian shared-private state space",
"state_dims": {
"shared": 8,
"private_each": 4
},
"posterior": "block-tridiagonal equivalent Kalman information filter + RTS smoother",
"sampling": "joint latent trajectories and missing emissions; observed features copied exactly",
"fit_objective": "train-only observed Gaussian marginal likelihood including log determinants",
"epochs": 8,
"frozen_before_teacher_student": true
},
"architecture": {
"projection": 32,
"bigru_hidden_each_direction": 16,
"cross_source_layers": 1,
"cross_time_read": true,
"rank": 4,
"reliability_gru": "directional hidden decay; reset applied before candidate map; update gate multiplied by rho",
"final_gate": "rho times bounded content score plus positive null prior",
"output": "neutral point mass plus sign-specific Beta magnitudes; K-path probabilities mixed before decoding"
},
"reliability_hyperparameters": {
"selected_per_model_on": "group-disjoint internal training reliability-validation slice",
"candidate_values": [
[
0.5,
0.0,
0.0,
0.0
],
[
0.5,
0.05,
0.05,
0.05
],
[
0.5,
0.1,
0.0,
0.0
],
[
0.3,
0.05,
0.05,
0.05
],
[
0.7,
0.05,
0.05,
0.05
]
],
"validation_scenarios": [
"0.0/natural",
"0.3/single",
"0.3/sync",
"0.5/async"
],
"selected_by_model": {
"teacher": [
0.3,
0.05,
0.05,
0.05
],
"C1": [
0.5,
0.05,
0.05,
0.05
],
"C2": [
0.5,
0.05,
0.05,
0.05
],
"C3": [
0.5,
0.0,
0.0,
0.0
],
"C4": [
0.7,
0.05,
0.05,
0.05
],
"C5": [
0.3,
0.05,
0.05,
0.05
],
"C6": [
0.3,
0.05,
0.05,
0.05
],
"C6_no_distance": [
0.5,
0.0,
0.0,
0.0
],
"C6_no_reconstruction": [
0.3,
0.05,
0.05,
0.05
],
"C6_pointmask": [
0.3,
0.05,
0.05,
0.05
],
"C7_distill": [
0.5,
0.05,
0.05,
0.05
],
"C7_group": [
0.5,
0.0,
0.0,
0.0
]
}
},
"group_risk_hyperparameters": {
"selection": "lambda_group and group_temperature jointly selected with reliability hyperparameters on fixed group-disjoint internal training scenarios",
"candidate_values": [
[
0.05,
0.1
],
[
0.1,
0.05
],
[
0.1,
0.1
],
[
0.1,
0.2
],
[
0.2,
0.1
]
],
"selected": [
0.1,
0.1
],
"selection_split": "reliability_validation"
},
"loss": {
"supervision": "negative log mixture of Beta interval masses plus scaled Huber mean term",
"delta_u": 0.027777499999999997,
"delta_u_source": "half the minimum positive spacing of nonzero absolute labels in fit only",
"lambda_y": 1.0,
"lambda_distill": 0.1,
"lambda_reconstruction": 0.05,
"lambda_group_default": 0.1,
"group_temperature_default": 0.1,
"selected_group_risk": [
0.1,
0.1
],
"distill_temperature": 2.0,
"distill_retention_exponent": 1.0,
"imputer_regularization": {
"emission_l2": 0.0001,
"transition_l2": 0.0001
},
"group_and_distill_separate": true
},
"calibration": {
"method": "temperature scaling on a group-disjoint internal official-train holdout, separated from reliability selection",
"temperature": 1.122980387832455,
"valid_used_for_selection": true,
"test_used_for_selection_or_calibration": false
},
"selected_model": "C6",
"attachment3_low_information_priors": {
"class_probability_method": "fit counts + one pseudocount per class",
"class_probability_values": [
0.2786020441806792,
0.2205736894164194,
0.5008242664029015
],
"negative_beta": [
1.1441766023635864,
1.8200817108154297
],
"positive_beta": [
1.271026611328125,
2.4681642055511475
]
},
"ablation_definitions": {
"C1": "masked BiGRU; no posterior imputation, explicit reliability or source gate",
"C2": "exact Gaussian posterior mean; no joint trajectory integral",
"C3": "joint trajectory integral plus final reliability/content fusion gate",
"C4": "C3 plus bounded cross-time source attention and null source",
"C5": "C4 plus reliability-modulated BiGRU update",
"C6": "C5 plus optional rank-4 CP residual",
"C6_no_distance": "C6 with uncertainty retained but both distance/span reliability penalties fixed to zero",
"C6_no_reconstruction": "C6 trained without the auxiliary hidden-feature reconstruction loss",
"C6_pointmask": "C6 trained with independent point masking instead of contiguous spans",
"C7_distill": "C6 plus entropy/retention-weighted teacher distillation only",
"C7_group": "C6 plus smooth worst-group risk only"
},
"masking": {
"rates": [
0.0,
0.1,
0.3,
0.5,
0.7
],
"patterns": [
"single",
"sync",
"partial",
"async"
],
"preserve_at_least_fraction_per_selected_modality": 0.2,
"controlled_sweep_split": "official validation",
"identical_masks_across_models": true,
"scenario_count": 42,
"scenario_seed": 20261833,
"training_mask_rng_seed": 20261227,
"reliability_scenario_seed": 20261830,
"controlled_torch_sampling_seed": 20261476,
"paired_control_bootstrap_seed": 20261477,
"mask_audit_file": null,
"mask_audit_omitted_reason": "Row-level audit omitted from the size-limited deliverable; regenerated by rerunning final.q2.math.train.",
"additional_one_factor_controls": [
"modality T/A/V and combinations",
"start/middle/end",
"one-long/multiple-short",
"sync/partial/async"
],
"semantic_position_control": "not run: unaligned_50 does not provide audited semantic boundary indices; raw text is prohibited in student inputs"
},
"final_test_metrics": {
"n": 727,
"accuracy": 0.672627235213205,
"macro_f1": 0.5484581573154621,
"negative_support": 207,
"neutral_support": 158,
"positive_support": 362,
"negative_recall": 0.7439613526570048,
"middle_recall": 0.10126582278481013,
"positive_recall": 0.8812154696132597,
"regression_mae": 0.7100059986114502,
"regression_rmse": 0.969292458045841,
"pearson": 0.6424147486686707,
"brier": 0.4367243729993746,
"classification_nll": 0.7538501024246216,
"ece_15": 0.047142177615237854,
"selection_nll": 2.8449835777282715,
"interval_90_coverage": 0.8954607977991746,
"interval_90_mean_width": 2.48697829246521,
"predictive_variance_mean_uncalibrated": 0.6001424193382263,
"within_trajectory_variance_mean": 0.6001414060592651,
"between_trajectory_variance_mean": 1.021712705551181e-06,
"predictive_mean_mean_calibrated": 0.18709982931613922,
"predictive_variance_mean_calibrated": 0.6344733238220215
},
"test_gate_diagnostics_file": null,
"test_gate_diagnostics_omitted_reason": "Row-level audit omitted from the size-limited deliverable; regenerated by rerunning final.q2.math.train.",
"attachment3_cases": 0,
"attachment3_labeled_metrics": null,
"completed_utc": "2026-09-25T12:10:16Z"
}
@@ -0,0 +1,27 @@
{
"n": 728,
"accuracy": 0.6085164835164835,
"macro_f1": 0.518781225964892,
"negative_support": 206,
"neutral_support": 184,
"positive_support": 338,
"negative_recall": 0.6796116504854369,
"middle_recall": 0.125,
"positive_recall": 0.8284023668639053,
"regression_mae": 0.6846789717674255,
"regression_rmse": 0.9208215740080159,
"pearson": 0.6101368069648743,
"brier": 0.4977383080922297,
"classification_nll": 0.8427478075027466,
"ece_15": 0.03950894476620705,
"selection_nll": 2.783693552017212,
"interval_90_coverage": 0.8873626373626373,
"interval_90_mean_width": 2.3910491466522217,
"predictive_variance_mean_uncalibrated": 0.5537729859352112,
"within_trajectory_variance_mean": 0.5537727475166321,
"between_trajectory_variance_mean": 2.9140662149984564e-07,
"predictive_mean_mean_calibrated": 0.2094990462064743,
"predictive_variance_mean_calibrated": 0.5816943049430847,
"selected_model": "C6",
"temperature": 1.122980387832455
}
@@ -0,0 +1,37 @@
{
"selected_method": "A0",
"provisional_seed42_method": "A2",
"candidate_methods_with_three_seeds": [
"A2",
"A0",
"A1"
],
"selection_rule": "lowest mean fixed four-scenario validation task loss across seeds 42, 3407, 2026",
"candidate_summary": [
{
"method": "A0",
"seed_losses": "[0.8643234267339601, 0.8756466648735842, 0.8574190991265433]",
"mean_validation_selection_loss": 0.8657963969113626,
"std_validation_selection_loss": 0.009202622948013908,
"seeds": 3,
"validation_only_selection": true
},
{
"method": "A1",
"seed_losses": "[0.8649765662439577, 0.8810990981675766, 0.8557776766163963]",
"mean_validation_selection_loss": 0.8672844470093102,
"std_validation_selection_loss": 0.012817501026466038,
"seeds": 3,
"validation_only_selection": true
},
{
"method": "A2",
"seed_losses": "[0.8634063961741689, 0.880271397449158, 0.862087192279952]",
"mean_validation_selection_loss": 0.8685883286344263,
"std_validation_selection_loss": 0.010139311979910028,
"seeds": 3,
"validation_only_selection": true
}
],
"attachment4_labels_used": false
}
@@ -0,0 +1,42 @@
{
"method": "A0",
"seed": 2026,
"best_epoch": 4,
"best_selection_loss": 0.8574190991265433,
"batch_size": 64,
"epoch_limit": 12,
"patience": 3,
"optimizer": "AdamW",
"learning_rate": 0.0003,
"weight_decay": 0.001,
"gradient_clip_norm": 1.0,
"training_mask_rates": [
0.0,
0.1,
0.3,
0.5,
0.7
],
"training_mask_patterns": [
"single",
"sync",
"partial",
"async"
],
"training_mask_seed_base": 20261227,
"same_orders_and_masks_across_methods_for_same_seed": true,
"config": {
"name": "A0_main_effects",
"low_rank": false,
"cross_attention": false,
"anchored": true,
"lambda_interaction": 0.001,
"lambda_mask": 0.0,
"rank": 4,
"hidden": 64,
"gru_hidden_per_direction": 32,
"attention_heads": 4,
"attention_ffn": 128,
"eta_init": 0.1
}
}
@@ -0,0 +1,42 @@
{
"method": "A0",
"seed": 3407,
"best_epoch": 3,
"best_selection_loss": 0.8756466648735842,
"batch_size": 64,
"epoch_limit": 12,
"patience": 3,
"optimizer": "AdamW",
"learning_rate": 0.0003,
"weight_decay": 0.001,
"gradient_clip_norm": 1.0,
"training_mask_rates": [
0.0,
0.1,
0.3,
0.5,
0.7
],
"training_mask_patterns": [
"single",
"sync",
"partial",
"async"
],
"training_mask_seed_base": 20261227,
"same_orders_and_masks_across_methods_for_same_seed": true,
"config": {
"name": "A0_main_effects",
"low_rank": false,
"cross_attention": false,
"anchored": true,
"lambda_interaction": 0.001,
"lambda_mask": 0.0,
"rank": 4,
"hidden": 64,
"gru_hidden_per_direction": 32,
"attention_heads": 4,
"attention_ffn": 128,
"eta_init": 0.1
}
}
@@ -0,0 +1,42 @@
{
"method": "A0",
"seed": 42,
"best_epoch": 4,
"best_selection_loss": 0.8643234267339601,
"batch_size": 64,
"epoch_limit": 12,
"patience": 3,
"optimizer": "AdamW",
"learning_rate": 0.0003,
"weight_decay": 0.001,
"gradient_clip_norm": 1.0,
"training_mask_rates": [
0.0,
0.1,
0.3,
0.5,
0.7
],
"training_mask_patterns": [
"single",
"sync",
"partial",
"async"
],
"training_mask_seed_base": 20261227,
"same_orders_and_masks_across_methods_for_same_seed": true,
"config": {
"name": "A0_main_effects",
"low_rank": false,
"cross_attention": false,
"anchored": true,
"lambda_interaction": 0.001,
"lambda_mask": 0.0,
"rank": 4,
"hidden": 64,
"gru_hidden_per_direction": 32,
"attention_heads": 4,
"attention_ffn": 128,
"eta_init": 0.1
}
}
@@ -0,0 +1,119 @@
{
"experiment": "ATI\u2013HO Q3 staged training and structural attribution audit",
"created_utc": "2026-09-26T07:33:41Z",
"device": "cuda",
"torch_version": "2.14.0+cu130",
"cuda_available": true,
"cuda_version": "13.0",
"gpu": "NVIDIA GeForce RTX 5070 Ti",
"python": "3.14.7",
"seeds": [
42,
3407,
2026
],
"epochs_max": 12,
"training_protocol": {
"batch_size": 64,
"early_stopping_patience": 3,
"optimizer": "AdamW",
"learning_rate": 0.0003,
"weight_decay": 0.001,
"gradient_clip_norm": 1.0,
"training_mask_rates": [
0.0,
0.1,
0.3,
0.5,
0.7
],
"training_mask_patterns": [
"single",
"sync",
"partial",
"async"
],
"validation_selection_scenarios": [
"0.0/none",
"0.3/single",
"0.3/sync",
"0.5/async"
],
"held_out_attachment4_touched_during_training": false
},
"ati_output": {
"parameter_vector": "3 centered class logits + r_negative + r_positive",
"intensity": "negative/positive magnitudes are 3*sigmoid(r); neutral class is exactly zero",
"loss": "cross entropy + conditional magnitude SmoothL1 + 0.2*Huber(delta=0.25) + configured regularizers",
"baseline_checkpoint_reuse": "No: retrain B0 and B1 on the fixed ATI split/mask schedule because existing Q2 checkpoints differ in seeds, batch size, and schedule.",
"calibration_temperature": 1.0
},
"data": {
"feature_file": "${FINAL_DATA_DIR}/attachment2/unaligned_50.pkl",
"feature_sha256": "77eda14a06be9749a96c52ae45470c7cffcfa7219011eae391d231a0664c3762",
"scaler_file": "final/experiments/q2/unaligned_deep_two_b128/unaligned_50_robust_stats.npz",
"scaler_max_abs_difference_from_train_only_recompute": 0.0,
"representation": "Q1 adapter Relative-Progress projection; 50 slots; not physical-time alignment",
"adapter": "final.adapter.adapt_official_split; shared train-only robust scaler retained from Q2 V2",
"dimensions": [
768,
74,
35
],
"train_samples": 3395,
"valid_samples": 728,
"train_source_video_groups": 1528,
"valid_source_video_groups": 239,
"test_samples": 727,
"test_source_video_groups": 381,
"source_video_overlap_counts": {
"train/valid": 0,
"train/test": 0,
"valid/test": 0
},
"adapter_audit": {
"train": {
"method": "shared_interval_overlap_on_normalized_progress",
"coordinate_mode": "relative",
"physical_time_alignment": false,
"samples": 3395,
"vision_length_conflict_samples": 618,
"vision_tail_ambiguous_samples": 618,
"nonzero_text_rows_outside_attention": 86078,
"observed_target_rows": {
"text": 169750,
"audio": 169750,
"vision": 163302
},
"mean_target_coverage": {
"text": 1.0,
"audio": 1.0,
"vision": 0.9540337701314328
},
"quality_fields_available": false,
"word_or_frame_timestamps_available": false
},
"valid": {
"method": "shared_interval_overlap_on_normalized_progress",
"coordinate_mode": "relative",
"physical_time_alignment": false,
"samples": 728,
"vision_length_conflict_samples": 141,
"vision_tail_ambiguous_samples": 141,
"nonzero_text_rows_outside_attention": 17772,
"observed_target_rows": {
"text": 36400,
"audio": 36400,
"vision": 35315
},
"mean_target_coverage": {
"text": 1.0,
"audio": 1.0,
"vision": 0.96090314748523
},
"quality_fields_available": false,
"word_or_frame_timestamps_available": false
}
}
}
}
+23
View File
@@ -0,0 +1,23 @@
"""Q2 model entry points, one file per paper scheme."""
from .crg import CRG, StructuredGaussianImputer
from .c0 import C0
from .c1 import C1
from .c2 import C2
from .c3 import C3
from .c4 import C4
from .c5 import C5
from .c6 import C6
from .c6_no_distance import C6NoDistance
from .c6_no_reconstruction import C6NoReconstruction
from .c6_pointmask import C6PointMask
from .c7_distill import C7Distill
from .c7_group import C7Group
from .early_concat import AlignedFusionModel
from .mofe import MixtureOfFusionExperts
__all__ = [
"CRG", "StructuredGaussianImputer", "C0", "C1", "C2", "C3", "C4", "C5", "C6",
"C6NoDistance", "C6NoReconstruction", "C6PointMask", "C7Distill", "C7Group",
"AlignedFusionModel", "MixtureOfFusionExperts",
]
+358
View File
@@ -0,0 +1,358 @@
from __future__ import annotations
import math
from typing import Any
import torch
import torch.nn.functional as F
from torch import nn
from .ati_ho_config import ATIConfig
PAIR_INDICES = ((0, 1), (0, 2), (1, 2))
PAIR_NAMES = ("TA", "TV", "AV")
def _center_class_parameters(value: torch.Tensor) -> torch.Tensor:
"""Apply C to the three class logits while leaving magnitude parameters alone."""
logits = value[..., :3]
logits = logits - logits.mean(dim=-1, keepdim=True)
return torch.cat((logits, value[..., 3:]), dim=-1)
class PrivateTemporalEncoder(nn.Module):
"""One modality-private projection, BiGRU(32 each way), and attention pool."""
def __init__(self, input_dim: int, hidden: int, gru_hidden: int) -> None:
super().__init__()
self.projection = nn.Sequential(
nn.Linear(input_dim, hidden), nn.GELU(), nn.LayerNorm(hidden)
)
self.temporal = nn.GRU(
input_size=hidden,
hidden_size=gru_hidden,
num_layers=1,
batch_first=True,
bidirectional=True,
)
self.pool_score = nn.Linear(hidden, 1)
self.output_dim = 2 * gru_hidden
def forward(self, x: torch.Tensor, observed: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
observed = observed.bool()
projected = self.projection(x)
projected = projected * observed.unsqueeze(-1).to(projected.dtype)
sequence, _ = self.temporal(projected)
sequence = sequence * observed.unsqueeze(-1).to(sequence.dtype)
scores = self.pool_score(torch.tanh(sequence)).squeeze(-1)
scores = scores.masked_fill(~observed, torch.finfo(scores.dtype).min)
has_any = observed.any(dim=1, keepdim=True)
weights = torch.softmax(scores, dim=1)
weights = torch.where(has_any, weights, torch.zeros_like(weights))
pooled = torch.sum(sequence * weights.unsqueeze(-1), dim=1)
return sequence, pooled
class MainEffectHead(nn.Module):
def __init__(self, hidden: int) -> None:
super().__init__()
self.network = nn.Sequential(nn.Linear(hidden, hidden), nn.GELU(), nn.Linear(hidden, 5))
def forward(self, pooled: torch.Tensor) -> torch.Tensor:
return _center_class_parameters(self.network(pooled))
class AnchoredPairBranch(nn.Module):
"""A pair reads only two private streams; its four-term anchor is explicit."""
def __init__(self, hidden: int, config: ATIConfig) -> None:
super().__init__()
self.low_rank_enabled = config.low_rank
self.cross_attention_enabled = config.cross_attention
self.rank = config.rank
if self.low_rank_enabled:
self.left_factor = nn.Linear(hidden, config.rank)
self.right_factor = nn.Linear(hidden, config.rank)
self.low_rank_out = nn.Linear(config.rank, 5, bias=False)
else:
self.left_factor = None
self.right_factor = None
self.low_rank_out = None
if self.cross_attention_enabled:
self.left_to_right = nn.MultiheadAttention(
hidden, config.attention_heads, batch_first=True
)
self.right_to_left = nn.MultiheadAttention(
hidden, config.attention_heads, batch_first=True
)
self.left_norm1 = nn.LayerNorm(hidden)
self.right_norm1 = nn.LayerNorm(hidden)
self.left_ffn = nn.Sequential(
nn.Linear(hidden, config.attention_ffn),
nn.GELU(),
nn.Linear(config.attention_ffn, hidden),
)
self.right_ffn = nn.Sequential(
nn.Linear(hidden, config.attention_ffn),
nn.GELU(),
nn.Linear(config.attention_ffn, hidden),
)
self.left_norm2 = nn.LayerNorm(hidden)
self.right_norm2 = nn.LayerNorm(hidden)
self.cross_out = nn.Linear(hidden * 2, 5, bias=False)
else:
self.left_to_right = None
self.right_to_left = None
self.left_norm1 = None
self.right_norm1 = None
self.left_ffn = None
self.right_ffn = None
self.left_norm2 = None
self.right_norm2 = None
self.cross_out = None
init = min(max(config.eta_init, 1e-5), 1 - 1e-5)
self.eta_logit = nn.Parameter(torch.tensor(math.log(init / (1.0 - init))))
# q(x0,y)=q(x,y0)=q(x0,y0)=offset. Four-term subtraction cancels it.
# D0 deliberately leaves this offset in the output as a leakage control.
self.anchor_offset = nn.Parameter(torch.zeros(5))
@staticmethod
def _masked_mean(sequence: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
weights = mask.to(sequence.dtype).unsqueeze(-1)
return (sequence * weights).sum(dim=1) / weights.sum(dim=1).clamp_min(1.0)
@staticmethod
def _safe_key_mask(mask: torch.Tensor) -> torch.Tensor:
safe = mask.clone()
empty = ~safe.any(dim=1)
if empty.any():
safe[empty, 0] = True
return safe
def _core(
self,
left: torch.Tensor,
right: torch.Tensor,
left_mask: torch.Tensor,
right_mask: torch.Tensor,
) -> torch.Tensor:
joint = left_mask.bool() & right_mask.bool()
values: list[torch.Tensor] = []
if self.low_rank_enabled:
assert self.left_factor is not None and self.right_factor is not None
assert self.low_rank_out is not None
product = torch.tanh(self.left_factor(left)) * torch.tanh(self.right_factor(right))
values.append(self.low_rank_out(self._masked_mean(product, joint)))
if self.cross_attention_enabled:
assert self.left_to_right is not None and self.right_to_left is not None
assert self.left_norm1 is not None and self.right_norm1 is not None
assert self.left_ffn is not None and self.right_ffn is not None
assert self.left_norm2 is not None and self.right_norm2 is not None
assert self.cross_out is not None
safe_left = self._safe_key_mask(left_mask.bool())
safe_right = self._safe_key_mask(right_mask.bool())
left_msg, _ = self.left_to_right(
left, right, right, key_padding_mask=~safe_right, need_weights=False
)
right_msg, _ = self.right_to_left(
right, left, left, key_padding_mask=~safe_left, need_weights=False
)
left_context = self.left_norm1(left + left_msg)
right_context = self.right_norm1(right + right_msg)
left_context = self.left_norm2(left_context + self.left_ffn(left_context))
right_context = self.right_norm2(right_context + self.right_ffn(right_context))
left_context = left_context * left_mask.unsqueeze(-1).to(left_context.dtype)
right_context = right_context * right_mask.unsqueeze(-1).to(right_context.dtype)
pooled = torch.cat(
(self._masked_mean(left_context, joint), self._masked_mean(right_context, joint)),
dim=-1,
)
cross = self.cross_out(pooled)
values.append(torch.sigmoid(self.eta_logit) * cross)
if not values:
return left.new_zeros((left.shape[0], 5))
# Each branch has a bias-free output and a joint-observation gate. Thus
# core(x, y0)=core(x0, y)=core(x0, y0)=0 exactly.
return torch.stack(values, dim=0).sum(dim=0)
def forward(
self,
left: torch.Tensor,
right: torch.Tensor,
left_mask: torch.Tensor,
right_mask: torch.Tensor,
*,
anchored: bool,
) -> torch.Tensor:
raw_xy = self._core(left, right, left_mask, right_mask) + self.anchor_offset
if anchored:
# Four-term difference:
# q(x,y)-q(x,x0)-q(x0,y)+q(x0,y0) = core(x,y).
# The three absent-modality terms equal anchor_offset by the
# joint gate and bias-free core, so they cancel algebraically.
value = raw_xy - self.anchor_offset
else:
value = raw_xy
return _center_class_parameters(value)
class ATIHOModel(nn.Module):
"""Five-parameter additive multimodal predictor with exact modality anchors."""
def __init__(self, dims: tuple[int, int, int], config: ATIConfig, steps: int = 50) -> None:
super().__init__()
self.dims = tuple(int(d) for d in dims)
self.steps = int(steps)
self.config = config
hidden = config.hidden
self.encoders = nn.ModuleList(
PrivateTemporalEncoder(dim, hidden, config.gru_hidden_per_direction) for dim in dims
)
self.main_heads = nn.ModuleList(MainEffectHead(hidden) for _ in dims)
self.mask_heads = nn.ModuleList(nn.Linear(hidden, 1) for _ in dims)
self.pair_branches = nn.ModuleList(
AnchoredPairBranch(hidden, config) for _ in PAIR_INDICES
)
self.baseline = nn.Parameter(torch.zeros(5))
def forward(
self,
xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
masks: torch.Tensor,
*,
return_details: bool = True,
) -> dict[str, Any]:
if len(xs) != 3:
raise ValueError("ATI–HO requires text, audio, and vision streams")
if masks.ndim != 3 or masks.shape[-1] != 3:
raise ValueError(f"masks must be B x T x 3, got {tuple(masks.shape)}")
if masks.shape[1] > self.steps:
raise ValueError(f"ATI–HO supports at most {self.steps} steps")
masks = masks.bool()
sequences: list[torch.Tensor] = []
pooled: list[torch.Tensor] = []
mask_logits: list[torch.Tensor] = []
main_effects: list[torch.Tensor] = []
for modality, (encoder, head, mask_head, x) in enumerate(
zip(self.encoders, self.main_heads, self.mask_heads, xs)
):
if x.shape[-1] != self.dims[modality]:
raise ValueError(
f"modality {modality} has {x.shape[-1]} features, expected {self.dims[modality]}"
)
sequence, representation = encoder(x, masks[..., modality])
# Missing-mask baseline has a zero pooled representation. Explicit
# subtraction makes every main effect zero at that baseline.
baseline_raw = head(representation.new_zeros(representation.shape))
effect = _center_class_parameters(head(representation) - baseline_raw)
sequences.append(sequence)
pooled.append(representation)
mask_logits.append(mask_head(sequence).squeeze(-1))
main_effects.append(effect)
pair_effects: list[torch.Tensor] = []
pair_penalties: list[torch.Tensor] = []
for branch, (left_idx, right_idx) in zip(self.pair_branches, PAIR_INDICES):
pair = branch(
sequences[left_idx],
sequences[right_idx],
masks[..., left_idx],
masks[..., right_idx],
anchored=self.config.anchored,
)
pair_effects.append(pair)
pair_penalties.append(pair.square().mean())
main_tensor = torch.stack(main_effects, dim=1)
pair_tensor = torch.stack(pair_effects, dim=1)
params = self.baseline.unsqueeze(0) + main_tensor.sum(dim=1) + pair_tensor.sum(dim=1)
logits = params[:, :3]
probabilities = torch.softmax(logits, dim=-1)
nu_negative = 3.0 * torch.sigmoid(params[:, 3])
nu_positive = 3.0 * torch.sigmoid(params[:, 4])
predicted_class = logits.argmax(dim=-1)
hard_intensity = torch.where(
predicted_class == 0,
-nu_negative,
torch.where(predicted_class == 2, nu_positive, torch.zeros_like(nu_positive)),
)
soft_intensity = probabilities[:, 2] * nu_positive - probabilities[:, 0] * nu_negative
result: dict[str, Any] = {
"logits": logits,
"probabilities": probabilities,
"predicted_class": predicted_class,
"intensity": hard_intensity,
"soft_intensity": soft_intensity,
"nu_negative": nu_negative,
"nu_positive": nu_positive,
"params": params,
"interaction_penalty": torch.stack(pair_penalties).mean(),
"mask_logits": torch.stack(mask_logits, dim=-1),
}
if return_details:
result.update(
{
"baseline": self.baseline.unsqueeze(0).expand(xs[0].shape[0], -1),
"main_effects": main_tensor,
"pair_effects": pair_tensor,
"main_sequences": torch.stack(sequences, dim=1),
}
)
return result
def task_loss(
output: dict[str, Any],
y_cls: torch.Tensor,
y_reg: torch.Tensor,
*,
lambda_interaction: float,
lambda_mask: float,
mask_target: torch.Tensor | None = None,
) -> tuple[torch.Tensor, dict[str, torch.Tensor]]:
"""CE + conditional polarity magnitude + low-weight continuous Huber."""
class_loss = F.cross_entropy(output["logits"], y_cls)
negative = y_reg < 0
positive = y_reg > 0
target_mag = torch.abs(y_reg) / 3.0
magnitude_parts: list[torch.Tensor] = []
if negative.any():
magnitude_parts.append(
F.smooth_l1_loss(output["nu_negative"][negative] / 3.0, target_mag[negative])
)
if positive.any():
magnitude_parts.append(
F.smooth_l1_loss(output["nu_positive"][positive] / 3.0, target_mag[positive])
)
magnitude_loss = torch.stack(magnitude_parts).mean() if magnitude_parts else class_loss.new_zeros(())
continuous_loss = F.huber_loss(
output["soft_intensity"] / 3.0, y_reg / 3.0, delta=0.25
)
interaction_loss = output["interaction_penalty"]
mask_loss = class_loss.new_zeros(())
if lambda_mask > 0:
if mask_target is None:
raise ValueError("mask_target is required when the visibility-mask auxiliary loss is enabled")
mask_loss = F.binary_cross_entropy_with_logits(
output["mask_logits"], mask_target.to(output["mask_logits"].dtype)
)
total = (
class_loss
+ magnitude_loss
+ 0.2 * continuous_loss
+ lambda_interaction * interaction_loss
+ lambda_mask * mask_loss
)
parts = {
"classification": class_loss,
"conditional_magnitude": magnitude_loss,
"continuous_huber": continuous_loss,
"interaction": interaction_loss,
"visibility_mask": mask_loss,
"total": total,
}
return total, parts
+35
View File
@@ -0,0 +1,35 @@
from __future__ import annotations
from dataclasses import asdict, dataclass
@dataclass(frozen=True)
class ATIConfig:
name: str
low_rank: bool = False
cross_attention: bool = False
anchored: bool = True
lambda_interaction: float = 1e-3
lambda_mask: float = 0.0
rank: int = 4
hidden: int = 64
gru_hidden_per_direction: int = 32
attention_heads: int = 4
attention_ffn: int = 128
eta_init: float = 0.1
def to_dict(self) -> dict[str, object]:
return asdict(self)
CONFIGS: dict[str, ATIConfig] = {
"A0": ATIConfig(name="A0_main_effects"),
"A1": ATIConfig(name="A1_low_rank_pairs", low_rank=True),
"A2": ATIConfig(name="A2_anchored_pairwise", low_rank=True, cross_attention=True),
"A3": ATIConfig(
name="A3_pairwise_mask_aux", low_rank=True, cross_attention=True, lambda_mask=0.05
),
"D0": ATIConfig(
name="D0_unanchored_diagnostic", low_rank=True, cross_attention=True, anchored=False
),
}
+57
View File
@@ -0,0 +1,57 @@
"""C0: observed statistics and mask baseline from E题V2, table 5.9."""
from __future__ import annotations
import numpy as np
from sklearn.linear_model import LogisticRegression, Ridge
MODALITIES = ("text", "audio", "vision")
def sample_statistics(arrays: dict[str, np.ndarray], mask: np.ndarray) -> np.ndarray:
"""Mean, standard deviation, missing fraction and longest gap per modality."""
mask = np.asarray(mask, dtype=bool)
if mask.ndim != 3 or mask.shape[-1] != 3:
raise ValueError("mask must have shape (N, T, 3)")
parts = []
for index, name in enumerate(MODALITIES):
x = np.asarray(arrays[name], dtype=np.float32)
if x.shape[:2] != mask.shape[:2]:
raise ValueError(f"{name}: feature and mask shapes disagree")
visible = mask[:, :, index]
count = visible.sum(axis=1, keepdims=True)
mean = (x * visible[:, :, None]).sum(axis=1) / np.maximum(count, 1)
variance = (((x - mean[:, None, :]) ** 2) * visible[:, :, None]).sum(axis=1) / np.maximum(count, 1)
missing = 1.0 - visible.mean(axis=1, keepdims=True)
max_gap = []
for row in visible:
longest = current = 0
for observed in row:
current = 0 if observed else current + 1
longest = max(longest, current)
max_gap.append(longest / max(len(row), 1))
parts.extend((mean, np.sqrt(variance), missing, np.asarray(max_gap, np.float32)[:, None]))
return np.concatenate(parts, axis=1).astype(np.float32)
class C0:
"""Logistic polarity classifier and Ridge intensity regressor."""
def __init__(self) -> None:
self.classifier = LogisticRegression(C=0.05, max_iter=2500, random_state=20260924)
self.regressor = Ridge(alpha=25.0)
def fit(self, arrays: dict[str, np.ndarray], mask: np.ndarray,
polarity: np.ndarray, intensity: np.ndarray) -> "C0":
features = sample_statistics(arrays, mask)
self.classifier.fit(features, polarity)
self.regressor.fit(features, intensity)
return self
def predict(self, arrays: dict[str, np.ndarray], mask: np.ndarray) -> dict[str, np.ndarray]:
features = sample_statistics(arrays, mask)
probabilities = np.zeros((len(features), 3), np.float64)
probabilities[:, self.classifier.classes_] = self.classifier.predict_proba(features)
return {
"probabilities": probabilities,
"intensity": np.clip(self.regressor.predict(features), -3.0, 3.0),
}
+9
View File
@@ -0,0 +1,9 @@
"""C1: masked BiGRU, without probabilistic completion or explicit gates."""
from .crg import CRG
class C1(CRG):
def __init__(self, **kwargs):
super().__init__(use_imputer=False, use_joint_draws=False,
use_final_gate=False, use_source_attention=False,
reliability_update=False, use_low_rank=False, **kwargs)
+9
View File
@@ -0,0 +1,9 @@
"""C2: Gaussian posterior mean completion, without trajectory integration."""
from .crg import CRG
class C2(CRG):
def __init__(self, **kwargs):
super().__init__(use_imputer=True, use_joint_draws=False,
use_final_gate=False, use_source_attention=False,
reliability_update=False, use_low_rank=False, **kwargs)
+9
View File
@@ -0,0 +1,9 @@
"""C3: joint trajectory integration and final reliability/content fusion."""
from .crg import CRG
class C3(CRG):
def __init__(self, **kwargs):
super().__init__(use_imputer=True, use_joint_draws=True,
use_final_gate=True, use_source_attention=False,
reliability_update=False, use_low_rank=False, **kwargs)
+9
View File
@@ -0,0 +1,9 @@
"""C4: C3 with bounded source attention and a null source."""
from .crg import CRG
class C4(CRG):
def __init__(self, **kwargs):
super().__init__(use_imputer=True, use_joint_draws=True,
use_final_gate=True, use_source_attention=True,
reliability_update=False, use_low_rank=False, **kwargs)
+9
View File
@@ -0,0 +1,9 @@
"""C5: C4 with reliability-modulated recurrent updates."""
from .crg import CRG
class C5(CRG):
def __init__(self, **kwargs):
super().__init__(use_imputer=True, use_joint_draws=True,
use_final_gate=True, use_source_attention=True,
reliability_update=True, use_low_rank=False, **kwargs)
+9
View File
@@ -0,0 +1,9 @@
"""C6: C5 with the optional rank-four CP interaction residual enabled."""
from .crg import CRG
class C6(CRG):
def __init__(self, **kwargs):
super().__init__(use_imputer=True, use_joint_draws=True,
use_final_gate=True, use_source_attention=True,
reliability_update=True, use_low_rank=True, **kwargs)
+7
View File
@@ -0,0 +1,7 @@
"""C6 diagnostic: disable uncertainty distance and span penalties."""
from .c6 import C6
class C6NoDistance(C6):
def __init__(self, **kwargs):
super().__init__(reliability_hparams=(0.5, 0.05, 0.0, 0.0), **kwargs)
@@ -0,0 +1,8 @@
"""C6 diagnostic: omit auxiliary hidden-feature reconstruction while fitting."""
from .c6 import C6
class C6NoReconstruction(C6):
"""Use the C6 forward pass and set reconstruction loss weight to zero."""
reconstruction_loss_weight = 0.0
+8
View File
@@ -0,0 +1,8 @@
"""C6 diagnostic: train with independent point masks instead of spans."""
from .c6 import C6
class C6PointMask(C6):
"""Use the C6 forward pass with mask_kind='point' during fitting."""
training_mask_kind = "point"
+41
View File
@@ -0,0 +1,41 @@
"""C7 distillation-only branch: C6 architecture plus teacher loss."""
from __future__ import annotations
import math
import numpy as np
import torch
from torch.nn import functional as F
from .c6 import C6
DISTILL_TEMPERATURE = 2.0
DISTILL_WEIGHT = 0.1
class C7Distill(C6):
"""Inference uses C6; training adds weighted teacher distillation."""
def distillation_per_sample(student: dict, teacher: dict,
original: np.ndarray, current: np.ndarray) -> torch.Tensor:
"""Entropy/retention-weighted KL and score term for Q2 distillation."""
temp = DISTILL_TEMPERATURE
p_teacher = teacher["tempered_probs_by_path"].mean(dim=0).detach().clamp_min(1e-8)
p_student = student["tempered_probs_by_path"].mean(dim=0).clamp_min(1e-8)
entropy = -(p_teacher * p_teacher.log()).sum(dim=-1)
confidence_weight = (1.0 - entropy / math.log(3.0)).clamp(0.0, 1.0)
orig_t = torch.as_tensor(original, device=p_teacher.device, dtype=torch.float32)
curr_t = torch.as_tensor(current, device=p_teacher.device, dtype=torch.float32)
retained = []
for modality in range(3):
denominator = orig_t[:, :, modality].sum(dim=1)
ratio = (orig_t[:, :, modality] * curr_t[:, :, modality]).sum(dim=1) / denominator.clamp_min(1.0)
retained.append(torch.where(denominator > 0, ratio, torch.ones_like(ratio)))
weight = confidence_weight * torch.stack(retained, dim=-1).mean(dim=-1)
kl = (p_teacher * (p_teacher.log() - p_student.log())).sum(dim=-1) * temp * temp
teacher_score = teacher["mixed_score"].detach()
student_score = student["mixed_score"]
regression = F.huber_loss((teacher_score - student_score) / 3.0,
torch.zeros_like(teacher_score), reduction="none", delta=0.25)
return weight * (kl + regression)
+32
View File
@@ -0,0 +1,32 @@
"""C7 group-risk-only branch: C6 architecture plus smooth worst-group loss."""
from __future__ import annotations
import numpy as np
import torch
from .c6 import C6
class C7Group(C6):
"""Inference uses C6; training adds smooth worst-group risk."""
def smooth_group_risk(losses: torch.Tensor, group_ids: np.ndarray,
lambda_group: float = 0.1,
group_temperature: float = 0.05) -> torch.Tensor:
"""Match the selected group penalty from the Q2 training protocol."""
if group_temperature <= 0 or not 0 <= lambda_group <= 1:
raise ValueError("invalid group risk parameters")
groups = torch.as_tensor(group_ids, device=losses.device, dtype=torch.long)
if groups.shape != losses.shape:
raise ValueError("group_ids must match per-sample losses")
group_losses, priors = [], []
for group in torch.unique(groups):
selected = groups == group
group_losses.append(losses[selected].mean())
priors.append(selected.float().mean())
values = torch.stack(group_losses)
prior = torch.stack(priors).clamp_min(1e-8)
expected = (prior * values).sum()
worst = group_temperature * torch.logsumexp(torch.log(prior) + values / group_temperature, dim=0)
return (1.0 - lambda_group) * expected + lambda_group * worst
+556
View File
@@ -0,0 +1,556 @@
"""Structured Gaussian imputation and reliability-aware CRG sequence model."""
from __future__ import annotations
import math
from typing import Sequence
import torch
from torch import nn
from torch.nn import functional as F
MODALITIES = ("text", "audio", "vision")
INPUT_DIMS = (768, 74, 35)
HIDDEN = 32
SHARED_STATE = 8
PRIVATE_STATE = 4
STATE_DIM = SHARED_STATE + len(MODALITIES) * PRIVATE_STATE
def _inv_softplus(value: float) -> float:
return math.log(math.expm1(value))
class StructuredGaussianImputer(nn.Module):
"""Linear-Gaussian shared/private state model with exact block-Gaussian inference.
The state is [shared(8), text-private(4), audio-private(4), vision-private(4)].
Each modality emits from the shared state and its own private state only. The
filtering likelihood uses the matrix determinant lemma, retaining its log-det
normalization without forming a covariance matrix in observation space.
"""
def __init__(self, input_dims: Sequence[int] = INPUT_DIMS) -> None:
super().__init__()
self.input_dims = tuple(int(x) for x in input_dims)
self.state_dim = STATE_DIM
transition_mask = torch.zeros(STATE_DIM, STATE_DIM)
blocks = [slice(0, SHARED_STATE)] + [
slice(SHARED_STATE + i * PRIVATE_STATE, SHARED_STATE + (i + 1) * PRIVATE_STATE)
for i in range(len(MODALITIES))
]
for block in blocks:
transition_mask[block, block] = 1.0
self.register_buffer("transition_mask", transition_mask)
self.transition_raw = nn.Parameter(0.8 * torch.eye(STATE_DIM))
self.mu0 = nn.Parameter(torch.zeros(STATE_DIM))
self.pi0_raw = nn.Parameter(torch.full((STATE_DIM,), _inv_softplus(1.0)))
self.q_raw = nn.Parameter(torch.full((STATE_DIM,), _inv_softplus(0.08)))
self.emission_raw = nn.ParameterList()
self.biases = nn.ParameterList()
self.r_raw = nn.ParameterList()
for index, dim in enumerate(self.input_dims):
mask = torch.zeros(dim, STATE_DIM)
mask[:, :SHARED_STATE] = 1.0
private_start = SHARED_STATE + index * PRIVATE_STATE
mask[:, private_start:private_start + PRIVATE_STATE] = 1.0
self.register_buffer(f"emission_mask_{index}", mask)
self.emission_raw.append(nn.Parameter(torch.randn(dim, STATE_DIM) * 0.025))
self.biases.append(nn.Parameter(torch.zeros(dim)))
self.r_raw.append(nn.Parameter(torch.full((dim,), _inv_softplus(0.5))))
def _transition(self) -> torch.Tensor:
matrix = self.transition_raw * self.transition_mask
norm = torch.linalg.matrix_norm(matrix, ord=2).clamp_min(1e-8)
return matrix * torch.clamp(0.98 / norm, max=1.0)
def _covariances(self) -> tuple[torch.Tensor, torch.Tensor]:
eye = torch.eye(self.state_dim, device=self.mu0.device, dtype=self.mu0.dtype)
p0 = torch.diag(F.softplus(self.pi0_raw) + 1e-4) + 1e-5 * eye
q = torch.diag(F.softplus(self.q_raw) + 1e-4) + 1e-5 * eye
return p0, q
def emissions(self) -> list[torch.Tensor]:
return [raw * getattr(self, f"emission_mask_{i}") for i, raw in enumerate(self.emission_raw)]
def _filter(
self,
xs: Sequence[torch.Tensor],
observed: torch.Tensor,
*,
calculate_log_likelihood: bool,
retain_states: bool,
) -> tuple[torch.Tensor | None, dict[str, list[torch.Tensor]] | None]:
# xs[m]: [B,T,Dm], observed: [B,T,3]
batch, steps, _ = observed.shape
transition = self._transition()
p0, process_noise = self._covariances()
emissions = self.emissions()
noise = [F.softplus(x) + 1e-4 for x in self.r_raw]
mu_prior = self.mu0.expand(batch, -1)
p_prior = p0.expand(batch, -1, -1)
total_nll = torch.zeros(batch, device=observed.device, dtype=mu_prior.dtype)
prior_means: list[torch.Tensor] = []
prior_covs: list[torch.Tensor] = []
filtered_means: list[torch.Tensor] = []
filtered_covs: list[torch.Tensor] = []
for t in range(steps):
if retain_states:
prior_means.append(mu_prior)
prior_covs.append(p_prior)
p_chol = torch.linalg.cholesky(p_prior + 1e-6 * torch.eye(self.state_dim, device=p_prior.device))
p_inv = torch.cholesky_inverse(p_chol)
information_parts: list[torch.Tensor] = []
vector_parts: list[torch.Tensor] = []
quadratic_parts: list[torch.Tensor] = []
logdet_r = torch.zeros(batch, device=p_prior.device, dtype=p_prior.dtype)
n_observed = torch.zeros_like(logdet_r)
for m, (x, emission, variance) in enumerate(zip(xs, emissions, noise)):
active = observed[:, t, m].to(dtype=mu_prior.dtype)
weights = active[:, None] / variance[None, :]
centered = x[:, t] - self.biases[m]
residual = centered - mu_prior @ emission.T
information_parts.append(torch.einsum("di,bd,dj->bij", emission, weights, emission))
vector_parts.append((residual * weights) @ emission)
quadratic_parts.append((residual.square() * weights).sum(dim=-1))
logdet_r = logdet_r + active * torch.log(variance).sum()
n_observed = n_observed + active * x.shape[-1]
information = torch.stack(information_parts).sum(dim=0)
innovation = torch.stack(vector_parts).sum(dim=0)
precision = p_inv + information
precision_chol = torch.linalg.cholesky(precision + 1e-6 * torch.eye(self.state_dim, device=precision.device))
p_filtered = torch.cholesky_inverse(precision_chol)
mu_filtered = mu_prior + torch.einsum("bij,bj->bi", p_filtered, innovation)
if calculate_log_likelihood:
logdet_p = 2.0 * torch.log(torch.diagonal(p_chol, dim1=-2, dim2=-1)).sum(dim=-1)
logdet_precision = 2.0 * torch.log(torch.diagonal(precision_chol, dim1=-2, dim2=-1)).sum(dim=-1)
quad = torch.stack(quadratic_parts).sum(dim=0)
correction = torch.einsum("bi,bij,bj->b", innovation, p_filtered, innovation)
log_likelihood = logdet_r + logdet_p + logdet_precision + (quad - correction).clamp_min(0.0)
log_likelihood = log_likelihood + n_observed * math.log(2.0 * math.pi)
total_nll = total_nll + 0.5 * log_likelihood
if retain_states:
filtered_means.append(mu_filtered)
filtered_covs.append(p_filtered)
mu_prior = mu_filtered @ transition.T
p_prior = transition @ p_filtered @ transition.T + process_noise
states = None
if retain_states:
states = {
"prior_mean": prior_means,
"prior_cov": prior_covs,
"filtered_mean": filtered_means,
"filtered_cov": filtered_covs,
"transition": [transition],
}
return (total_nll if calculate_log_likelihood else None), states
def observed_nll(self, xs: Sequence[torch.Tensor], observed: torch.Tensor) -> torch.Tensor:
"""Exact observed-data Gaussian NLL, including covariance log determinants."""
nll, _ = self._filter(xs, observed, calculate_log_likelihood=True, retain_states=False)
assert nll is not None
return nll
@staticmethod
def _draw(mean: torch.Tensor, covariance: torch.Tensor, paths: int) -> torch.Tensor:
chol = torch.linalg.cholesky(covariance + 1e-5 * torch.eye(covariance.shape[-1], device=covariance.device))
noise = torch.randn((paths, *mean.shape), dtype=mean.dtype, device=mean.device)
return mean.unsqueeze(0) + torch.einsum("bij,kbj->kbi", chol, noise)
@torch.no_grad()
def complete(
self,
xs: Sequence[torch.Tensor],
observed: torch.Tensor,
paths: int,
*,
joint_draws: bool,
) -> tuple[list[torch.Tensor], list[torch.Tensor]]:
"""RTS smooth, draw joint latent trajectories, then draw missing emissions."""
_, stored = self._filter(xs, observed, calculate_log_likelihood=False, retain_states=True)
assert stored is not None
fm, fc = stored["filtered_mean"], stored["filtered_cov"]
pm, pc = stored["prior_mean"], stored["prior_cov"]
transition = stored["transition"][0]
steps = len(fm)
smoother_gains: list[torch.Tensor] = [torch.empty(0, device=observed.device)] * max(0, steps - 1)
smooth_cov: list[torch.Tensor] = [torch.empty(0, device=observed.device)] * steps
smooth_cov[-1] = fc[-1]
for t in range(steps - 2, -1, -1):
next_chol = torch.linalg.cholesky(pc[t + 1] + 1e-6 * torch.eye(self.state_dim, device=observed.device))
gain = torch.cholesky_solve((fc[t] @ transition.T).transpose(-1, -2), next_chol).transpose(-1, -2)
smoother_gains[t] = gain
smooth_cov[t] = fc[t] + gain @ (smooth_cov[t + 1] - pc[t + 1]) @ gain.transpose(-1, -2)
smooth_cov[t] = 0.5 * (smooth_cov[t] + smooth_cov[t].transpose(-1, -2))
if joint_draws:
state = torch.empty((paths, observed.shape[0], steps, self.state_dim), device=observed.device, dtype=fm[0].dtype)
state[:, :, -1] = self._draw(fm[-1], fc[-1], paths)
for t in range(steps - 2, -1, -1):
gain = smoother_gains[t]
conditional_mean = fm[t].unsqueeze(0) + torch.einsum(
"bij,kbj->kbi", gain, state[:, :, t + 1] - pm[t + 1].unsqueeze(0)
)
conditional_cov = fc[t] - gain @ pc[t + 1] @ gain.transpose(-1, -2)
conditional_cov = 0.5 * (conditional_cov + conditional_cov.transpose(-1, -2))
chol = torch.linalg.cholesky(conditional_cov + 1e-5 * torch.eye(self.state_dim, device=observed.device))
eps = torch.randn_like(conditional_mean)
state[:, :, t] = conditional_mean + torch.einsum("bij,kbj->kbi", chol, eps)
else:
means = torch.stack(fm, dim=1)
covs = torch.stack(smooth_cov, dim=1)
smoothed_means = [fm[-1]] * steps
smoothed_means[-1] = fm[-1]
for t in range(steps - 2, -1, -1):
smoothed_means[t] = fm[t] + torch.einsum(
"bij,bj->bi", smoother_gains[t], smoothed_means[t + 1] - pm[t + 1]
)
state = torch.stack(smoothed_means, dim=1).unsqueeze(0).expand(paths, -1, -1, -1)
completed: list[torch.Tensor] = []
variances: list[torch.Tensor] = []
for m, (x, emission) in enumerate(zip(xs, self.emissions())):
mean = torch.einsum("kbti,di->kbtd", state, emission) + self.biases[m]
if joint_draws:
noise = torch.randn_like(mean) * torch.sqrt(F.softplus(self.r_raw[m]) + 1e-4)
draws = mean + noise
else:
draws = mean
visible = observed[:, :, m].unsqueeze(0).unsqueeze(-1)
completed.append(torch.where(visible, x.unsqueeze(0), draws))
projected_cov = torch.einsum("di,btij,dj->btd", emission, torch.stack(smooth_cov, dim=1), emission)
variance = projected_cov + (F.softplus(self.r_raw[m]) + 1e-4)
variances.append(torch.where(observed[:, :, m, None], torch.zeros_like(variance), variance.clamp_min(1e-6)))
return completed, variances
class ReliabilityGRU(nn.Module):
"""One-layer BiGRU with directional time decay and rho-scaled updates."""
def __init__(self, input_dim: int, hidden: int = 16) -> None:
super().__init__()
self.hidden = hidden
self.x_proj = nn.Linear(input_dim, 3 * hidden)
self.h_proj = nn.Linear(hidden, 2 * hidden, bias=False)
self.candidate_h = nn.Linear(hidden, hidden, bias=False)
self.decay_raw = nn.Parameter(torch.full((hidden,), -3.0))
def _one_direction(
self,
x: torch.Tensor,
rho: torch.Tensor,
distance: torch.Tensor,
reverse: bool,
reliability_update: bool,
) -> torch.Tensor:
batch, steps, _ = x.shape
state = torch.zeros(batch, self.hidden, dtype=x.dtype, device=x.device)
x_parts = self.x_proj(x).chunk(3, dim=-1)
output: list[torch.Tensor | None] = [None] * steps
indices = range(steps - 1, -1, -1) if reverse else range(steps)
for t in indices:
if reliability_update:
decay = torch.exp(-F.softplus(self.decay_raw)[None, :] * distance[:, t:t + 1])
decayed_state = decay * state
else:
decayed_state = state
hz, hr = self.h_proj(decayed_state).chunk(2, dim=-1)
z = torch.sigmoid(x_parts[0][:, t] + hz)
r = torch.sigmoid(x_parts[1][:, t] + hr)
candidate = torch.tanh(x_parts[2][:, t] + self.candidate_h(r * decayed_state))
effective_z = rho[:, t:t + 1] * z if reliability_update else z
state = (1.0 - effective_z) * decayed_state + effective_z * candidate
output[t] = state
return torch.stack([v for v in output if v is not None], dim=1)
def forward(
self,
x: torch.Tensor,
rho: torch.Tensor,
dminus: torch.Tensor,
dplus: torch.Tensor,
reliability_update: bool,
) -> torch.Tensor:
if not reliability_update:
rho = torch.ones_like(rho)
return torch.cat((
self._one_direction(x, rho, dminus, False, reliability_update),
self._one_direction(x, rho, dplus, True, reliability_update),
), dim=-1)
class CRG(nn.Module):
"""Quality-aware multimodal sequence predictor for a configured ablation."""
def __init__(
self,
imputer: StructuredGaussianImputer | None = None,
input_dims: Sequence[int] = INPUT_DIMS,
*,
use_imputer: bool = True,
use_joint_draws: bool = True,
use_final_gate: bool = True,
use_source_attention: bool = True,
reliability_update: bool = True,
use_low_rank: bool = True,
reliability_hparams: tuple[float, float, float, float] = (0.5, 0.05, 0.05, 0.05),
) -> None:
super().__init__()
self.use_imputer = use_imputer
self.use_joint_draws = use_joint_draws
self.use_final_gate = use_final_gate
self.use_source_attention = use_source_attention
self.reliability_update = reliability_update
self.use_low_rank = use_low_rank
self.imputer = imputer if imputer is not None else StructuredGaussianImputer(input_dims)
self.projections = nn.ModuleList(
nn.Sequential(nn.Linear(d, HIDDEN), nn.LayerNorm(HIDDEN), nn.GELU()) for d in input_dims
)
recurrent_input = HIDDEN + 13
self.temporal = nn.ModuleList(ReliabilityGRU(recurrent_input, 16) for _ in MODALITIES)
rho_imp, lambda_u, lambda_gap, lambda_span = reliability_hparams
if not 0.0 < rho_imp < 1.0 or min(lambda_u, lambda_gap, lambda_span) < 0.0:
raise ValueError("reliability requires 0<rho_imp<1 and nonnegative distance/uncertainty penalties")
self.register_buffer("rho_imp", torch.tensor(float(rho_imp)))
self.register_buffer("rel_u", torch.full((3,), float(lambda_u)))
self.register_buffer("rel_gap", torch.full((3,), float(lambda_gap)))
self.register_buffer("rel_span", torch.full((3,), float(lambda_span)))
self.query = nn.Linear(HIDDEN, HIDDEN, bias=False)
self.key = nn.Linear(HIDDEN, HIDDEN, bias=False)
self.value = nn.Linear(HIDDEN, HIDDEN, bias=False)
self.relative_bias = nn.Embedding(99, 1)
nn.init.zeros_(self.relative_bias.weight)
self.cross_base = nn.Linear(HIDDEN, HIDDEN)
self.cross_out = nn.Linear(HIDDEN, HIDDEN, bias=False)
self.cross_eta_logit = nn.Parameter(torch.tensor(-1.0))
self.content_score = nn.Sequential(nn.Linear(HIDDEN, 16), nn.Tanh(), nn.Linear(16, 1, bias=False))
self.null_expert = nn.Parameter(torch.zeros(HIDDEN))
self.pool_hidden = nn.Linear(HIDDEN, 16)
self.pool_score = nn.Linear(16, 1, bias=False)
self.reconstruction_heads = nn.ModuleList(nn.Linear(HIDDEN, d) for d in input_dims)
self.cp_factors = nn.ModuleList(nn.Linear(HIDDEN + 1, 4, bias=False) for _ in MODALITIES)
self.cp_output = nn.Parameter(torch.randn(4, HIDDEN) * 0.02)
self.low_rank_output = nn.Linear(HIDDEN, HIDDEN, bias=False)
self.low_rank_eta_logit = nn.Parameter(torch.tensor(-4.0))
self.head = nn.Sequential(nn.Linear(HIDDEN + 18, 64), nn.GELU(), nn.Dropout(0.2))
self.classifier = nn.Linear(64, 3)
self.magnitude_mean = nn.Linear(64, 2)
self.concentration_raw = nn.Parameter(torch.full((2,), _inv_softplus(6.0)))
@staticmethod
def _gap_features(observed: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
# Time positions are valid sequence locations even when all three sources are missing.
batch, steps, modalities = observed.shape
device = observed.device
positions = torch.arange(steps, device=device).view(1, steps).expand(batch, -1)
previous = torch.full((batch, modalities), -1, device=device, dtype=torch.long)
before, before_edge = [], []
for t in range(steps):
before_edge.append(previous < 0)
before.append(torch.where(previous < 0, torch.ones_like(previous, dtype=torch.float32), (t - previous).float() / max(1, steps - 1)))
previous = torch.where(observed[:, t], torch.full_like(previous, t), previous)
following = torch.full((batch, modalities), steps, device=device, dtype=torch.long)
after, after_edge = [None] * steps, [None] * steps
for t in range(steps - 1, -1, -1):
after_edge[t] = following >= steps
after[t] = torch.where(following >= steps, torch.ones_like(following, dtype=torch.float32), (following - t).float() / max(1, steps - 1))
following = torch.where(observed[:, t], torch.full_like(following, t), following)
dminus = torch.stack(before, dim=1)
dplus = torch.stack([x for x in after if x is not None], dim=1)
edge_minus = torch.stack(before_edge, dim=1)
edge_plus = torch.stack([x for x in after_edge if x is not None], dim=1)
dminus = torch.where(observed, torch.zeros_like(dminus), dminus)
dplus = torch.where(observed, torch.zeros_like(dplus), dplus)
edge_minus = edge_minus & ~observed
edge_plus = edge_plus & ~observed
missing = ~observed
left_run = torch.zeros((batch, steps, modalities), device=device, dtype=torch.float32)
run = torch.zeros((batch, modalities), device=device, dtype=torch.float32)
for t in range(steps):
run = torch.where(missing[:, t], run + 1.0, torch.zeros_like(run))
left_run[:, t] = run
right_run = torch.zeros_like(left_run)
run.zero_()
for t in range(steps - 1, -1, -1):
run = torch.where(missing[:, t], run + 1.0, torch.zeros_like(run))
right_run[:, t] = run
span = torch.where(missing, (left_run + right_run - 1.0) / max(1, steps), torch.zeros_like(left_run))
return dminus, dplus, span, torch.stack((edge_minus, edge_plus), dim=-1).float()
def _reliability(
self, observed: torch.Tensor, uncertainty: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
dminus, dplus, span, edges = self._gap_features(observed)
gap = torch.minimum(dminus, dplus)
gap = torch.where(observed, torch.zeros_like(gap), gap)
u = torch.where(observed, torch.zeros_like(uncertainty), uncertainty).clamp_min(0.0)
qstar = observed.float() # External Q2 quality is unavailable: q*=1 only for visible rows; J=0.
rho_missing = self.rho_imp.clamp(1e-4, 0.999) * torch.exp(
-self.rel_u[None, None, :] * u
-self.rel_gap[None, None, :] * gap
-self.rel_span[None, None, :] * span
)
rho = torch.where(observed, qstar, rho_missing).clamp(1e-4, 1.0)
return rho, u, gap, span, dminus, dplus, torch.cat((qstar.unsqueeze(-1), torch.zeros_like(qstar).unsqueeze(-1), edges), dim=-1)
def _cross_source(self, hidden: torch.Tensor, rho: torch.Tensor) -> torch.Tensor:
# hidden [B,T,M,H]; each query reads every legal time in each other source.
batch, steps, modalities, width = hidden.shape
outputs = []
q = self.query(hidden)
k = self.key(hidden)
v = torch.tanh(self.value(hidden))
loc = torch.arange(steps, device=hidden.device)
relative_index = (loc[None, :] - loc[:, None] + 49).clamp(0, 98)
relative = self.relative_bias(relative_index).squeeze(-1)
for target in range(modalities):
numerator = torch.zeros((batch, steps, width), device=hidden.device, dtype=hidden.dtype)
denominator = torch.ones((batch, steps, 1), device=hidden.device, dtype=hidden.dtype)
for source in range(modalities):
if source == target:
continue
raw = torch.matmul(q[:, :, target], k[:, :, source].transpose(-1, -2)) / math.sqrt(width)
scores = 2.0 * torch.tanh(raw + relative)
base = 1.0 / (max(1, modalities - 1) * steps)
weights = base * rho[:, None, :, source] * torch.exp(scores.clamp(-2.0, 2.0))
numerator = numerator + torch.matmul(weights, v[:, :, source])
denominator = denominator + weights.sum(dim=-1, keepdim=True)
context = numerator / denominator
eta = torch.sigmoid(self.cross_eta_logit)
outputs.append(torch.tanh(self.cross_base(hidden[:, :, target]) + eta * self.cross_out(context)))
return torch.stack(outputs, dim=2)
def _low_rank_residual(self, gated: torch.Tensor) -> torch.Tensor:
# Linear CP factors use [1; z_m] and subtract their constant all-zero term.
batch, steps, modalities, width = gated.shape
one = torch.ones((batch, steps, 1), device=gated.device, dtype=gated.dtype)
products = torch.ones((batch, steps, 4), device=gated.device, dtype=gated.dtype)
constant = torch.ones(4, device=gated.device, dtype=gated.dtype)
for m in range(modalities):
factor = self.cp_factors[m](torch.cat((one, gated[:, :, m]), dim=-1))
products = products * factor
zero_input = torch.zeros((1, 1, width + 1), device=gated.device, dtype=gated.dtype)
zero_input[..., 0] = 1.0
constant = constant * self.cp_factors[m](zero_input)[0, 0]
residual = (products - constant) @ self.cp_output
return torch.sigmoid(self.low_rank_eta_logit) * self.low_rank_output(torch.tanh(residual))
def forward(
self,
xs: Sequence[torch.Tensor],
observed_mask: torch.Tensor,
*,
paths: int = 4,
joint_draws: bool | None = None,
) -> dict[str, torch.Tensor | list[torch.Tensor]]:
batch, steps, modalities = observed_mask.shape
if joint_draws is None:
joint_draws = self.use_joint_draws
if self.use_imputer:
completed, variance = self.imputer.complete(xs, observed_mask, paths, joint_draws=joint_draws)
else:
completed = [torch.where(observed_mask[:, :, m, None], x, torch.zeros_like(x)).unsqueeze(0) for m, x in enumerate(xs)]
variance = [torch.zeros_like(x) for x in xs]
paths = completed[0].shape[0]
uncertainty_parts = [v.mean(dim=-1) for v in variance]
uncertainty = torch.stack(uncertainty_parts, dim=-1)
rho, u, gap, span, dminus, dplus, quality_fields = self._reliability(observed_mask, uncertainty)
position = torch.linspace(0.0, 1.0, steps, device=observed_mask.device, dtype=xs[0].dtype)
pe = torch.stack((torch.sin(2 * math.pi * position), torch.cos(2 * math.pi * position),
torch.sin(4 * math.pi * position), torch.cos(4 * math.pi * position)), dim=-1)
encoded_paths: list[torch.Tensor] = []
reconstructed_paths: list[list[torch.Tensor]] = []
logits_paths: list[torch.Tensor] = []
beta_paths: list[torch.Tensor] = []
fusion_weight_paths: list[torch.Tensor] = []
null_weight_paths: list[torch.Tensor] = []
time_pool_weight_paths: list[torch.Tensor] = []
for path_index in range(paths):
enc = [projection(completed[m][path_index]) for m, projection in enumerate(self.projections)]
hmods, reconstruction = [], []
for m, encoder in enumerate(self.temporal):
if self.use_final_gate or self.use_source_attention or self.reliability_update:
scalar = torch.cat((observed_mask[:, :, m:m + 1].float(), quality_fields[:, :, m],
torch.log1p(u[:, :, m:m + 1]), dminus[:, :, m:m + 1],
dplus[:, :, m:m + 1], span[:, :, m:m + 1],
pe.unsqueeze(0).expand(batch, -1, -1)), dim=-1)
else:
# C1/C2 receive only the visibility mask and legal position code.
scalar = torch.zeros((batch, steps, 13), dtype=pe.dtype, device=pe.device)
scalar[:, :, 0] = observed_mask[:, :, m].float()
scalar[:, :, -4:] = pe.unsqueeze(0)
# q*, J_Q, edge flags, directional gaps, uncertainty and span are explicit.
seq = torch.cat((enc[m], scalar), dim=-1)
h = encoder(seq, rho[:, :, m], dminus[:, :, m], dplus[:, :, m], self.reliability_update)
hmods.append(h)
reconstruction.append(self.reconstruction_heads[m](h))
hidden = torch.stack(hmods, dim=2)
if self.use_source_attention:
enhanced = self._cross_source(hidden, rho)
else:
enhanced = hidden
if self.use_final_gate:
content = 2.0 * torch.tanh(self.content_score(enhanced).squeeze(-1))
weights_unnorm = rho * torch.exp(content.clamp(-2.0, 2.0))
denom = 1.0 + weights_unnorm.sum(dim=-1, keepdim=True)
alpha = weights_unnorm / denom
null_alpha = 1.0 / denom.squeeze(-1)
gated = enhanced * alpha.unsqueeze(-1)
fused = gated.sum(dim=2) + null_alpha.unsqueeze(-1) * self.null_expert
else:
if self.use_imputer:
alpha = torch.full_like(observed_mask.float(), 1.0 / modalities)
else:
alpha = observed_mask.float() / observed_mask.float().sum(dim=-1, keepdim=True).clamp_min(1.0)
null_alpha = torch.zeros((batch, steps), device=observed_mask.device, dtype=alpha.dtype)
fused = (enhanced * alpha.unsqueeze(-1)).sum(dim=2)
gated = enhanced * alpha.unsqueeze(-1)
if self.use_low_rank:
fused = fused + self._low_rank_residual(gated)
pool_logits = 2.0 * torch.tanh(self.pool_score(torch.tanh(self.pool_hidden(fused))).squeeze(-1))
pool_weight = torch.softmax(pool_logits, dim=1)
pooled = (pool_weight.unsqueeze(-1) * fused).sum(dim=1)
missing_rate = 1.0 - observed_mask.float().mean(dim=1)
mean_rho = rho.mean(dim=1) if (self.use_final_gate or self.use_source_attention or self.reliability_update) else observed_mask.float().mean(dim=1)
max_gap = gap.max(dim=1).values
max_span = span.max(dim=1).values
edge_rate = self._gap_features(observed_mask)[3].mean(dim=1).reshape(batch, -1)
stats = torch.cat((missing_rate, mean_rho, max_gap, max_span, edge_rate), dim=-1)
representation = torch.cat((pooled, stats), dim=-1)
feature = self.head(representation)
logits_paths.append(self.classifier(feature))
mean_fraction = torch.sigmoid(self.magnitude_mean(feature)).clamp(1e-4, 1.0 - 1e-4)
concentration = F.softplus(self.concentration_raw).clamp_min(1e-3)
alpha_beta = torch.stack((mean_fraction * concentration, (1.0 - mean_fraction) * concentration), dim=-1)
beta_paths.append(alpha_beta)
reconstructed_paths.append(reconstruction)
encoded_paths.append(hidden)
fusion_weight_paths.append(alpha)
null_weight_paths.append(null_alpha)
time_pool_weight_paths.append(pool_weight)
class_logits = torch.stack(logits_paths, dim=0)
beta_params = torch.stack(beta_paths, dim=0)
class_probs_by_path = torch.softmax(class_logits, dim=-1)
beta_mean = beta_params[..., 0] / beta_params.sum(dim=-1)
conditional_mean = 3.0 * (class_probs_by_path[..., 2] * beta_mean[..., 1] - class_probs_by_path[..., 0] * beta_mean[..., 0])
return {
"class_logits": class_logits,
"class_probs_by_path": class_probs_by_path,
"class_probs": class_probs_by_path.mean(dim=0),
"tempered_probs_by_path": torch.softmax(class_logits / 2.0, dim=-1),
"beta_params": beta_params,
"beta_mean": beta_mean,
"mixed_score": conditional_mean.mean(dim=0),
"reconstructions": [torch.stack([reconstructed_paths[k][m] for k in range(paths)], dim=0) for m in range(modalities)],
"reliability": rho,
"imputation_uncertainty": uncertainty,
"gap": gap,
"span": span,
"distance_before": dminus,
"distance_after": dplus,
"fusion_weights_by_path": torch.stack(fusion_weight_paths, dim=0),
"null_weights_by_path": torch.stack(null_weight_paths, dim=0),
"time_pool_weights_by_path": torch.stack(time_pool_weight_paths, dim=0),
"low_rank_scale": torch.sigmoid(self.low_rank_eta_logit),
}
+67
View File
@@ -0,0 +1,67 @@
from __future__ import annotations
import torch
from torch import nn
class AlignedFusionModel(nn.Module):
"""Early concatenation + BiGRU model for the supplied aligned sequence."""
def __init__(
self,
kind: str,
dims: tuple[int, int, int],
steps: int = 50,
hidden: int = 128,
dropout: float = 0.15,
) -> None:
super().__init__()
if kind != "concat":
raise ValueError(f"only the selected EarlyConcat model is maintained; got: {kind}")
self.kind = kind
self.hidden = hidden
self.projections = nn.ModuleList(
nn.Sequential(nn.Linear(size, hidden), nn.GELU(), nn.LayerNorm(hidden))
for size in dims
)
self.position = nn.Parameter(torch.randn(1, steps, hidden) * 0.02)
self.modality = nn.Parameter(torch.randn(1, 1, 3, hidden) * 0.02)
self.dropout = nn.Dropout(dropout)
self.fusion = nn.Sequential(
nn.Linear(hidden * 3 + 3, hidden), nn.GELU(), nn.LayerNorm(hidden), nn.Dropout(dropout)
)
self.temporal = nn.GRU(
input_size=hidden,
hidden_size=hidden // 2,
num_layers=1,
batch_first=True,
bidirectional=True,
)
self.head = nn.Sequential(nn.Linear(hidden, hidden // 2), nn.GELU(), nn.Dropout(dropout))
self.classifier = nn.Linear(hidden // 2, 3)
self.regressor = nn.Linear(hidden // 2, 1)
def forward(self, xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor], masks: torch.Tensor):
masks = masks.bool()
pos = self.position[:, :masks.shape[1]]
encoded = []
for modality, (projection, x) in enumerate(zip(self.projections, xs)):
token = projection(x)
token = token + pos + self.modality[:, :, modality, :]
token = token * masks[:, :, modality, None]
encoded.append(token)
stack = torch.stack(encoded, dim=2) # B x T x M x D
availability = masks.to(stack.dtype)
fused = self.fusion(torch.cat((stack.flatten(2), availability), dim=-1))
temporal, _ = self.temporal(self.dropout(fused))
time_weight = masks.any(dim=-1).to(temporal.dtype)
empty_time = time_weight.sum(dim=1, keepdim=True) <= 0
if empty_time.any():
time_weight[empty_time.squeeze(1), 0] = 1.0
pooled = (temporal * time_weight[..., None]).sum(dim=1)
pooled = pooled / time_weight.sum(dim=1, keepdim=True).clamp_min(1.0)
hidden = self.head(pooled)
logits = self.classifier(hidden)
intensity = 3.0 * torch.tanh(self.regressor(hidden).squeeze(-1))
return {"logits": logits, "intensity": intensity}
+224
View File
@@ -0,0 +1,224 @@
from __future__ import annotations
from typing import Any
import torch
import torch.nn.functional as F
from torch import nn
SUBSETS: dict[str, tuple[int, ...]] = {
"T": (0,),
"A": (1,),
"V": (2,),
"TA": (0, 1),
"TV": (0, 2),
"AV": (1, 2),
"TAV": (0, 1, 2),
}
EXPERT_NAMES = tuple(SUBSETS)
EXPERT_BITS = {
name: tuple(int(i in indices) for i in range(3))
for name, indices in SUBSETS.items()
}
class MixtureOfFusionExperts(nn.Module):
"""Seven-subset, hard-availability MoFE with the selected MLP router.
Each modality has a private projection. Experts only receive the private
projections belonging to their subset. The weighted result is passed
through one shared temporal backbone and one shared prediction head.
"""
def __init__(
self,
dims: tuple[int, int, int],
router: str = "mlp",
expert_names: tuple[str, ...] = EXPERT_NAMES,
availability_mode: str = "hard",
steps: int = 50,
latent_dim: int = 64,
hidden: int = 128,
dropout: float = 0.15,
) -> None:
super().__init__()
if router != "mlp":
raise ValueError(f"only the selected MLP router is maintained; got: {router}")
if availability_mode != "hard":
raise ValueError(f"only hard availability masking is maintained; got: {availability_mode}")
if tuple(expert_names) != EXPERT_NAMES:
raise ValueError("the selected MoFE uses all seven modality-subset experts")
self.dims = dims
self.router_kind = router
self.expert_names = tuple(expert_names)
self.availability_mode = availability_mode
self.steps = steps
self.latent_dim = latent_dim
self.hidden = hidden
# These projections are private to each modality and are not tied.
self.private_projections = nn.ModuleList(
nn.Sequential(nn.Linear(size, latent_dim), nn.GELU()) for size in dims
)
self.experts = nn.ModuleDict()
for name in self.expert_names:
n_modalities = len(SUBSETS[name])
self.experts[name] = nn.Sequential(
nn.Linear(n_modalities * latent_dim, hidden),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(hidden, latent_dim),
nn.LayerNorm(latent_dim),
)
router_input_dim = 9
self.router = nn.Sequential(
nn.Linear(router_input_dim, 16),
nn.GELU(),
nn.Linear(16, len(self.expert_names)),
)
# Shared early-fusion projection, BiGRU, and task heads.
self.all_missing_token = nn.Parameter(torch.zeros(1, 1, latent_dim))
self.input_projection = nn.Sequential(
nn.Linear(latent_dim + 3, hidden),
nn.GELU(),
nn.LayerNorm(hidden),
nn.Dropout(dropout),
)
self.dropout = nn.Dropout(dropout)
self.temporal = nn.GRU(
input_size=hidden,
hidden_size=hidden // 2,
num_layers=1,
batch_first=True,
bidirectional=True,
)
self.head = nn.Sequential(nn.Linear(hidden, hidden // 2), nn.GELU(), nn.Dropout(dropout))
self.classifier = nn.Linear(hidden // 2, 3)
self.regressor = nn.Linear(hidden // 2, 1)
@staticmethod
def _availability(masks: torch.Tensor, names: tuple[str, ...]) -> torch.Tensor:
masks = masks.bool()
columns = [masks[..., list(SUBSETS[name])].all(dim=-1) for name in names]
return torch.stack(columns, dim=-1)
def _router_features(
self,
private: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
masks: torch.Tensor,
) -> torch.Tensor:
observed = masks.to(dtype=private[0].dtype)
magnitude = torch.stack(
[torch.sqrt(x.square().mean(dim=-1) + 1e-8) for x in private], dim=-1
)
local_ratio = F.avg_pool1d(
observed.transpose(1, 2), kernel_size=5, stride=1, padding=2, count_include_pad=False
).transpose(1, 2)
return torch.cat((observed, torch.log1p(magnitude), local_ratio), dim=-1)
def _route(
self,
router_features: torch.Tensor,
availability: torch.Tensor,
force_expert: str | None,
) -> torch.Tensor:
scores = self.router(router_features)
scores = scores.masked_fill(~availability, -1e4)
weights = torch.softmax(scores, dim=-1) * availability.to(scores.dtype)
# In the full seven-expert model this is exactly the all-modalities-
# missing case. It also safely handles ablations with no eligible set.
has_expert = availability.any(dim=-1, keepdim=True)
weights = weights * has_expert.to(weights.dtype)
weights = weights / weights.sum(dim=-1, keepdim=True).clamp_min(1e-8)
if force_expert is not None:
if force_expert not in self.expert_names:
raise ValueError(f"expert {force_expert} is not enabled in this model")
expert_idx = self.expert_names.index(force_expert)
forced = torch.zeros_like(weights)
forced[..., expert_idx] = 1.0
# Force the requested expert where its modality subset is present;
# where it is unavailable, use the learned router over eligible
# experts instead of replacing observed information with zeros.
return torch.where(availability[..., expert_idx, None], forced, weights)
return weights
def forward(
self,
xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
masks: torch.Tensor,
force_expert: str | None = None,
) -> dict[str, Any]:
masks = masks.bool()
if masks.ndim != 3 or masks.shape[-1] != 3:
raise ValueError(f"masks must have shape B x T x 3, got {tuple(masks.shape)}")
if masks.shape[1] > self.steps:
raise ValueError(f"sequence has {masks.shape[1]} steps, model supports {self.steps}")
private_values = []
for modality, (projector, x) in enumerate(zip(self.private_projections, xs)):
projected = projector(x)
projected = projected * masks[..., modality, None].to(projected.dtype)
private_values.append(projected)
private = tuple(private_values)
router_features = self._router_features(private, masks)
availability = self._availability(masks, self.expert_names)
local_expert_outputs = []
for name in self.expert_names:
indices = SUBSETS[name]
expert_input = torch.cat([private[i] for i in indices], dim=-1)
local_expert_outputs.append(self.experts[name](expert_input))
expert_stack = torch.stack(local_expert_outputs, dim=-2)
alpha_local = self._route(router_features, availability, force_expert)
fused = (expert_stack * alpha_local[..., None]).sum(dim=-2)
has_expert = availability.any(dim=-1)
fused = torch.where(
has_expert[..., None], fused, self.all_missing_token.expand_as(fused)
)
# Restore a stable seven-column interface for saved diagnostics,
# including expert-set ablations.
alpha = masks.new_zeros((*masks.shape[:2], len(EXPERT_NAMES)), dtype=private[0].dtype)
expert_outputs = private[0].new_zeros((*masks.shape[:2], len(EXPERT_NAMES), self.latent_dim))
for local_idx, name in enumerate(self.expert_names):
global_idx = EXPERT_NAMES.index(name)
alpha[..., global_idx] = alpha_local[..., local_idx]
expert_outputs[..., global_idx, :] = expert_stack[..., local_idx, :]
fused_with_masks = torch.cat((fused, masks.to(fused.dtype)), dim=-1)
encoded = self.input_projection(fused_with_masks)
temporal, _ = self.temporal(self.dropout(encoded))
time_weight = masks.any(dim=-1).to(temporal.dtype)
empty_time = time_weight.sum(dim=1, keepdim=True) <= 0
if empty_time.any():
time_weight[empty_time.squeeze(1), 0] = 1.0
pooled = (temporal * time_weight[..., None]).sum(dim=1)
pooled = pooled / time_weight.sum(dim=1, keepdim=True).clamp_min(1.0)
hidden = self.head(pooled)
logits = self.classifier(hidden)
intensity = 3.0 * torch.tanh(self.regressor(hidden).squeeze(-1))
bits = torch.tensor(
[EXPERT_BITS[name] for name in EXPERT_NAMES],
dtype=alpha.dtype,
device=alpha.device,
)
utility = torch.einsum("bte,em->btm", alpha, bits)
return {
"logits": logits,
"intensity": intensity,
"fused": fused,
"alpha": alpha,
"utility": utility,
"availability": availability,
"expert_outputs": expert_outputs,
"fallback": ~has_expert,
"router_features": router_features,
}

Some files were not shown because too many files have changed in this diff Show More