Skip to content

Commit 78ed5ad

Browse files
authored
Decode audio in the Blocks API, as RawAudioSamples (#1666)
1 parent 7a07432 commit 78ed5ad

16 files changed

Lines changed: 602 additions & 92 deletions

benchmarks/bench_blocks.py

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -15,7 +15,7 @@
1515
import torch
1616

1717
from torchcodec.decoders import VideoDecoder
18-
from torchcodec.decoders._blocks import ColorConverter, PacketDecoder, VideoDemuxer
18+
from torchcodec.decoders._blocks import ColorConverter, VideoDemuxer, VideoPacketDecoder
1919

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

114114
def _decode_sequential(path, device="cpu"):
115115
demuxer = VideoDemuxer(path)
116-
decoder = PacketDecoder(demuxer, device=device)
116+
decoder = VideoPacketDecoder(demuxer, device=device)
117117
converter = ColorConverter(device=device)
118118
_consume(_convert(converter, _decode(decoder, _demux(demuxer))))
119119

120120

121121
def _decode_prefetch_frames(path, device="cpu"):
122122
# [demux + decode] on one thread || [color-convert] on another.
123123
demuxer = VideoDemuxer(path)
124-
decoder = PacketDecoder(demuxer, device=device)
124+
decoder = VideoPacketDecoder(demuxer, device=device)
125125
converter = ColorConverter(device=device)
126126
frames = prefetch(_decode(decoder, _demux(demuxer)))
127127
_consume(_convert(converter, frames))
@@ -130,7 +130,7 @@ def _decode_prefetch_frames(path, device="cpu"):
130130
def _decode_prefetch_packets(path, device="cpu"):
131131
# [demux] on one thread || [decode + color-convert] on another.
132132
demuxer = VideoDemuxer(path)
133-
decoder = PacketDecoder(demuxer, device=device)
133+
decoder = VideoPacketDecoder(demuxer, device=device)
134134
converter = ColorConverter(device=device)
135135
packets = prefetch(_demux(demuxer))
136136
_consume(_convert(converter, _decode(decoder, packets)))
@@ -139,7 +139,7 @@ def _decode_prefetch_packets(path, device="cpu"):
139139
def _decode_prefetch_packets_and_frames(path, device="cpu"):
140140
# [demux] || [decode] || [color-convert], each on its own thread.
141141
demuxer = VideoDemuxer(path)
142-
decoder = PacketDecoder(demuxer, device=device)
142+
decoder = VideoPacketDecoder(demuxer, device=device)
143143
converter = ColorConverter(device=device)
144144
packets = prefetch(_demux(demuxer))
145145
frames = prefetch(_decode(decoder, packets))

examples/decoding/blocks.py

Lines changed: 13 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -21,7 +21,7 @@
2121
2222
.. code-block::
2323
24-
VideoDemuxer -> PacketDecoder -> ColorConverter
24+
VideoDemuxer -> VideoPacketDecoder -> ColorConverter
2525
Packet RawFrame RGB Frame
2626
2727
The blocks are passive: they never create threads, and they release the GIL.
@@ -63,14 +63,14 @@
6363
# it can output a frame, and it buffers a few frames that ``drain()`` returns
6464
# at the end.
6565
#
66-
# ``PacketDecoder`` and ``ColorConverter`` both accept ``device="cuda"``:
66+
# ``VideoPacketDecoder`` and ``ColorConverter`` both accept ``device="cuda"``:
6767
# decoding then runs on NVDEC and the color conversion on the GPU, and the
6868
# frames never leave the device. Demuxing always happens on the CPU. Left
6969
# unspecified, ``device`` is the current default device.
70-
from torchcodec.decoders._blocks import ColorConverter, VideoDemuxer, PacketDecoder
70+
from torchcodec.decoders._blocks import ColorConverter, VideoDemuxer, VideoPacketDecoder
7171

7272
demuxer = VideoDemuxer(video_path)
73-
packet_decoder = PacketDecoder(demuxer, device=device)
73+
packet_decoder = VideoPacketDecoder(demuxer, device=device)
7474
color_converter = ColorConverter(device=device)
7575

7676
frames = []
@@ -134,15 +134,15 @@ def drain():
134134
def sequential():
135135
# demux -> decode -> color-convert, all on the calling thread.
136136
demuxer = VideoDemuxer(video_path)
137-
packet_decoder = PacketDecoder(demuxer, device=device)
137+
packet_decoder = VideoPacketDecoder(demuxer, device=device)
138138
color_converter = ColorConverter(device=device)
139139
return color_convert(color_converter, decode(packet_decoder, demux(demuxer)))
140140

141141

142142
def convert_on_own_thread():
143143
# [demux + decode] on one thread || [color-convert] on another.
144144
demuxer = VideoDemuxer(video_path)
145-
packet_decoder = PacketDecoder(demuxer, device=device)
145+
packet_decoder = VideoPacketDecoder(demuxer, device=device)
146146
color_converter = ColorConverter(device=device)
147147
raw_frames = prefetch(decode(packet_decoder, demux(demuxer)))
148148
return color_convert(color_converter, raw_frames)
@@ -153,7 +153,7 @@ def demux_on_own_thread():
153153
# natural split on CUDA: demuxing is CPU and I/O work, while decoding and
154154
# color conversion both happen on the GPU, so they belong together.
155155
demuxer = VideoDemuxer(video_path)
156-
packet_decoder = PacketDecoder(demuxer, device=device)
156+
packet_decoder = VideoPacketDecoder(demuxer, device=device)
157157
color_converter = ColorConverter(device=device)
158158
packets = prefetch(demux(demuxer))
159159
return color_convert(color_converter, decode(packet_decoder, packets))
@@ -179,9 +179,9 @@ def demux_on_own_thread():
179179
# drop them until you reach the timestamp you asked for.
180180
#
181181
# The seek also invalidates the frames the decoder is holding on to, so the
182-
# ``PacketDecoder`` must be ``reset()``.
182+
# ``VideoPacketDecoder`` must be ``reset()``.
183183
demuxer = VideoDemuxer(video_path)
184-
packet_decoder = PacketDecoder(demuxer, device=device)
184+
packet_decoder = VideoPacketDecoder(demuxer, device=device)
185185
color_converter = ColorConverter(device=device)
186186

187187
seconds = 2.5
@@ -212,7 +212,7 @@ def demux_on_own_thread():
212212
# costs one pass over the file, and it leaves the demuxer back at the start.
213213
demuxer = VideoDemuxer(video_path)
214214
index = demuxer.scan()
215-
packet_decoder = PacketDecoder(demuxer, device=device)
215+
packet_decoder = VideoPacketDecoder(demuxer, device=device)
216216
color_converter = ColorConverter(device=device)
217217

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

287287
Y, U, V = raw_frame.planes
@@ -353,7 +353,7 @@ def upsample(plane):
353353
)
354354

