Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 5 additions & 5 deletions benchmarks/bench_blocks.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@
import torch

from torchcodec.decoders import VideoDecoder
from torchcodec.decoders._blocks import ColorConverter, PacketDecoder, VideoDemuxer
from torchcodec.decoders._blocks import ColorConverter, VideoDemuxer, VideoPacketDecoder

# Kept minimal on purpose; the filename is derived from exactly these.
_DURATION_S = 10
Expand Down Expand Up @@ -113,15 +113,15 @@ def _consume(frames):

def _decode_sequential(path, device="cpu"):
demuxer = VideoDemuxer(path)
decoder = PacketDecoder(demuxer, device=device)
decoder = VideoPacketDecoder(demuxer, device=device)
converter = ColorConverter(device=device)
_consume(_convert(converter, _decode(decoder, _demux(demuxer))))


def _decode_prefetch_frames(path, device="cpu"):
# [demux + decode] on one thread || [color-convert] on another.
demuxer = VideoDemuxer(path)
decoder = PacketDecoder(demuxer, device=device)
decoder = VideoPacketDecoder(demuxer, device=device)
converter = ColorConverter(device=device)
frames = prefetch(_decode(decoder, _demux(demuxer)))
_consume(_convert(converter, frames))
Expand All @@ -130,7 +130,7 @@ def _decode_prefetch_frames(path, device="cpu"):
def _decode_prefetch_packets(path, device="cpu"):
# [demux] on one thread || [decode + color-convert] on another.
demuxer = VideoDemuxer(path)
decoder = PacketDecoder(demuxer, device=device)
decoder = VideoPacketDecoder(demuxer, device=device)
converter = ColorConverter(device=device)
packets = prefetch(_demux(demuxer))
_consume(_convert(converter, _decode(decoder, packets)))
Expand All @@ -139,7 +139,7 @@ def _decode_prefetch_packets(path, device="cpu"):
def _decode_prefetch_packets_and_frames(path, device="cpu"):
# [demux] || [decode] || [color-convert], each on its own thread.
demuxer = VideoDemuxer(path)
decoder = PacketDecoder(demuxer, device=device)
decoder = VideoPacketDecoder(demuxer, device=device)
converter = ColorConverter(device=device)
packets = prefetch(_demux(demuxer))
frames = prefetch(_decode(decoder, packets))
Expand Down
26 changes: 13 additions & 13 deletions examples/decoding/blocks.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@

.. code-block::

VideoDemuxer -> PacketDecoder -> ColorConverter
VideoDemuxer -> VideoPacketDecoder -> ColorConverter
Packet RawFrame RGB Frame

The blocks are passive: they never create threads, and they release the GIL.
Expand Down Expand Up @@ -63,14 +63,14 @@
# it can output a frame, and it buffers a few frames that ``drain()`` returns
# at the end.
#
# ``PacketDecoder`` and ``ColorConverter`` both accept ``device="cuda"``:
# ``VideoPacketDecoder`` and ``ColorConverter`` both accept ``device="cuda"``:
# decoding then runs on NVDEC and the color conversion on the GPU, and the
# frames never leave the device. Demuxing always happens on the CPU. Left
# unspecified, ``device`` is the current default device.
from torchcodec.decoders._blocks import ColorConverter, VideoDemuxer, PacketDecoder
from torchcodec.decoders._blocks import ColorConverter, VideoDemuxer, VideoPacketDecoder

demuxer = VideoDemuxer(video_path)
packet_decoder = PacketDecoder(demuxer, device=device)
packet_decoder = VideoPacketDecoder(demuxer, device=device)
color_converter = ColorConverter(device=device)

frames = []
Expand Down Expand Up @@ -134,15 +134,15 @@ def drain():
def sequential():
# demux -> decode -> color-convert, all on the calling thread.
demuxer = VideoDemuxer(video_path)
packet_decoder = PacketDecoder(demuxer, device=device)
packet_decoder = VideoPacketDecoder(demuxer, device=device)
color_converter = ColorConverter(device=device)
return color_convert(color_converter, decode(packet_decoder, demux(demuxer)))


