diff --git a/mlx/export.cpp b/mlx/export.cpp index 972febc2..01655fca 100644 --- a/mlx/export.cpp +++ b/mlx/export.cpp @@ -716,7 +716,7 @@ void FunctionExporter::export_with_callback( if (arr.has_primitive() || input_set.find(arr.id()) != input_set.end()) { continue; } - if (constants.insert(arr.id()).second) { + if (constants.insert({arr.id(), arr}).second) { new_constants.emplace_back(namer.get_name(arr), arr); } } @@ -848,7 +848,7 @@ void FunctionExporter::export_function(const Args& args, const Kwargs& kwargs) { if (input_set.find(arr.id()) == input_set.end()) { serialize(os, true); // Save constant data if not already saved - if (constants.insert(arr.id()).second) { + if (constants.insert({arr.id(), arr}).second) { serialize(os, arr.shape()); serialize(os, arr.dtype()); os.write(arr.data(), arr.nbytes()); diff --git a/mlx/export_impl.h b/mlx/export_impl.h index 0e781898..be215aaa 100644 --- a/mlx/export_impl.h +++ b/mlx/export_impl.h @@ -72,7 +72,7 @@ struct FunctionExporter { const std::vector& outputs, const std::vector& tape, const std::vector& kwarg_keys); - std::set constants; + std::unordered_map constants; int count{0}; bool closed{false}; std::shared_ptr ftable; diff --git a/python/tests/test_export_import.py b/python/tests/test_export_import.py index 1aa251b6..6e280a79 100644 --- a/python/tests/test_export_import.py +++ b/python/tests/test_export_import.py @@ -575,6 +575,27 @@ class TestExportImport(mlx_tests.MLXTestCase): out = imported(a)[0] self.assertTrue(mx.allclose(expected, out)) + def test_export_import_multi_with_constants(self): + + path = os.path.join(self.test_dir, "fn.mlxfn") + + def fun(y): + i = y.shape[0] + x = mx.array(i) + for j in range(10): + x = x + mx.array(i + j) + return x * y.sum() + + ys = [mx.array([1]), mx.array([1, 1]), mx.array([1, 1, 1])] + + with mx.exporter(path, fun) as exporter: + for y in ys: + exporter(y) + + imported = mx.import_function(path) + for y in ys: + self.assertEqual(imported(y)[0].item(), fun(y).item()) + if __name__ == "__main__": mlx_tests.MLXTestRunner()