355355
hdr_demuxer = VideoDemuxer(hdr_video_path)
356-
hdr_packet_decoder = PacketDecoder(hdr_demuxer, device=device)
356+
hdr_packet_decoder = VideoPacketDecoder(hdr_demuxer, device=device)
357357
hdr_raw = next(decode(hdr_packet_decoder, demux(hdr_demuxer)))
358358

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

419419
frames = []
Lines changed: 38 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,38 @@
1+
// Copyright (c) Meta Platforms, Inc. and affiliates.
2+
// All rights reserved.
3+
//
4+
// This source code is licensed under the BSD-style license found in the
5+
// LICENSE file in the root directory of this source tree.
6+
7+
#include "AudioCommon.h"
8+
9+
namespace facebook::torchcodec {
10+
11+
torch::headeronly::ScalarType sample_format_dtype(
12+
AVSampleFormat sample_format) {
13+
switch (av_get_packed_sample_fmt(sample_format)) {
14+
case AV_SAMPLE_FMT_U8:
15+
return kStableUInt8;
16+
case AV_SAMPLE_FMT_S16:
17+
return kStableInt16;
18+
case AV_SAMPLE_FMT_S32:
19+
return kStableInt32;
20+
case AV_SAMPLE_FMT_S64:
21+
return kStableInt64;
22+
case AV_SAMPLE_FMT_FLT:
23+
return kStableFloat32;
24+
case AV_SAMPLE_FMT_DBL:
25+
return kStableFloat64;
26+
default:
27+
break;
28+
}
29+
const char* name = av_get_sample_fmt_name(sample_format);
30+
STD_TORCH_CHECK(
31+
false,
32+
"Unsupported sample format '",
33+
name == nullptr ? "unknown" : name,
34+
"'.");
35+
return kStableUInt8;
36+
}
37+
38+
} // namespace facebook::torchcodec

