Fix stale transform copy-chain leaks (#3290)

This commit is contained in:
LongYinan
2026-03-24 14:15:23 -07:00
committed by GitHub
parent e40ada3fe2
commit 604c825538
2 changed files with 78 additions and 8 deletions
+54
View File
@@ -17,6 +17,24 @@
using namespace mlx::core;
namespace {
int count_graph_nodes(const array& x, const std::string& node_name) {
std::ostringstream oss;
print_graph(oss, x);
auto graph = oss.str();
int count = 0;
size_t pos = 0;
while ((pos = graph.find(node_name, pos)) != std::string::npos) {
count++;
pos += node_name.size();
}
return count;
}
} // namespace
TEST_CASE("test stop gradient") {
auto x = zeros({5, 5});
auto y = stop_gradient(x);
@@ -328,6 +346,42 @@ TEST_CASE("test grad") {
}
}
TEST_CASE("test transform container reuse does not accumulate stale wrappers") {
auto x = ones({128});
SUBCASE("grad reuses a single copy wrapper") {
std::vector<array> container = {array(1.0f)};
auto grad_fn = grad([&container](const std::vector<array>& inputs) {
container[0] = inputs[0];
return sum(inputs[1]);
});
for (int i = 0; i < 5; ++i) {
auto grads = grad_fn({container[0], x});
eval(grads);
}
CHECK_EQ(count_graph_nodes(container[0], "Copy "), 1);
}
SUBCASE("jvp reuses a single copy wrapper") {
std::vector<array> container = {array(1.0f)};
auto fun = [&container](const std::vector<array>& inputs) {
container[0] = inputs[0];
return std::vector<array>{sum(inputs[1])};
};
for (int i = 0; i < 5; ++i) {
auto [outputs, tangents] =
jvp(fun, {container[0], x}, {array(1.0f), ones({128})});
eval(outputs);
eval(tangents);
}
CHECK_EQ(count_graph_nodes(container[0], "Copy "), 1);
}
}
TEST_CASE("test creation grads") {
// Test astype
{