@@ -112,6 +112,22 @@ namespace pybindings {
112112
113113namespace {
114114
115+ void * mutable_tensor_data_ptr_no_cow (at::Tensor& tensor) {
116+ if (tensor.numel () == 0 ) {
117+ return nullptr ;
118+ }
119+
120+ auto * storage_data = static_cast <char *>(tensor.unsafeGetTensorImpl ()
121+ ->unsafe_storage ()
122+ .unsafeGetStorageImpl ()
123+ ->_mutable_data_ptr_no_checks ()
124+ .mutable_get ());
125+ ET_CHECK_MSG (
126+ storage_data != nullptr ,
127+ " Tensor has a non-zero number of elements, but its data is not allocated" );
128+ return storage_data + tensor.storage_offset () * tensor.itemsize ();
129+ }
130+
115131void write_data_to_file (const std::string& path, void * buf, size_t size) {
116132 FILE * f = fopen (path.c_str (), " w+" );
117133 if (!f) {
@@ -1117,12 +1133,14 @@ struct PyMethod final {
11171133 " should be contiguous or channels-last." ;
11181134 throw std::runtime_error (error_msg);
11191135 }
1120- TensorPtr tensor =
1121- for_blob (at_tensor.data_ptr (), std::move (sizes), type)
1122- .strides (std::move (strides))
1123- .dim_order (std::move (dim_order))
1124- .dynamism (aten::TensorShapeDynamism::STATIC )
1125- .make_tensor_ptr ();
1136+ TensorPtr tensor = for_blob (
1137+ mutable_tensor_data_ptr_no_cow (at_tensor),
1138+ std::move (sizes),
1139+ type)
1140+ .strides (std::move (strides))
1141+ .dim_order (std::move (dim_order))
1142+ .dynamism (aten::TensorShapeDynamism::STATIC )
1143+ .make_tensor_ptr ();
11261144 input_tensors.push_back (tensor);
11271145 EValue evalue (input_tensors.back ());
11281146#endif
0 commit comments