def convert_on_own_thread():
# [demux + decode] on one thread || [color-convert] on another.
demuxer = VideoDemuxer(video_path)
packet_decoder = PacketDecoder(demuxer, device=device)
packet_decoder = VideoPacketDecoder(demuxer, device=device)
color_converter = ColorConverter(device=device)
raw_frames = prefetch(decode(packet_decoder, demux(demuxer)))
return color_convert(color_converter, raw_frames)
Expand All @@ -153,7 +153,7 @@ def demux_on_own_thread():
# natural split on CUDA: demuxing is CPU and I/O work, while decoding and
# color conversion both happen on the GPU, so they belong together.
demuxer = VideoDemuxer(video_path)
packet_decoder = PacketDecoder(demuxer, device=device)
packet_decoder = VideoPacketDecoder(demuxer, device=device)
color_converter = ColorConverter(device=device)
packets = prefetch(demux(demuxer))
return color_convert(color_converter, decode(packet_decoder, packets))
Expand All @@ -179,9 +179,9 @@ def demux_on_own_thread():
# drop them until you reach the timestamp you asked for.
#
# The seek also invalidates the frames the decoder is holding on to, so the
# ``PacketDecoder`` must be ``reset()``.
# ``VideoPacketDecoder`` must be ``reset()``.
demuxer = VideoDemuxer(video_path)
packet_decoder = PacketDecoder(demuxer, device=device)
packet_decoder = VideoPacketDecoder(demuxer, device=device)
color_converter = ColorConverter(device=device)

seconds = 2.5
Expand Down Expand Up @@ -212,7 +212,7 @@ def demux_on_own_thread():
# costs one pass over the file, and it leaves the demuxer back at the start.
demuxer = VideoDemuxer(video_path)
index = demuxer.scan()
packet_decoder = PacketDecoder(demuxer, device=device)
packet_decoder = VideoPacketDecoder(demuxer, device=device)
color_converter = ColorConverter(device=device)

print(f"{len(index)} frames at {index.average_fps} fps, "
Expand Down Expand Up @@ -281,7 +281,7 @@ def demux_on_own_thread():
# Color conversion is optional. A ``RawFrame`` can hand out the decoder's own
# planes as tensor views, with no copy and no conversion.
demuxer = VideoDemuxer(video_path)
packet_decoder = PacketDecoder(demuxer, device=device)
packet_decoder = VideoPacketDecoder(demuxer, device=device)
raw_frame = next(decode(packet_decoder, demux(demuxer)))

Y, U, V = raw_frame.planes
Expand Down Expand Up @@ -353,7 +353,7 @@ def upsample(plane):
)

hdr_demuxer = VideoDemuxer(hdr_video_path)
hdr_packet_decoder = PacketDecoder(hdr_demuxer, device=device)
hdr_packet_decoder = VideoPacketDecoder(hdr_demuxer, device=device)
hdr_raw = next(decode(hdr_packet_decoder, demux(hdr_demuxer)))

hdr_Y = hdr_raw.planes[0]
Expand Down Expand Up @@ -413,7 +413,7 @@ def start_live_stream():
# The blocks just stream it, and we stop whenever we want:
ffmpeg = start_live_stream()
demuxer = VideoDemuxer(fifo_path)
packet_decoder = PacketDecoder(demuxer, device=device)
packet_decoder = VideoPacketDecoder(demuxer, device=device)
color_converter = ColorConverter(device=device)

frames = []
Expand Down
38 changes: 38 additions & 0 deletions src/torchcodec/_core/AudioCommon.cpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,38 @@
// Copyright (c) Meta Platforms, Inc. and affiliates.
// All rights reserved.
//
// This source code is licensed under the BSD-style license found in the
// LICENSE file in the root directory of this source tree.

