@@ -88,15 +88,15 @@ STABLE_TORCH_LIBRARY_FRAGMENT(torchcodec_ns, m) {
8888 m.def (" _blocks_packet_decoder_send_eof(Tensor(a!) decoder) -> int" );
8989 m.def (" _blocks_packet_decoder_reset(Tensor(a!) decoder) -> ()" );
9090 m.def (
91- " _blocks_packet_decoder_receive_frame(Tensor(a!) decoder) -> (Tensor, int, float, float, str , Tensor)" );
91+ " _blocks_packet_decoder_receive_frame(Tensor(a!) decoder) -> (Tensor, int, float, float, Device , Tensor)" );
9292 m.def (
9393 " _blocks_create_color_converter(str device=\" cpu\" , str output_dtype=\" uint8\" ) -> Tensor" );
9494 m.def (
95- " _blocks_convert_frame(Tensor(a!) converter, Tensor frame, str device) -> Tensor" );
95+ " _blocks_convert_frame(Tensor(a!) converter, Tensor frame, Device device) -> Tensor" );
9696 m.def (
9797 " _blocks_frame_metadata(Tensor frame) -> (str, str, str, int, int, int, float)" );
9898 m.def (
99- " _blocks_frame_planes(Tensor frame, str device) -> (Tensor, Tensor, Tensor, Tensor)" );
99+ " _blocks_frame_planes(Tensor frame, Device device) -> (Tensor, Tensor, Tensor, Tensor)" );
100100 m.def (" _get_key_frame_indices(Tensor(a!) decoder) -> Tensor" );
101101 m.def (" get_json_metadata(Tensor(a!) decoder) -> str" );
102102 m.def (" get_container_json_metadata(Tensor(a!) decoder) -> str" );
@@ -904,7 +904,7 @@ using OpsReceiveFrameOutput = std::tuple<
904904 int64_t ,
905905 double ,
906906 double ,
907- std::string ,
907+ StableDevice ,
908908 torch::stable::Tensor>;
909909
910910OpsReceiveFrameOutput _blocks_packet_decoder_receive_frame (
@@ -919,7 +919,7 @@ OpsReceiveFrameOutput _blocks_packet_decoder_receive_frame(
919919 static_cast <int64_t >(status),
920920 0.0 ,
921921 0.0 ,
922- std::string ( " cpu " ),
922+ StableDevice ( kStableCPU ),
923923 torch::stable::empty ({int64_t (0 )}, kStableUInt8 ));
924924 }
925925 AVRational time_base = decoder_ptr->time_base ();
@@ -929,7 +929,7 @@ OpsReceiveFrameOutput _blocks_packet_decoder_receive_frame(
929929 // ones a CUDA decoder had to decode on the CPU and upload.
930930 // TODO_API_BREAKDOWN DESIGN P1: Not sure we need to return the device at all,
931931 // the device should always match the device parameter now.
932- std::string device = device_to_string ( decoder_ptr->device () );
932+ StableDevice device = decoder_ptr->device ();
933933 torch::stable::Tensor storage =
934934 decoder_ptr->get_frame_storage (*av_frame).value_or (
935935 torch::stable::empty ({int64_t (0 )}, kStableUInt8 ));
@@ -956,7 +956,7 @@ torch::stable::Tensor _blocks_create_color_converter(
956956torch::stable::Tensor _blocks_convert_frame (
957957 torch::stable::Tensor& converter,
958958 torch::stable::Tensor& frame,
959- std::string device) {
959+ StableDevice device) {
960960 ColorConverter* converter_ptr =
961961 unwrap_tensor_to_pointer<ColorConverter>(converter);
962962 return converter_ptr->convert (
@@ -994,10 +994,10 @@ using OpsFramePlanesOutput = std::tuple<
994994
995995OpsFramePlanesOutput _blocks_frame_planes (
996996 torch::stable::Tensor& tensor_handle,
997- std::string device) {
997+ StableDevice device) {
998998 AVFrame* av_frame = unwrap_tensor_to_pointer<AVFrame>(tensor_handle);
999999 std::vector<torch::stable::Tensor> planes =
1000- frame_planes (*av_frame, StableDevice ( device) , tensor_handle);
1000+ frame_planes (*av_frame, device, tensor_handle);
10011001
10021002 // Op schema wants a fixed number of planes, so we pad with empty tensors that
10031003 // then get removed at the Python level.
0 commit comments