(self, tensor)
| 124 | mgx.save(self.model, path, format="msgpack") |
| 125 | |
| 126 | def tensor_to_arg(self, tensor): |
| 127 | mgx_shape = mgx.shape(type=self.torch_to_mgx_dtype_dict[tensor.dtype], |
| 128 | lens=list(tensor.size()), |
| 129 | strides=list(tensor.stride())) |
| 130 | return mgx.argument_from_pointer(mgx_shape, tensor.data_ptr()) |
| 131 | |
| 132 | def prealloc_buffers(self, param_names): |
| 133 | for param_name in param_names: |
no test coverage detected