#include "AudioCommon.h"

namespace facebook::torchcodec {

torch::headeronly::ScalarType sample_format_dtype(
AVSampleFormat sample_format) {
switch (av_get_packed_sample_fmt(sample_format)) {
case AV_SAMPLE_FMT_U8:
return kStableUInt8;
case AV_SAMPLE_FMT_S16:
return kStableInt16;
case AV_SAMPLE_FMT_S32:
return kStableInt32;
case AV_SAMPLE_FMT_S64:
return kStableInt64;
case AV_SAMPLE_FMT_FLT:
return kStableFloat32;
case AV_SAMPLE_FMT_DBL:
return kStableFloat64;
default:
break;
}
const char* name = av_get_sample_fmt_name(sample_format);
STD_TORCH_CHECK(
false,
"Unsupported sample format '",
name == nullptr ? "unknown" : name,
"'.");
return kStableUInt8;
}

} // namespace facebook::torchcodec
23 changes: 23 additions & 0 deletions src/torchcodec/_core/AudioCommon.h
Original file line number Diff line number Diff line change
@@ -0,0 +1,23 @@
// Copyright (c) Meta Platforms, Inc. and affiliates.
// All rights reserved.
//
// This source code is licensed under the BSD-style license found in the
// LICENSE file in the root directory of this source tree.

#pragma once

#include "FFMPEGCommon.h"
#include "StableABICompat.h"

// Where FFmpeg's audio samples meet torch tensors. FFMPEGCommon deliberately
// knows nothing about tensors, and these helpers are shared by the decode, the
// conversion and the SingleStreamDecoder paths, so they live on their own.

namespace facebook::torchcodec {

// The dtype that holds `sample_format`'s samples exactly. Planar and packed
// variants of a format share a sample type, which is why this doesn't care
// which one it is given.
torch::headeronly::ScalarType sample_format_dtype(AVSampleFormat sample_format);

} // namespace facebook::torchcodec
108 changes: 102 additions & 6 deletions src/torchcodec/_core/PacketDecoder.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,10 @@

#include "PacketDecoder.h"

#include "AudioCommon.h"

#include <algorithm>
#include <cstring>

