Skip to content

Commit 970c488

Browse files
committed
Let StreamingEncoder.open() accept the destination
1 parent 9497790 commit 970c488

7 files changed

Lines changed: 137 additions & 117 deletions

File tree

src/torchcodec/_core/Encoder.cpp

Lines changed: 16 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -1055,8 +1055,12 @@ MultiStreamEncoder::~MultiStreamEncoder() {
10551055
close();
10561056
}
10571057

1058-
MultiStreamEncoder::MultiStreamEncoder(std::string_view fileName) {
1058+
MultiStreamEncoder::MultiStreamEncoder() {
10591059
setFFmpegLogLevel();
1060+
}
1061+
1062+
void MultiStreamEncoder::open(std::string_view fileName) {
1063+
STD_TORCH_CHECK(!headerWritten_, "open() was already called.");
10601064

10611065
AVFormatContext* avFormatContext = nullptr;
10621066
int status = avformat_alloc_output_context2(
@@ -1078,13 +1082,17 @@ MultiStreamEncoder::MultiStreamEncoder(std::string_view fileName) {
10781082
fileName,
10791083
", make sure it's a valid path? ",
10801084
getFFMPEGErrorStringFromErrorCode(status));
1085+
1086+
openStreamsAndWriteHeader();
10811087
}
10821088

1083-
MultiStreamEncoder::MultiStreamEncoder(
1089+
void MultiStreamEncoder::open(
10841090
std::string_view formatName,
1085-
std::unique_ptr<AVIOContextHolder> avioContextHolder)
1086-
: avioContextHolder_(std::move(avioContextHolder)) {
1087-
setFFmpegLogLevel();
1091+
std::unique_ptr<AVIOContextHolder> avioContextHolder) {
1092+
STD_TORCH_CHECK(!headerWritten_, "open() was already called.");
1093+
1094+
avioContextHolder_ = std::move(avioContextHolder);
1095+
10881096
// Map mkv -> matroska when used as format name
10891097
formatName = (formatName == "mkv") ? "matroska" : formatName;
10901098
AVFormatContext* avFormatContext = nullptr;
@@ -1101,6 +1109,8 @@ MultiStreamEncoder::MultiStreamEncoder(
11011109
avFormatContext_.reset(avFormatContext);
11021110

11031111
avFormatContext_->pb = avioContextHolder_->getAVIOContext();
1112+
1113+
openStreamsAndWriteHeader();
11041114
}
11051115

11061116
void MultiStreamEncoder::addVideoStream(
@@ -1383,8 +1393,7 @@ void MultiStreamEncoder::initializeAudioStream() {
13831393
audioStream.avAudioFifo.reset(avAudioFifo);
13841394
}
13851395

1386-
void MultiStreamEncoder::open() {
1387-
STD_TORCH_CHECK(!headerWritten_, "open() was already called.");
1396+
void MultiStreamEncoder::openStreamsAndWriteHeader() {
13881397
STD_TORCH_CHECK(
13891398
videoStream_.has_value() || audioStream_.has_value(),
13901399
"Call addVideoStream() or addAudioStream() before open().");

src/torchcodec/_core/Encoder.h

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -187,10 +187,7 @@ class FORCE_PUBLIC_VISIBILITY MultiStreamEncoder {
187187
MultiStreamEncoder(MultiStreamEncoder&&) = delete;
188188
MultiStreamEncoder& operator=(MultiStreamEncoder&&) = delete;
189189

190-
MultiStreamEncoder(std::string_view fileName);
191-
MultiStreamEncoder(
192-
std::string_view formatName,
193-
std::unique_ptr<AVIOContextHolder> avioContextHolder);
190+
MultiStreamEncoder();
194191

195192
void addVideoStream(
196193
int height,
@@ -207,7 +204,10 @@ class FORCE_PUBLIC_VISIBILITY MultiStreamEncoder {
207204
int sampleRate,
208205
int numChannels,
209206
std::optional<int> bitRate = std::nullopt);
210-
void open();
207+
void open(std::string_view fileName);
208+
void open(
209+
std::string_view formatName,
210+
std::unique_ptr<AVIOContextHolder> avioContextHolder);
211211
void addFrames(const torch::stable::Tensor& frames);
212212
void addSamples(const torch::stable::Tensor& samples);
213213
void close();
@@ -237,6 +237,7 @@ class FORCE_PUBLIC_VISIBILITY MultiStreamEncoder {
237237
};
238238

239239
void initializeVideoStream();
240+
void openStreamsAndWriteHeader();
240241
void encodeVideoFrame(
241242
AutoAVPacket& autoAVPacket,
242243
const UniqueAVFrame& avFrame);

src/torchcodec/_core/__init__.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -18,8 +18,7 @@
1818
_get_nvdec_cache_size,
1919
_test_frame_pts_equality,
2020
core_library_path,
21-
create_streaming_encoder_to_file,
22-
create_streaming_encoder_to_file_like,
21+
create_streaming_encoder,
2322
create_wav_decoder_from_file,
2423
encode_audio_to_file,
2524
encode_audio_to_file_like,
@@ -47,5 +46,6 @@
4746
streaming_encoder_add_samples,
4847
streaming_encoder_add_video_stream,
4948
streaming_encoder_close,
50-
streaming_encoder_open,
49+
streaming_encoder_open_file,
50+
streaming_encoder_open_file_like,
5151
)

src/torchcodec/_core/custom_ops.cpp

Lines changed: 33 additions & 34 deletions
Original file line numberDiff line numberDiff line change
@@ -89,15 +89,15 @@ STABLE_TORCH_LIBRARY(torchcodec_ns, m) {
8989
m.def(
9090
"_test_frame_pts_equality(Tensor(a!) decoder, *, int frame_index, float pts_seconds_to_test) -> bool");
9191
m.def("scan_all_streams_to_update_metadata(Tensor(a!) decoder) -> ()");
92-
m.def("create_streaming_encoder_to_file(str filename) -> Tensor");
93-
m.def(
94-
"create_streaming_encoder_to_file_like(str format, int file_like_context) -> Tensor");
92+
m.def("create_streaming_encoder() -> Tensor");
9593
m.def("streaming_encoder_close(Tensor(a!) encoder) -> ()");
9694
m.def(
9795
"streaming_encoder_add_video_stream(Tensor(a!) encoder, int height, int width, float frame_rate, str device=\"cpu\", str? codec=None, str? pixel_format=None, float? crf=None, str? preset=None, str[]? extra_options=None) -> ()");
9896
m.def(
9997
"streaming_encoder_add_audio_stream(Tensor(a!) encoder, int sample_rate, int num_channels, int? bit_rate=None) -> ()");
100-
m.def("streaming_encoder_open(Tensor(a!) encoder) -> ()");
98+
m.def("streaming_encoder_open_file(Tensor(a!) encoder, str filename) -> ()");
99+
m.def(
100+
"streaming_encoder_open_file_like(Tensor(a!) encoder, str format, int file_like_context) -> ()");
101101
m.def(
102102
"streaming_encoder_add_frames(Tensor(a!) encoder, Tensor frames) -> ()");
103103
m.def(
@@ -1202,22 +1202,8 @@ void scan_all_streams_to_update_metadata(torch::stable::Tensor& decoder) {
12021202
videoDecoder->scanFileAndUpdateMetadataAndIndex();
12031203
}
12041204

1205-
torch::stable::Tensor create_streaming_encoder_to_file(std::string file_name) {
1206-
auto encoder = std::make_unique<MultiStreamEncoder>(file_name);
1207-
return wrapMultiStreamEncoderPointerToTensor(std::move(encoder));
1208-
}
1209-
1210-
torch::stable::Tensor create_streaming_encoder_to_file_like(
1211-
std::string format,
1212-
int64_t file_like_context) {
1213-
auto fileLikeContext =
1214-
reinterpret_cast<AVIOFileLikeContext*>(file_like_context);
1215-
STD_TORCH_CHECK(
1216-
fileLikeContext != nullptr, "file_like_context must be a valid pointer");
1217-
std::unique_ptr<AVIOFileLikeContext> avioContextHolder(fileLikeContext);
1218-
1219-
auto encoder = std::make_unique<MultiStreamEncoder>(
1220-
format, std::move(avioContextHolder));
1205+
torch::stable::Tensor create_streaming_encoder() {
1206+
auto encoder = std::make_unique<MultiStreamEncoder>();
12211207
return wrapMultiStreamEncoderPointerToTensor(std::move(encoder));
12221208
}
12231209

@@ -1252,8 +1238,23 @@ void streaming_encoder_add_video_stream(
12521238
std::move(extraOptionsMap));
12531239
}
12541240

1255-
void streaming_encoder_open(torch::stable::Tensor& encoder) {
1256-
unwrapTensorToGetMultiStreamEncoder(encoder)->open();
1241+
void streaming_encoder_open_file(
1242+
torch::stable::Tensor& encoder,
1243+
std::string filename) {
1244+
unwrapTensorToGetMultiStreamEncoder(encoder)->open(filename);
1245+
}
1246+
1247+
void streaming_encoder_open_file_like(
1248+
torch::stable::Tensor& encoder,
1249+
std::string format,
1250+
int64_t file_like_context) {
1251+
auto fileLikeContext =
1252+
reinterpret_cast<AVIOFileLikeContext*>(file_like_context);
1253+
STD_TORCH_CHECK(
1254+
fileLikeContext != nullptr, "file_like_context must be a valid pointer");
1255+
std::unique_ptr<AVIOFileLikeContext> avioContextHolder(fileLikeContext);
1256+
unwrapTensorToGetMultiStreamEncoder(encoder)->open(
1257+
format, std::move(avioContextHolder));
12571258
}
12581259

12591260
void streaming_encoder_add_audio_stream(
@@ -1357,19 +1358,18 @@ STABLE_TORCH_LIBRARY_IMPL(torchcodec_ns, BackendSelect, m) {
13571358
m.impl("encode_video_to_file", TORCH_BOX(&encode_video_to_file));
13581359
m.impl("encode_video_to_tensor", TORCH_BOX(&encode_video_to_tensor));
13591360
m.impl("_encode_video_to_file_like", TORCH_BOX(&_encode_video_to_file_like));
1360-
m.impl(
1361-
"create_streaming_encoder_to_file",
1362-
TORCH_BOX(&create_streaming_encoder_to_file));
1363-
m.impl(
1364-
"create_streaming_encoder_to_file_like",
1365-
TORCH_BOX(&create_streaming_encoder_to_file_like));
1361+
m.impl("create_streaming_encoder", TORCH_BOX(&create_streaming_encoder));
13661362
m.impl(
13671363
"streaming_encoder_add_video_stream",
13681364
TORCH_BOX(&streaming_encoder_add_video_stream));
13691365
m.impl(
13701366
"streaming_encoder_add_audio_stream",
13711367
TORCH_BOX(&streaming_encoder_add_audio_stream));
1372-
m.impl("streaming_encoder_open", TORCH_BOX(&streaming_encoder_open));
1368+
m.impl(
1369+
"streaming_encoder_open_file", TORCH_BOX(&streaming_encoder_open_file));
1370+
m.impl(
1371+
"streaming_encoder_open_file_like",
1372+
TORCH_BOX(&streaming_encoder_open_file_like));
13731373
m.impl(
13741374
"streaming_encoder_add_frames", TORCH_BOX(&streaming_encoder_add_frames));
13751375
m.impl(
@@ -1420,20 +1420,19 @@ STABLE_TORCH_LIBRARY_IMPL(torchcodec_ns, CPU, m) {
14201420
TORCH_BOX(&scan_all_streams_to_update_metadata));
14211421

14221422
m.impl("_get_backend_details", TORCH_BOX(&get_backend_details));
1423+
m.impl("create_streaming_encoder", TORCH_BOX(&create_streaming_encoder));
14231424
m.impl(
1424-
"create_streaming_encoder_to_file",
1425-
TORCH_BOX(&create_streaming_encoder_to_file));
1425+
"streaming_encoder_open_file", TORCH_BOX(&streaming_encoder_open_file));
14261426
m.impl(
1427-
"create_streaming_encoder_to_file_like",
1428-
TORCH_BOX(&create_streaming_encoder_to_file_like));
1427+
"streaming_encoder_open_file_like",
1428+
TORCH_BOX(&streaming_encoder_open_file_like));
14291429
m.impl("streaming_encoder_close", TORCH_BOX(&streaming_encoder_close));
14301430
m.impl(
14311431
"streaming_encoder_add_video_stream",
14321432
TORCH_BOX(&streaming_encoder_add_video_stream));
14331433
m.impl(
14341434
"streaming_encoder_add_audio_stream",
14351435
TORCH_BOX(&streaming_encoder_add_audio_stream));
1436-
m.impl("streaming_encoder_open", TORCH_BOX(&streaming_encoder_open));
14371436
m.impl(
14381437
"streaming_encoder_add_frames", TORCH_BOX(&streaming_encoder_add_frames));
14391438
m.impl(

src/torchcodec/_core/ops.py

Lines changed: 25 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -138,11 +138,8 @@ def add_video_stream(
138138
torch.ops.torchcodec_ns._get_json_ffmpeg_library_versions.default
139139
)
140140
_get_backend_details = torch.ops.torchcodec_ns._get_backend_details.default
141-
create_streaming_encoder_to_file = torch._dynamo.disallow_in_graph(
142-
torch.ops.torchcodec_ns.create_streaming_encoder_to_file.default
143-
)
144-
_create_streaming_encoder_to_file_like = torch._dynamo.disallow_in_graph(
145-
torch.ops.torchcodec_ns.create_streaming_encoder_to_file_like.default
141+
create_streaming_encoder = torch._dynamo.disallow_in_graph(
142+
torch.ops.torchcodec_ns.create_streaming_encoder.default
146143
)
147144
streaming_encoder_close = torch.ops.torchcodec_ns.streaming_encoder_close.default
148145
streaming_encoder_add_video_stream = (
@@ -151,7 +148,12 @@ def add_video_stream(
151148
streaming_encoder_add_audio_stream = (
152149
torch.ops.torchcodec_ns.streaming_encoder_add_audio_stream.default
153150
)
154-
streaming_encoder_open = torch.ops.torchcodec_ns.streaming_encoder_open.default
151+
streaming_encoder_open_file = (
152+
torch.ops.torchcodec_ns.streaming_encoder_open_file.default
153+
)
154+
_streaming_encoder_open_file_like = (
155+
torch.ops.torchcodec_ns.streaming_encoder_open_file_like.default
156+
)
155157
streaming_encoder_add_frames = (
156158
torch.ops.torchcodec_ns.streaming_encoder_add_frames.default
157159
)
@@ -271,14 +273,16 @@ def encode_video_to_file_like(
271273
)
272274

273275

274-
def create_streaming_encoder_to_file_like(
276+
def streaming_encoder_open_file_like(
277+
encoder: torch.Tensor,
275278
format: str,
276279
file_like: io.RawIOBase | io.BufferedIOBase,
277-
) -> torch.Tensor:
280+
) -> None:
278281
assert _pybind_ops is not None
279-
return _create_streaming_encoder_to_file_like(
282+
_streaming_encoder_open_file_like(
283+
encoder,
280284
format,
281-
_pybind_ops.create_file_like_context(file_like, True), # True means for writing
285+
_pybind_ops.create_file_like_context(file_like, True),
282286
)
283287

284288

@@ -621,18 +625,8 @@ def _get_backend_details_abstract(decoder: torch.Tensor) -> str:
621625
return ""
622626

623627

624-
@register_fake("torchcodec_ns::create_streaming_encoder_to_file")
625-
def _create_streaming_encoder_to_file_abstract(
626-
filename: str,
627-
) -> torch.Tensor:
628-
return torch.empty([], dtype=torch.long)
629-
630-
631-
@register_fake("torchcodec_ns::create_streaming_encoder_to_file_like")
632-
def _create_streaming_encoder_to_file_like_abstract(
633-
format: str,
634-
file_like_context: int,
635-
) -> torch.Tensor:
628+
@register_fake("torchcodec_ns::create_streaming_encoder")
629+
def _create_streaming_encoder_abstract() -> torch.Tensor:
636630
return torch.empty([], dtype=torch.long)
637631

638632

@@ -667,8 +661,15 @@ def streaming_encoder_add_audio_stream_abstract(
667661
return
668662

669663

670-
@register_fake("torchcodec_ns::streaming_encoder_open")
671-
def streaming_encoder_open_abstract(encoder: torch.Tensor) -> None:
664+
@register_fake("torchcodec_ns::streaming_encoder_open_file")
665+
def streaming_encoder_open_file_abstract(encoder: torch.Tensor, filename: str) -> None:
666+
return
667+
668+
669+
@register_fake("torchcodec_ns::streaming_encoder_open_file_like")
670+
def streaming_encoder_open_file_like_abstract(
671+
encoder: torch.Tensor, format: str, file_like_context: int
672+
) -> None:
672673
return
673674

674675

src/torchcodec/encoders/_multi_stream_encoder.py

Lines changed: 7 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -22,13 +22,8 @@ def write(self, samples: Tensor) -> None:
2222

2323

2424
class StreamingEncoder:
25-
def __init__(self, dest, *, format: str | None = None):
26-
if format is not None:
27-
self._encoder_tensor = _core.create_streaming_encoder_to_file_like(
28-
format, dest
29-
)
30-
else:
31-
self._encoder_tensor = _core.create_streaming_encoder_to_file(str(dest))
25+
def __init__(self):
26+
self._encoder_tensor = _core.create_streaming_encoder()
3227

3328
def add_video(
3429
self,
@@ -75,8 +70,11 @@ def add_audio(
7570
)
7671
return _AudioStream(self._encoder_tensor)
7772

78-
def open(self) -> None:
79-
_core.streaming_encoder_open(self._encoder_tensor)
73+
def open(self, dest, *, format: str | None = None) -> None:
74+
if format is not None:
75+
_core.streaming_encoder_open_file_like(self._encoder_tensor, format, dest)
76+
else:
77+
_core.streaming_encoder_open_file(self._encoder_tensor, str(dest))
8078

8179
def close(self) -> None:
8280
_core.streaming_encoder_close(self._encoder_tensor)

0 commit comments

Comments
 (0)