"""Frozen true-time frame extraction and JPEG materialization.""" from __future__ import annotations from bisect import bisect_left, bisect_right from collections.abc import Sequence from dataclasses import dataclass, replace from fractions import Fraction from hashlib import sha256 from pathlib import Path from typing import Any NATIVE_FRAME_STRIDE = 5 SELECTED_FRAME_COUNT = 256 DECORD_NUM_THREADS = 1 JPEG_QUALITY = 95 DECORD_PTS_MAX_ABS_ERROR_S = 0.005 @dataclass(frozen=True) class StrideFrame: """One native frame retained by the frozen full-video stride.""" preselection_ordinal: int clip_stride_ordinal: int source_frame_ordinal: int pts: int time_base: Fraction source_time: Fraction def _fraction_json(value: Fraction) -> dict[str, int]: return {"numerator": value.numerator, "denominator": value.denominator} def _sha256_file(path: Path) -> str: return sha256(path.read_bytes()).hexdigest() def endpoint_inclusive_indices(count: int) -> list[int]: """Select 256 evenly spaced positions while retaining both endpoints.""" if count < SELECTED_FRAME_COUNT: raise ValueError(f"need at least {SELECTED_FRAME_COUNT} frames") denominator = SELECTED_FRAME_COUNT - 1 rounding = denominator // 2 return [ (index * (count - 1) + rounding) // denominator for index in range(SELECTED_FRAME_COUNT) ] def build_full_stride_ledger( source_path: Path, ) -> tuple[dict[str, int], list[StrideFrame]]: """Map every fifth decoded source frame to its true PyAV PTS.""" import av import decord reader = decord.VideoReader( str(source_path), ctx=decord.cpu(0), num_threads=DECORD_NUM_THREADS, ) declared_frame_count = len(reader) ledger: list[StrideFrame] = [] decoded_count = 0 last_source_time: Fraction | None = None container = av.open(str(source_path), mode="r") try: if not container.streams.video: raise ValueError(f"source has no video stream: {source_path}") stream = container.streams.video[0] stream.codec_context.thread_count = 1 for source_ordinal, frame in enumerate(container.decode(video=0)): decoded_count = source_ordinal + 1 if frame.pts is None or frame.time_base is None: raise ValueError(f"frame {source_ordinal} has no true PTS/time_base") time_base = Fraction( frame.time_base.numerator, frame.time_base.denominator, ) source_time = int(frame.pts) * time_base if last_source_time is not None and source_time < last_source_time: raise ValueError( "PyAV presentation PTS regressed at source ordinal " f"{source_ordinal}" ) last_source_time = source_time if source_ordinal % NATIVE_FRAME_STRIDE == 0: preselection_ordinal = source_ordinal // NATIVE_FRAME_STRIDE ledger.append( StrideFrame( preselection_ordinal=preselection_ordinal, clip_stride_ordinal=preselection_ordinal, source_frame_ordinal=source_ordinal, pts=int(frame.pts), time_base=time_base, source_time=source_time, ) ) finally: container.close() if decoded_count != declared_frame_count: raise ValueError( "Decord/PyAV frame-count mismatch prevents ordinal mapping: " f"decord={declared_frame_count}, pyav={decoded_count}" ) expected_stride_count = ( declared_frame_count + NATIVE_FRAME_STRIDE - 1 ) // NATIVE_FRAME_STRIDE if len(ledger) != expected_stride_count: raise AssertionError( "decoded stride ledger length does not match native frame count" ) return { "decord_native_frame_count": declared_frame_count, "pyav_decoded_frame_count": decoded_count, "official_full_video_preselection_count": expected_stride_count, }, ledger def clip_stride_ledger( full_ledger: Sequence[StrideFrame], clip_start_s: Fraction, clip_end_s: Fraction, *, source_times: Sequence[Fraction] | None = None, ) -> list[StrideFrame]: """Apply inclusive clip bounds and renumber retained clip positions.""" if clip_end_s < clip_start_s: raise ValueError("clip end precedes clip start") if source_times is None: start = bisect_left( full_ledger, clip_start_s, key=lambda entry: entry.source_time, ) stop = bisect_right( full_ledger, clip_end_s, key=lambda entry: entry.source_time, ) else: if len(source_times) != len(full_ledger): raise ValueError("source_times and full_ledger must have equal length") if any( source_times[index] > source_times[index + 1] for index in range(len(source_times) - 1) ): raise ValueError("source_times must be monotonically non-decreasing") if any( value != entry.source_time for value, entry in zip(source_times, full_ledger, strict=True) ): raise ValueError("source_times must correspond exactly to full_ledger") start = bisect_left(source_times, clip_start_s) stop = bisect_right(source_times, clip_end_s) clipped = [ replace(entry, clip_stride_ordinal=index) for index, entry in enumerate(full_ledger[start:stop]) ] if len(clipped) < SELECTED_FRAME_COUNT: raise ValueError( f"clip has only {len(clipped)} official stride frames; " f"need {SELECTED_FRAME_COUNT}" ) return clipped def build_stride_ledger( source_path: Path, clip_start_s: Fraction, clip_end_s: Fraction, ) -> tuple[dict[str, int], list[StrideFrame]]: """Build the full true-time ledger, then apply inclusive clip bounds.""" counts, full_ledger = build_full_stride_ledger(source_path) ledger = clip_stride_ledger(full_ledger, clip_start_s, clip_end_s) return {**counts, "clip_preselection_count": len(ledger)}, ledger def select_stride_frames(ledger: Sequence[StrideFrame]) -> list[StrideFrame]: """Apply the frozen endpoint-inclusive 256-frame selection.""" positions = endpoint_inclusive_indices(len(ledger)) selected = [ledger[position] for position in positions] if any( frame.clip_stride_ordinal != position for frame, position in zip(selected, positions, strict=True) ): raise AssertionError( "clip stride ordinal does not match selected ledger position" ) return selected def decode_and_write_selected_jpegs( source_path: Path, selected: Sequence[StrideFrame], frame_dir: Path, ) -> tuple[list[dict[str, Any]], dict[str, float]]: """Decode selected Decord ordinals and write official OpenCV JPEG bytes.""" if len(selected) != SELECTED_FRAME_COUNT: raise ValueError(f"expected {SELECTED_FRAME_COUNT} selected frames") import cv2 import decord import numpy as np reader = decord.VideoReader( str(source_path), ctx=decord.cpu(0), num_threads=DECORD_NUM_THREADS, ) source_ordinals = [frame.source_frame_ordinal for frame in selected] raw_timestamps = reader.get_frame_timestamp(source_ordinals) if hasattr(raw_timestamps, "asnumpy"): raw_timestamps = raw_timestamps.asnumpy() timestamp_array = np.asarray(raw_timestamps) expected_timestamp_shape = (SELECTED_FRAME_COUNT, 2) if timestamp_array.shape != expected_timestamp_shape: raise ValueError( f"unexpected Decord timestamp shape {timestamp_array.shape}, " f"expected {expected_timestamp_shape}" ) raw_frames = reader.get_batch(source_ordinals) if hasattr(raw_frames, "asnumpy"): raw_frames = raw_frames.asnumpy() frames = np.asarray(raw_frames) if ( frames.shape[0] != SELECTED_FRAME_COUNT or frames.ndim != 4 or frames.shape[-1] != 3 ): raise ValueError(f"unexpected Decord pixel shape: {frames.shape}") if frames.dtype != np.uint8: raise ValueError(f"unexpected Decord pixel dtype: {frames.dtype}") frame_dir.mkdir(parents=True, exist_ok=False) rows: list[dict[str, Any]] = [] mapping_errors: list[float] = [] jpeg_options = [int(cv2.IMWRITE_JPEG_QUALITY), JPEG_QUALITY] for selected_position, (entry, rgb_frame, timestamp) in enumerate( zip(selected, frames, timestamp_array, strict=True) ): decord_start_s = float(timestamp[0]) mapping_error = abs(decord_start_s - float(entry.source_time)) mapping_errors.append(mapping_error) if mapping_error > DECORD_PTS_MAX_ABS_ERROR_S: raise ValueError( "Decord/PyAV PTS mismatch at source ordinal " f"{entry.source_frame_ordinal}: {mapping_error:.9f}s > " f"{DECORD_PTS_MAX_ABS_ERROR_S:.9f}s" ) bgr_frame = cv2.cvtColor(rgb_frame, cv2.COLOR_RGB2BGR) filename = f"frame{entry.preselection_ordinal:07d}.jpg" output_path = frame_dir / filename if not cv2.imwrite(str(output_path), bgr_frame, jpeg_options): raise RuntimeError(f"OpenCV failed to write {output_path}") encoded_ok, encoded = cv2.imencode( ".jpg", bgr_frame, jpeg_options, ) if not encoded_ok or output_path.read_bytes() != encoded.tobytes(): raise RuntimeError( f"OpenCV JPEG encoding was not byte-stable for {output_path}" ) rows.append( { "selected_position": selected_position, "preselection_ordinal": entry.preselection_ordinal, "clip_stride_ordinal": entry.clip_stride_ordinal, "source_frame_ordinal": entry.source_frame_ordinal, "pts": entry.pts, "time_base_numerator": entry.time_base.numerator, "time_base_denominator": entry.time_base.denominator, "source_time": _fraction_json(entry.source_time), "decord_timestamp_start_s": decord_start_s, "decord_pyav_abs_error_s": mapping_error, "jpeg_relpath": filename, "jpeg_size_bytes": output_path.stat().st_size, "jpeg_sha256": _sha256_file(output_path), } ) return rows, { "max_abs_error_s": max(mapping_errors), "mean_abs_error_s": sum(mapping_errors) / len(mapping_errors), "threshold_s": DECORD_PTS_MAX_ABS_ERROR_S, }