src/torchcodec/_core/AudioCommon.h

Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,23 @@
1+
// Copyright (c) Meta Platforms, Inc. and affiliates.
2+
// All rights reserved.
3+
//
4+
// This source code is licensed under the BSD-style license found in the
5+
// LICENSE file in the root directory of this source tree.
6+
7+
#pragma once
8+
9+
#include "FFMPEGCommon.h"
10+
#include "StableABICompat.h"
11+
12+
// Where FFmpeg's audio samples meet torch tensors. FFMPEGCommon deliberately
13+
// knows nothing about tensors, and these helpers are shared by the decode, the
14+
// conversion and the SingleStreamDecoder paths, so they live on their own.
15+
16+
namespace facebook::torchcodec {
17+
18+
// The dtype that holds `sample_format`'s samples exactly. Planar and packed
19+
// variants of a format share a sample type, which is why this doesn't care
20+
// which one it is given.
21+
torch::headeronly::ScalarType sample_format_dtype(AVSampleFormat sample_format);
22+
23+
} // namespace facebook::torchcodec

src/torchcodec/_core/PacketDecoder.cpp

Lines changed: 102 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,10 @@
66

77
#include "PacketDecoder.h"
88

9+
#include "AudioCommon.h"
10+
911
#include <algorithm>
12+
#include <cstring>
1013

