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
23 changes: 16 additions & 7 deletions src/torchcodec/_core/Encoder.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1055,8 +1055,12 @@ MultiStreamEncoder::~MultiStreamEncoder() {
close();
}

MultiStreamEncoder::MultiStreamEncoder(std::string_view fileName) {
MultiStreamEncoder::MultiStreamEncoder() {
setFFmpegLogLevel();
}

void MultiStreamEncoder::open(std::string_view fileName) {
STD_TORCH_CHECK(!headerWritten_, "open() was already called.");

AVFormatContext* avFormatContext = nullptr;
int status = avformat_alloc_output_context2(
Expand All @@ -1078,13 +1082,17 @@ MultiStreamEncoder::MultiStreamEncoder(std::string_view fileName) {
fileName,
", make sure it's a valid path? ",
getFFMPEGErrorStringFromErrorCode(status));

openStreamsAndWriteHeader();
}

MultiStreamEncoder::MultiStreamEncoder(
void MultiStreamEncoder::open(
std::string_view formatName,
std::unique_ptr<AVIOContextHolder> avioContextHolder)
: avioContextHolder_(std::move(avioContextHolder)) {
setFFmpegLogLevel();
std::unique_ptr<AVIOContextHolder> avioContextHolder) {
STD_TORCH_CHECK(!headerWritten_, "open() was already called.");

avioContextHolder_ = std::move(avioContextHolder);

// Map mkv -> matroska when used as format name
formatName = (formatName == "mkv") ? "matroska" : formatName;
AVFormatContext* avFormatContext = nullptr;
Expand All @@ -1101,6 +1109,8 @@ MultiStreamEncoder::MultiStreamEncoder(
avFormatContext_.reset(avFormatContext);

avFormatContext_->pb = avioContextHolder_->getAVIOContext();

openStreamsAndWriteHeader();
}

void MultiStreamEncoder::addVideoStream(
Expand Down Expand Up @@ -1383,8 +1393,7 @@ void MultiStreamEncoder::initializeAudioStream() {
audioStream.avAudioFifo.reset(avAudioFifo);
}

void MultiStreamEncoder::open() {
STD_TORCH_CHECK(!headerWritten_, "open() was already called.");
void MultiStreamEncoder::openStreamsAndWriteHeader() {
STD_TORCH_CHECK(
videoStream_.has_value() || audioStream_.has_value(),
"Call addVideoStream() or addAudioStream() before open().");
Expand Down
11 changes: 6 additions & 5 deletions src/torchcodec/_core/Encoder.h
Original file line number Diff line number Diff line change
Expand Up @@ -187,10 +187,7 @@ class FORCE_PUBLIC_VISIBILITY MultiStreamEncoder {
MultiStreamEncoder(MultiStreamEncoder&&) = delete;
MultiStreamEncoder& operator=(MultiStreamEncoder&&) = delete;

MultiStreamEncoder(std::string_view fileName);
MultiStreamEncoder(
std::string_view formatName,
std::unique_ptr<AVIOContextHolder> avioContextHolder);
MultiStreamEncoder();

void addVideoStream(
int height,
Expand All @@ -207,7 +204,10 @@ class FORCE_PUBLIC_VISIBILITY MultiStreamEncoder {
int sampleRate,
int numChannels,
std::optional<int> bitRate = std::nullopt);
void open();
void open(std::string_view fileName);
void open(
std::string_view formatName,
std::unique_ptr<AVIOContextHolder> avioContextHolder);
void addFrames(const torch::stable::Tensor& frames);
void addSamples(const torch::stable::Tensor& samples);
void close();
Expand Down Expand Up @@ -237,6 +237,7 @@ class FORCE_PUBLIC_VISIBILITY MultiStreamEncoder {
};

void initializeVideoStream();
void openStreamsAndWriteHeader();
void encodeVideoFrame(
AutoAVPacket& autoAVPacket,
const UniqueAVFrame& avFrame);
Expand Down
6 changes: 3 additions & 3 deletions src/torchcodec/_core/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,8 +18,7 @@
_get_nvdec_cache_size,
_test_frame_pts_equality,
core_library_path,
create_streaming_encoder_to_file,
create_streaming_encoder_to_file_like,
create_streaming_encoder,
create_wav_decoder_from_file,
encode_audio_to_file,
encode_audio_to_file_like,
Expand Down Expand Up @@ -47,5 +46,6 @@
streaming_encoder_add_samples,
streaming_encoder_add_video_stream,
streaming_encoder_close,
streaming_encoder_open,
streaming_encoder_open_file,
streaming_encoder_open_file_like,
)
67 changes: 33 additions & 34 deletions src/torchcodec/_core/custom_ops.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -89,15 +89,15 @@ STABLE_TORCH_LIBRARY(torchcodec_ns, m) {
m.def(
"_test_frame_pts_equality(Tensor(a!) decoder, *, int frame_index, float pts_seconds_to_test) -> bool");
m.def("scan_all_streams_to_update_metadata(Tensor(a!) decoder) -> ()");
m.def("create_streaming_encoder_to_file(str filename) -> Tensor");
m.def(
"create_streaming_encoder_to_file_like(str format, int file_like_context) -> Tensor");
m.def("create_streaming_encoder() -> Tensor");
m.def("streaming_encoder_close(Tensor(a!) encoder) -> ()");
m.def(
"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) -> ()");
m.def(
"streaming_encoder_add_audio_stream(Tensor(a!) encoder, int sample_rate, int num_channels, int? bit_rate=None) -> ()");
m.def("streaming_encoder_open(Tensor(a!) encoder) -> ()");
m.def("streaming_encoder_open_file(Tensor(a!) encoder, str filename) -> ()");
m.def(
"streaming_encoder_open_file_like(Tensor(a!) encoder, str format, int file_like_context) -> ()");
m.def(
"streaming_encoder_add_frames(Tensor(a!) encoder, Tensor frames) -> ()");
m.def(
Expand Down Expand Up @@ -1202,22 +1202,8 @@ void scan_all_streams_to_update_metadata(torch::stable::Tensor& decoder) {
videoDecoder->scanFileAndUpdateMetadataAndIndex();
}

torch::stable::Tensor create_streaming_encoder_to_file(std::string file_name) {
auto encoder = std::make_unique<MultiStreamEncoder>(file_name);
return wrapMultiStreamEncoderPointerToTensor(std::move(encoder));
}

torch::stable::Tensor create_streaming_encoder_to_file_like(
std::string format,
int64_t file_like_context) {
auto fileLikeContext =
reinterpret_cast<AVIOFileLikeContext*>(file_like_context);
STD_TORCH_CHECK(
fileLikeContext != nullptr, "file_like_context must be a valid pointer");
std::unique_ptr<AVIOFileLikeContext> avioContextHolder(fileLikeContext);

auto encoder = std::make_unique<MultiStreamEncoder>(
format, std::move(avioContextHolder));
torch::stable::Tensor create_streaming_encoder() {
auto encoder = std::make_unique<MultiStreamEncoder>();
return wrapMultiStreamEncoderPointerToTensor(std::move(encoder));
}

Expand Down Expand Up @@ -1252,8 +1238,23 @@ void streaming_encoder_add_video_stream(
std::move(extraOptionsMap));
}

void streaming_encoder_open(torch::stable::Tensor& encoder) {
unwrapTensorToGetMultiStreamEncoder(encoder)->open();
void streaming_encoder_open_file(
torch::stable::Tensor& encoder,
std::string filename) {
unwrapTensorToGetMultiStreamEncoder(encoder)->open(filename);
}

void streaming_encoder_open_file_like(
torch::stable::Tensor& encoder,
std::string format,
int64_t file_like_context) {
auto fileLikeContext =
reinterpret_cast<AVIOFileLikeContext*>(file_like_context);
STD_TORCH_CHECK(
fileLikeContext != nullptr, "file_like_context must be a valid pointer");
std::unique_ptr<AVIOFileLikeContext> avioContextHolder(fileLikeContext);
unwrapTensorToGetMultiStreamEncoder(encoder)->open(
format, std::move(avioContextHolder));
}

void streaming_encoder_add_audio_stream(
Expand Down Expand Up @@ -1357,19 +1358,18 @@ STABLE_TORCH_LIBRARY_IMPL(torchcodec_ns, BackendSelect, m) {
m.impl("encode_video_to_file", TORCH_BOX(&encode_video_to_file));
m.impl("encode_video_to_tensor", TORCH_BOX(&encode_video_to_tensor));
m.impl("_encode_video_to_file_like", TORCH_BOX(&_encode_video_to_file_like));
m.impl(
"create_streaming_encoder_to_file",
TORCH_BOX(&create_streaming_encoder_to_file));
m.impl(
"create_streaming_encoder_to_file_like",
TORCH_BOX(&create_streaming_encoder_to_file_like));
m.impl("create_streaming_encoder", TORCH_BOX(&create_streaming_encoder));
m.impl(
"streaming_encoder_add_video_stream",
TORCH_BOX(&streaming_encoder_add_video_stream));
m.impl(
"streaming_encoder_add_audio_stream",
TORCH_BOX(&streaming_encoder_add_audio_stream));
m.impl("streaming_encoder_open", TORCH_BOX(&streaming_encoder_open));
m.impl(
"streaming_encoder_open_file", TORCH_BOX(&streaming_encoder_open_file));
m.impl(
"streaming_encoder_open_file_like",
TORCH_BOX(&streaming_encoder_open_file_like));
m.impl(
"streaming_encoder_add_frames", TORCH_BOX(&streaming_encoder_add_frames));
m.impl(
Expand Down Expand Up @@ -1420,20 +1420,19 @@ STABLE_TORCH_LIBRARY_IMPL(torchcodec_ns, CPU, m) {
TORCH_BOX(&scan_all_streams_to_update_metadata));

m.impl("_get_backend_details", TORCH_BOX(&get_backend_details));
m.impl("create_streaming_encoder", TORCH_BOX(&create_streaming_encoder));
m.impl(
"create_streaming_encoder_to_file",
TORCH_BOX(&create_streaming_encoder_to_file));
"streaming_encoder_open_file", TORCH_BOX(&streaming_encoder_open_file));
m.impl(
"create_streaming_encoder_to_file_like",
TORCH_BOX(&create_streaming_encoder_to_file_like));
"streaming_encoder_open_file_like",
TORCH_BOX(&streaming_encoder_open_file_like));
m.impl("streaming_encoder_close", TORCH_BOX(&streaming_encoder_close));
m.impl(
"streaming_encoder_add_video_stream",
TORCH_BOX(&streaming_encoder_add_video_stream));
m.impl(
"streaming_encoder_add_audio_stream",
TORCH_BOX(&streaming_encoder_add_audio_stream));
m.impl("streaming_encoder_open", TORCH_BOX(&streaming_encoder_open));
m.impl(
"streaming_encoder_add_frames", TORCH_BOX(&streaming_encoder_add_frames));
m.impl(
Expand Down
49 changes: 25 additions & 24 deletions src/torchcodec/_core/ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -138,11 +138,8 @@ def add_video_stream(
torch.ops.torchcodec_ns._get_json_ffmpeg_library_versions.default
)
_get_backend_details = torch.ops.torchcodec_ns._get_backend_details.default
create_streaming_encoder_to_file = torch._dynamo.disallow_in_graph(
torch.ops.torchcodec_ns.create_streaming_encoder_to_file.default
)
_create_streaming_encoder_to_file_like = torch._dynamo.disallow_in_graph(
torch.ops.torchcodec_ns.create_streaming_encoder_to_file_like.default
create_streaming_encoder = torch._dynamo.disallow_in_graph(
torch.ops.torchcodec_ns.create_streaming_encoder.default
)
streaming_encoder_close = torch.ops.torchcodec_ns.streaming_encoder_close.default
streaming_encoder_add_video_stream = (
Expand All @@ -151,7 +148,12 @@ def add_video_stream(
streaming_encoder_add_audio_stream = (
torch.ops.torchcodec_ns.streaming_encoder_add_audio_stream.default
)
streaming_encoder_open = torch.ops.torchcodec_ns.streaming_encoder_open.default
streaming_encoder_open_file = (
torch.ops.torchcodec_ns.streaming_encoder_open_file.default
)
_streaming_encoder_open_file_like = (
torch.ops.torchcodec_ns.streaming_encoder_open_file_like.default
)
streaming_encoder_add_frames = (
torch.ops.torchcodec_ns.streaming_encoder_add_frames.default
)
Expand Down Expand Up @@ -271,14 +273,16 @@ def encode_video_to_file_like(
)


def create_streaming_encoder_to_file_like(
def streaming_encoder_open_file_like(
encoder: torch.Tensor,
format: str,
file_like: io.RawIOBase | io.BufferedIOBase,
) -> torch.Tensor:
) -> None:
assert _pybind_ops is not None
return _create_streaming_encoder_to_file_like(
_streaming_encoder_open_file_like(
encoder,
format,
_pybind_ops.create_file_like_context(file_like, True), # True means for writing
_pybind_ops.create_file_like_context(file_like, True),
)


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


@register_fake("torchcodec_ns::create_streaming_encoder_to_file")
def _create_streaming_encoder_to_file_abstract(
filename: str,
) -> torch.Tensor:
return torch.empty([], dtype=torch.long)


@register_fake("torchcodec_ns::create_streaming_encoder_to_file_like")
def _create_streaming_encoder_to_file_like_abstract(
format: str,
file_like_context: int,
) -> torch.Tensor:
@register_fake("torchcodec_ns::create_streaming_encoder")
def _create_streaming_encoder_abstract() -> torch.Tensor:
return torch.empty([], dtype=torch.long)


Expand Down Expand Up @@ -667,8 +661,15 @@ def streaming_encoder_add_audio_stream_abstract(
return


@register_fake("torchcodec_ns::streaming_encoder_open")
def streaming_encoder_open_abstract(encoder: torch.Tensor) -> None:
@register_fake("torchcodec_ns::streaming_encoder_open_file")
def streaming_encoder_open_file_abstract(encoder: torch.Tensor, filename: str) -> None:
return


@register_fake("torchcodec_ns::streaming_encoder_open_file_like")
def streaming_encoder_open_file_like_abstract(
encoder: torch.Tensor, format: str, file_like_context: int
) -> None:
return


Expand Down
18 changes: 9 additions & 9 deletions src/torchcodec/encoders/_multi_stream_encoder.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,13 +22,8 @@ def write(self, samples: Tensor) -> None:


class StreamingEncoder:
def __init__(self, dest, *, format: str | None = None):
if format is not None:
self._encoder_tensor = _core.create_streaming_encoder_to_file_like(
format, dest
)
else:
self._encoder_tensor = _core.create_streaming_encoder_to_file(str(dest))
def __init__(self):
self._encoder_tensor = _core.create_streaming_encoder()

def add_video(
self,
Expand Down Expand Up @@ -75,8 +70,13 @@ def add_audio(
)
return _AudioStream(self._encoder_tensor)

def open(self) -> None:
_core.streaming_encoder_open(self._encoder_tensor)
# TODO MultiStreamEncoder: Maybe there should 2 separate methods, one for
# file, one for file-like.
def open(self, dest, *, format: str | None = None) -> None:
if format is not None:
_core.streaming_encoder_open_file_like(self._encoder_tensor, format, dest)
else:
_core.streaming_encoder_open_file(self._encoder_tensor, str(dest))

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