File size: 3,035 Bytes
82416f7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
"""PyAV-based video reader for qwen-vl-utils.

The Space runtime ships torchvision >= 0.28, which removed
`torchvision.io.read_video` — the only decoder qwen-vl-utils 0.0.14 can call
without extra system libraries (torchcodec needs system FFmpeg shared libs;
decord has no cp310 wheels). PyAV wheels bundle their own FFmpeg, so this
reader reproduces qwen-vl-utils' official torchvision backend exactly:

  * frame-count/average-fps metadata from the container stream,
  * `smart_nframes` for the sampled-frame count (fps 2.0, min 4 / max 20,
    floor to factor 2 — identical because we reuse their own function),
  * `torch.linspace(0, total - 1, nframes).round().long()` index selection,
  * `sample_fps = nframes / total * video_fps`,
  * `video_metadata` dict with fps / frames_indices / total_num_frames.

`fetch_video` then handles smart_resize/BICUBIC/antialias exactly as before —
only the decode step is swapped, via `VIDEO_READER_BACKENDS["torchvision"]`.
"""

from __future__ import annotations

import time
from typing import Any, Dict, Tuple

import av
import torch


def read_video_pyav(ele: Dict[str, Any]) -> Tuple[torch.Tensor, Dict[str, Any], float]:
    """Decode a video with PyAV, returning (TCHW uint8 tensor, metadata, sample_fps).

    Mirrors `qwen_vl_utils.vision_process._read_video_torchvision`.
    """
    from qwen_vl_utils.vision_process import smart_nframes

    video_path = ele["video"]
    started = time.time()
    with av.open(video_path) as container:
        stream = container.streams.video[0]
        total_frames = stream.frames
        video_fps = float(stream.average_rate) if stream.average_rate else 0.0
        if not total_frames:  # some containers don't report a frame count
            total_frames = sum(1 for _ in container.decode(video=0))
            container.seek(0)

        nframes = smart_nframes(ele, total_frames=total_frames, video_fps=video_fps)
        idx = torch.linspace(0, total_frames - 1, nframes).round().long().tolist()
        sample_fps = nframes / max(total_frames, 1e-6) * video_fps

        wanted = set(idx)
        frames = []
        for position, frame in enumerate(container.decode(video=0)):
            if position > idx[-1]:
                break
            if position in wanted:
                # to_ndarray(format="rgb24") -> HWC uint8, matching read_video
                array = frame.to_ndarray(format="rgb24")
                frames.append(torch.from_numpy(array).permute(2, 0, 1))
        video = torch.stack(frames).contiguous()

    if video.shape[0] != nframes:  # truncated file: pad by repeating the last frame
        pad = torch.stack([video[-1]] * (nframes - video.shape[0]))
        video = torch.cat([video, pad])

    metadata = {
        "fps": video_fps,
        "frames_indices": idx,
        "total_num_frames": total_frames,
        "video_backend": "pyav",
    }
    print(f"pyav: decoded {video.shape[0]} frames of {total_frames} in {time.time() - started:.2f}s", flush=True)
    return video, metadata, sample_fps