1114
namespace facebook::torchcodec {
1215

@@ -57,26 +60,48 @@ const AVCodec* find_decoder(
5760
PacketDecoder::PacketDecoder(
5861
const Demuxer& demuxer,
5962
const StableDevice& device,
60-
std::optional<int> ffmpeg_thread_count) {
63+
std::optional<int> ffmpeg_thread_count)
64+
: media_type_(demuxer.media_type()) {
65+
bool is_audio = media_type_ == AVMEDIA_TYPE_AUDIO;
66+
STD_TORCH_CHECK(
67+
!is_audio || device.type() == kStableCPU,
68+
"Audio can only be decoded on the CPU.");
69+
6170
device_interface_ = create_device_interface(device);
6271
STD_TORCH_CHECK(
6372
device_interface_ != nullptr,
6473
"Failed to create device interface. This should never happen, please report.");
6574

6675
AVStream* stream = demuxer.active_stream();
6776
time_base_ = stream->time_base;
77+
6878
is_mpeg_ps_ =
6979
std::string_view(demuxer.format_context()->iformat->name) == "mpeg";
70-
if (const int32_t* matrix = get_display_matrix_from_stream(stream)) {
80+
81+
if (is_audio) {
82+
// Audio codecs are hardcoded to a single FFmpeg thread, see
83+
// https://github.com/pytorch/torchcodec/issues/1253.
84+
ffmpeg_thread_count = 1;
85+
} else if (const int32_t* matrix = get_display_matrix_from_stream(stream)) {
7186
display_matrix_.emplace();
7287
std::copy(
7388
matrix, matrix + display_matrix_->size(), display_matrix_->begin());
7489
}
90+
7591
const AVCodec* av_codec = find_decoder(stream, device_interface_.get());
7692
codec_context_ = create_and_open_codec_context(
7793
stream, av_codec, device_interface_.get(), ffmpeg_thread_count);
7894
device_interface_->initialize(codec_context_);
7995

96+
if (is_audio) {
97+
// Nothing else to set up: unlike video, we hand out the samples in the
98+
// codec's own format, so no conversion state is needed here. Note we
99+
// deliberately do NOT set request_sample_fmt: what SingleStreamDecoder
100+
// asks for (FLTP) is an optimization for its own conversion, and here it
101+
// would hide what the codec natively produces.
102+
return;
103+
}
104+
80105
const AVPixFmtDescriptor* stream_desc =
81106
av_pix_fmt_desc_get(codec_context_->pix_fmt);
82107
int stream_bit_depth = stream_desc ? stream_desc->comp[0].depth : 8;
@@ -133,10 +158,12 @@ int PacketDecoder::receive_frame(UniqueAVFrame& av_frame) {
133158
int status = device_interface_->receive_frame(av_frame);
134159
if (status == AVSUCCESS) {
135160
device_interface_->make_frame_standalone(av_frame);
136-
// Attach a copy of the display matrix to the frame, so the ColorConverter
137-
// can use it.
138-
set_display_matrix_on_frame(
139-
*av_frame, display_matrix_ ? display_matrix_->data() : nullptr);
161+
if (media_type_ == AVMEDIA_TYPE_VIDEO) {
162+
// Attach a copy of the display matrix to the frame, so the ColorConverter
163+
// can use it.
164+
set_display_matrix_on_frame(
165+
*av_frame, display_matrix_ ? display_matrix_->data() : nullptr);
166+
}
140167
}
141168
return status;
142169
}
@@ -240,4 +267,73 @@ std::vector<torch::stable::Tensor> frame_planes(
240267
return planes;
241268
}
242269

270+
namespace {
271+
// Scatters `num_channels`-interleaved samples into one contiguous row per
272+
// channel. Templated on an integer of the right width rather than the actual
273+
// sample type: we're only moving bytes around, so all that matters is size.
274+
template <typename T>
275+
void deinterleave(
276+
const uint8_t* src,
277+
uint8_t* dst,
278+
int num_channels,
279+
int num_samples) {
280+
const T* in = reinterpret_cast<const T*>(src);
281+
T* out = reinterpret_cast<T*>(dst);
282+
for (int channel = 0; channel < num_channels; ++channel) {
283+
T* row = out + static_cast<int64_t>(channel) * num_samples;
284+
for (int sample = 0; sample < num_samples; ++sample) {
285+
row[sample] = in[static_cast<int64_t>(sample) * num_channels + channel];
286+
}
287+
}
288+
}
289+
} // namespace
290+
291+
torch::stable::Tensor audio_samples(const AVFrame& av_frame) {
292+
auto sample_format = static_cast<AVSampleFormat>(av_frame.format);
293+
int num_channels = get_num_channels(av_frame);
294+
int64_t num_samples = av_frame.nb_samples;
295+
296+
torch::stable::Tensor samples = torch::stable::empty(
297+
{num_channels, num_samples}, sample_format_dtype(sample_format));
298+
if (num_samples == 0) {
299+
return samples;
300+
}
301+
302+
int bytes_per_sample = av_get_bytes_per_sample(sample_format);
303+
auto* dst = static_cast<uint8_t*>(samples.mutable_data_ptr());
304+
int64_t bytes_per_channel = num_samples * bytes_per_sample;
305+
306+
if (av_sample_fmt_is_planar(sample_format)) {
307+
for (int channel = 0; channel < num_channels; ++channel) {
308+
// extended_data rather than data: the latter only holds
309+
// AV_NUM_DATA_POINTERS (8) pointers, and we support more channels.
310+
std::memcpy(
311+
dst + channel * bytes_per_channel,
312+
av_frame.extended_data[channel],
313+
bytes_per_channel);
314+
}
315+
} else {
316+
const uint8_t* src = av_frame.extended_data[0];
317+
int num_samples_int = static_cast<int>(num_samples);
318+
switch (bytes_per_sample) {
319+
case 1:
320+
deinterleave<uint8_t>(src, dst, num_channels, num_samples_int);
321+
break;
322+
case 2:
323+
deinterleave<uint16_t>(src, dst, num_channels, num_samples_int);
324+
break;
325+
case 4:
326+
deinterleave<uint32_t>(src, dst, num_channels, num_samples_int);
327+
break;
328+
case 8:
329+
deinterleave<uint64_t>(src, dst, num_channels, num_samples_int);
330+
break;
331+
default:
332+
STD_TORCH_CHECK(
333+
false, "Unexpected sample width: ", bytes_per_sample, " bytes.");
334+
}
335+
}
336+
return samples;
337+
}
338+
243339
} // namespace facebook::torchcodec

0 commit comments

Comments
 (0)