namespace facebook::torchcodec {

Expand Down Expand Up @@ -57,26 +60,48 @@ const AVCodec* find_decoder(
PacketDecoder::PacketDecoder(
const Demuxer& demuxer,
const StableDevice& device,
std::optional<int> ffmpeg_thread_count) {
std::optional<int> ffmpeg_thread_count)
: media_type_(demuxer.media_type()) {
bool is_audio = media_type_ == AVMEDIA_TYPE_AUDIO;
STD_TORCH_CHECK(
!is_audio || device.type() == kStableCPU,
"Audio can only be decoded on the CPU.");

device_interface_ = create_device_interface(device);
STD_TORCH_CHECK(
device_interface_ != nullptr,
"Failed to create device interface. This should never happen, please report.");

AVStream* stream = demuxer.active_stream();
time_base_ = stream->time_base;

is_mpeg_ps_ =
std::string_view(demuxer.format_context()->iformat->name) == "mpeg";
if (const int32_t* matrix = get_display_matrix_from_stream(stream)) {

if (is_audio) {
// Audio codecs are hardcoded to a single FFmpeg thread, see
// https://github.com/pytorch/torchcodec/issues/1253.
ffmpeg_thread_count = 1;
} else if (const int32_t* matrix = get_display_matrix_from_stream(stream)) {
display_matrix_.emplace();
std::copy(
matrix, matrix + display_matrix_->size(), display_matrix_->begin());
}

const AVCodec* av_codec = find_decoder(stream, device_interface_.get());
codec_context_ = create_and_open_codec_context(
stream, av_codec, device_interface_.get(), ffmpeg_thread_count);
device_interface_->initialize(codec_context_);

if (is_audio) {
// Nothing else to set up: unlike video, we hand out the samples in the
// codec's own format, so no conversion state is needed here. Note we
// deliberately do NOT set request_sample_fmt: what SingleStreamDecoder
// asks for (FLTP) is an optimization for its own conversion, and here it
// would hide what the codec natively produces.
return;
}

const AVPixFmtDescriptor* stream_desc =
av_pix_fmt_desc_get(codec_context_->pix_fmt);
int stream_bit_depth = stream_desc ? stream_desc->comp[0].depth : 8;
Expand Down Expand Up @@ -133,10 +158,12 @@ int PacketDecoder::receive_frame(UniqueAVFrame& av_frame) {
int status = device_interface_->receive_frame(av_frame);
if (status == AVSUCCESS) {
device_interface_->make_frame_standalone(av_frame);
// Attach a copy of the display matrix to the frame, so the ColorConverter
// can use it.
set_display_matrix_on_frame(
*av_frame, display_matrix_ ? display_matrix_->data() : nullptr);
if (media_type_ == AVMEDIA_TYPE_VIDEO) {
// Attach a copy of the display matrix to the frame, so the ColorConverter
// can use it.
set_display_matrix_on_frame(
*av_frame, display_matrix_ ? display_matrix_->data() : nullptr);
}
}
return status;
}
Expand Down Expand Up @@ -240,4 +267,73 @@ std::vector<torch::stable::Tensor> frame_planes(
return planes;
}

namespace {
// Scatters `num_channels`-interleaved samples into one contiguous row per
// channel. Templated on an integer of the right width rather than the actual
// sample type: we're only moving bytes around, so all that matters is size.
template <typename T>
void deinterleave(
const uint8_t* src,
uint8_t* dst,
int num_channels,
int num_samples) {
const T* in = reinterpret_cast<const T*>(src);
T* out = reinterpret_cast<T*>(dst);
for (int channel = 0; channel < num_channels; ++channel) {
T* row = out + static_cast<int64_t>(channel) * num_samples;
for (int sample = 0; sample < num_samples; ++sample) {
row[sample] = in[static_cast<int64_t>(sample) * num_channels + channel];
}
}
}
} // namespace

torch::stable::Tensor audio_samples(const AVFrame& av_frame) {
auto sample_format = static_cast<AVSampleFormat>(av_frame.format);
int num_channels = get_num_channels(av_frame);
int64_t num_samples = av_frame.nb_samples;

torch::stable::Tensor samples = torch::stable::empty(
{num_channels, num_samples}, sample_format_dtype(sample_format));
if (num_samples == 0) {
return samples;
}

int bytes_per_sample = av_get_bytes_per_sample(sample_format);
auto* dst = static_cast<uint8_t*>(samples.mutable_data_ptr());
int64_t bytes_per_channel = num_samples * bytes_per_sample;

if (av_sample_fmt_is_planar(sample_format)) {
for (int channel = 0; channel < num_channels; ++channel) {
// extended_data rather than data: the latter only holds
// AV_NUM_DATA_POINTERS (8) pointers, and we support more channels.
std::memcpy(
dst + channel * bytes_per_channel,
av_frame.extended_data[channel],
bytes_per_channel);
}
} else {
const uint8_t* src = av_frame.extended_data[0];
int num_samples_int = static_cast<int>(num_samples);
switch (bytes_per_sample) {
case 1:
deinterleave<uint8_t>(src, dst, num_channels, num_samples_int);
break;
case 2:
deinterleave<uint16_t>(src, dst, num_channels, num_samples_int);
break;
case 4:
deinterleave<uint32_t>(src, dst, num_channels, num_samples_int);
break;
case 8:
deinterleave<uint64_t>(src, dst, num_channels, num_samples_int);
break;
default:
STD_TORCH_CHECK(
false, "Unexpected sample width: ", bytes_per_sample, " bytes.");
}
}
return samples;
}

} // namespace facebook::torchcodec
Loading
Loading