231 lines
9.2 KiB
Python
231 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
|
|
sys.path.insert(0, str(Q1_DIR))
|
|
from q1_io import csr_from_sample, load_sample # noqa: E402
|
|
|
|
|
|
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 / "results" / "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()
|