| """TransNetV2 shot detection wrapper (GPU-accelerated, lazy singleton).""" | |
| from typing import List | |
| _model = None # lazy-loaded singleton | |
| def _get_model(): | |
| global _model | |
| if _model is None: | |
| from transnetv2_pytorch import TransNetV2 # lazy import avoids startup cost | |
| _model = TransNetV2(device='auto') | |
| _model.eval() | |
| return _model | |
| def detect_transnet(video_path: str, threshold: float = 0.5) -> List[int]: | |
| """Return frame indices of cut points (start of each new shot after the first). | |
| Uses detect_scenes(), which returns a list of dicts with 'start_frame' and | |
| 'end_frame' keys (integer frame indices, not timestamps). | |
| If the model detects 1 scene (single-shot video) returns []. | |
| Args: | |
| video_path: Path to the video file. | |
| threshold: Shot-boundary probability threshold (default 0.5). | |
| Returns: | |
| List of integer frame indices where new shots begin (excluding frame 0). | |
| Empty list if the video contains only one shot. | |
| """ | |
| model = _get_model() | |
| # detect_scenes returns List[Dict] with keys: | |
| # shot_id, start_frame, end_frame, probability, start_time, end_time | |
| # This is the cleanest high-level path: no manual numpy conversion needed. | |
| scenes = model.detect_scenes(video_path, threshold=threshold) | |
| # Single-shot case: no cuts to report | |
| if len(scenes) <= 1: | |
| return [] | |
| # Cut boundaries = start_frame of every scene except the first | |
| return [int(scene['start_frame']) for scene in scenes[1:]] | |