Download bundle/extractor.py from gwd200/orena-procedure-runtime: direct link, hf CLI and curl.
- Browser
- Download file 10.9 kB
-
https://huggingface.co/gwd200/orena-procedure-runtime/resolve/main/bundle/extractor.py
- Command line
-
hf download hf://gwd200/orena-procedure-runtime/bundle/extractor.py
-
curl -L -o extractor.py https://huggingface.co/gwd200/orena-procedure-runtime/resolve/main/bundle/extractor.py
10.9 kB
| """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 | |
| 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, | |
| } | |