Files
modeling_zhaocui/final/q1/visualize_alignment.py
T

230 lines
9.2 KiB
Python

"""Render a selected Q1 word-to-source alignment example."""
from __future__ import annotations
import argparse
import json
import sys
from pathlib import Path
from typing import Any
import matplotlib
matplotlib.use("Agg")
matplotlib.rcParams["font.family"] = ["FandolHei", "DejaVu Sans"]
matplotlib.rcParams["axes.unicode_minus"] = False
import matplotlib.pyplot as plt
import numpy as np
from matplotlib.colors import PowerNorm
Q1_DIR = Path(__file__).resolve().parent
from .q1_io import csr_from_sample, load_sample
def time_edges(centers: np.ndarray, intervals: np.ndarray, duration_s: float) -> np.ndarray:
"""Make non-overlapping display bins centered on each native source row."""
centers = np.asarray(centers, dtype=np.float64).reshape(-1)
intervals = np.asarray(intervals, dtype=np.float64)
if not len(centers):
return np.asarray([0.0, duration_s], dtype=np.float64)
if len(centers) == 1:
half_width = max(float(intervals[0, 1] - intervals[0, 0]) / 2, 1e-3)
return np.asarray(
[max(0.0, centers[0] - half_width), min(duration_s, centers[0] + half_width)],
dtype=np.float64,
)
if np.any(np.diff(centers) <= 0):
raise ValueError("native source timestamps must be strictly increasing")
middle = (centers[:-1] + centers[1:]) / 2
first = max(0.0, centers[0] - (middle[0] - centers[0]))
last = min(duration_s, centers[-1] + (centers[-1] - middle[-1]))
return np.concatenate(([first], middle, [last]))
def read_records() -> list[dict[str, Any]]:
path = Q1_DIR / "features_v2" / "manifest_q1.jsonl"
return [json.loads(line) for line in path.read_text(encoding="utf-8").splitlines() if line.strip()]
def choose_sample(sample_id: str) -> tuple[dict[str, Any], dict[str, Any], Any, Any]:
record = next((item for item in read_records() if item["sample_id"] == sample_id), None)
if record is None:
raise ValueError(f"sample_id not found in Q1 manifest: {sample_id}")
sample = load_sample(sample_id)
audio_map = csr_from_sample(sample, "query_word_audio_H_time")
vision_map = csr_from_sample(sample, "query_word_vision_H_time")
if not audio_map.nnz or not vision_map.nnz:
raise ValueError(f"sample has no valid audio or vision query rows: {sample_id}")
return record, sample, audio_map, vision_map
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--sample-id", default="-iRBcNs9oI8/8")
parser.add_argument("--output", type=Path, default=None)
args = parser.parse_args()
record, sample, audio_map, vision_map = choose_sample(args.sample_id)
words = [str(word) for word in sample["native_text_words"]]
if audio_map.shape[0] != len(words) or vision_map.shape[0] != len(words):
raise ValueError("query matrix word rows do not match the stored transcript")
duration_s = float(sample["_meta"]["duration_s"])
modalities = [
(
"原始音频位置",
audio_map,
sample["native_audio_intervals"],
"音频窗中心时间 (s)",
"音频窗",
),
(
"原始视频帧位置",
vision_map,
sample["native_vision_intervals"],
"视频帧中心时间 (s)",
"视频帧",
),
]
matrices = [item[1].toarray().astype(np.float32) for item in modalities]
for matrix in (audio_map, vision_map):
row_sums = np.asarray(matrix.sum(axis=1)).reshape(-1)
populated = row_sums > 0
if populated.any() and not np.allclose(row_sums[populated], 1.0, atol=2e-3):
raise ValueError("each populated word-query row should sum to one")
nonzero = np.concatenate([matrix[matrix > 0] for matrix in matrices if np.any(matrix > 0)])
color_max = max(float(nonzero.max()), 1e-6)
height = max(10.5, min(26.0, 4.0 + 0.22 * len(words)))
tick_size = max(5.5, min(8.5, 8.8 - 0.045 * len(words)))
fig = plt.figure(figsize=(15, height))
grid = fig.add_gridspec(4, 1, height_ratios=(4.0, 1.35, 4.0, 1.35), hspace=0.35)
y_edges = np.arange(len(words) + 1, dtype=np.float64)
text_intervals = np.asarray(sample["native_text_intervals"], dtype=np.float64)
text_valid = text_intervals[:, 1] > text_intervals[:, 0]
text_centers = text_intervals.mean(axis=1)
meshes = []
heat_axes = []
for plot_index, ((title, _sparse_matrix, intervals, _xlabel, row_name), matrix) in enumerate(zip(modalities, matrices)):
axis = fig.add_subplot(grid[2 * plot_index])
text_axis = fig.add_subplot(grid[2 * plot_index + 1], sharex=axis)
heat_axes.append(axis)
intervals = np.asarray(intervals, dtype=np.float64)
centers = intervals.mean(axis=1)
x_edges = time_edges(centers, intervals, duration_s)
mesh = axis.pcolormesh(
x_edges,
y_edges,
matrix,
shading="flat",
cmap="magma",
norm=PowerNorm(gamma=0.5, vmin=0.0, vmax=color_max),
rasterized=True,
)
meshes.append(mesh)
axis.set_xlim(0.0, duration_s)
axis.set_ylim(len(words), 0)
axis.set_xticks([])
axis.set_yticks([])
axis.set_xlabel("")
axis.set_ylabel("")
axis.grid(False)
axis.set_frame_on(False)
for spine in axis.spines.values():
spine.set_visible(False)
axis.text(
0.0,
1.025,
f"{title}(列为{row_name};有效源行的词内权重和为 1)",
transform=axis.transAxes,
ha="left",
va="bottom",
fontsize=11,
clip_on=False,
)
for word_index, word in enumerate(words):
axis.text(
-0.012,
word_index + 0.5,
word,
transform=axis.get_yaxis_transform(),
ha="right",
va="center",
fontsize=tick_size,
clip_on=False,
)
# Place actual transcript words along the source-time axis using the
# center of each stored CTC word interval.
text_axis.set_xlim(0.0, duration_s)
text_axis.set_ylim(0.0, 1.0)
text_axis.set_xticks([])
text_axis.set_yticks([])
text_axis.set_frame_on(False)
for spine in text_axis.spines.values():
spine.set_visible(False)
for center, word, valid in zip(text_centers, words, text_valid):
if valid:
text_axis.text(
center,
0.96,
word,
rotation=90,
ha="center",
va="top",
fontsize=tick_size,
clip_on=False,
)
text_axis.grid(False)
fig.subplots_adjust(left=0.17, right=0.88, top=0.91, bottom=0.08)
colorbar = fig.colorbar(meshes[0], ax=heat_axes, fraction=0.025, pad=0.02)
colorbar.ax.set_title("重叠\n权重", fontsize=8, pad=8)
sample_id = record["sample_id"]
fig.suptitle(
f"词到原始音频/视频位置查询矩阵|样本 {sample_id}",
fontsize=14,
y=0.97,
)
fig.text(
0.17,
0.02,
"横轴为真实源时间 (s),每个文本标注位于对应词区间中心。颜色表示每个词内归一化的物理时间交叠权重,不是学习得到的内容相似度。",
fontsize=8,
ha="left",
)
output = args.output or (Q1_DIR.parent / "output" / "q1" / "alignment_query_example.png")
output = output.resolve()
output.parent.mkdir(parents=True, exist_ok=True)
fig.savefig(output, dpi=220, facecolor="white")
plt.close(fig)
metadata = {
"sample_id": sample_id,
"selection": "manually selected example with valid audio and vision query maps",
"duration_s": duration_s,
"word_count": len(words),
"words": words,
"audio_query_shape": list(audio_map.shape),
"audio_query_nnz": int(audio_map.nnz),
"audio_query_mapped_words": int(np.count_nonzero(np.asarray(audio_map.sum(axis=1)).reshape(-1))),
"vision_query_shape": list(vision_map.shape),
"vision_query_nnz": int(vision_map.nnz),
"vision_query_mapped_words": int(np.count_nonzero(np.asarray(vision_map.sum(axis=1)).reshape(-1))),
"audio_query_feature_validity_channel": "log_energy (audio feature dimension 66)",
"vision_query_feature_validity_channel": "OpenFace AU12 intensity (vision feature dimension 10)",
"alignment_semantics": "per-word normalized overlap between stored CTC word intervals and native source intervals; no learned content similarity",
"coordinate_axes_shown": False,
"color_scale": "shared square-root intensity transform for visibility; original overlap weights retained",
"word_label_positions": "actual transcript labels placed at their CTC word-interval centers on the source-time axis",
"image": str(output),
}
metadata_path = output.with_suffix(".json")
metadata_path.write_text(json.dumps(metadata, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
print(json.dumps(metadata, ensure_ascii=False, indent=2))
if __name__ == "__main__":
main()