"""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()