Skip to content

Commit bb2683b

Browse files
Metal backend: stop tracking blob memory once a model is destroyed (pytorch#22954)
### Summary Fixes pytorch#22953 Tensors created from a model's constants blob are tracked in `memory_to_n_tensor` as `NOT_OWN`, and `aoti_torch_delete_tensor_object` leaves those entries in place when it deletes such a tensor. `cleanup_memory()` cleared `tensors` but not that map, so the addresses stayed tracked after the model was gone. The next model whose constants landed on a reused address then failed to load with "Memory address ... is already being tracked by another tensor". Loading six small Metal models twice in one process failed 7 of 12 loads. `cleanup_memory()` now clears `memory_to_n_tensor` once every tensor has been deleted. The entry is not erased in the `NOT_OWN` branch of `aoti_torch_delete_tensor_object` because several tensors can alias one `NOT_OWN` address without a count, and erasing on the first delete would make the later ones fail with "memory not found during deletion". ### Test plan New `backends/apple/metal/runtime/test/test_memory.cpp`, wired through `et_cxx_test` under `BUILD_TESTING` the same way `backends/cuda` does it: ``` cmake -S . -B cmake-out-test -DEXECUTORCH_BUILD_METAL=ON -DAOTI_METAL=ON \ -DEXECUTORCH_BUILD_EXTENSION_TENSOR=ON -DEXECUTORCH_BUILD_EXTENSION_FLAT_TENSOR=ON \ -DEXECUTORCH_BUILD_EXTENSION_DATA_LOADER=ON -DEXECUTORCH_BUILD_EXTENSION_MODULE=ON \ -DEXECUTORCH_BUILD_EXTENSION_NAMED_DATA_MAP=ON -DEXECUTORCH_BUILD_TESTS=ON cmake --build cmake-out-test --target test_metal_memory cmake-out-test/backends/apple/metal/test_metal_memory ``` - Before the change: `BlobAddressCanBeReusedAfterCleanup` and `CleanupLeavesNoTrackedMemory` fail; `BlobAddressIsTrackedWhileTensorIsAlive` (the control) passes. - After: `[ PASSED ] 3 tests.` - End to end, the load/destroy/load loop from the issue went from `12 loads, 7 failures` to `12 loads, 0 failures` in three consecutive runs. - `python -m unittest backends.apple.metal.tests.test_modules.TestMetalBackendModules`: `Ran 134 tests ... OK`. - `lintrunner` clean on the touched files.
1 parent 32fb206 commit bb2683b

3 files changed

Lines changed: 109 additions & 0 deletions

File tree

‎backends/apple/metal/CMakeLists.txt‎

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -139,3 +139,12 @@ install(
139139
EXPORT ExecuTorchTargets
140140
DESTINATION lib
141141
)
142+
143+
if(BUILD_TESTING)
144+
include(${EXECUTORCH_ROOT}/tools/cmake/Test.cmake)
145+
146+
et_cxx_test(
147+
test_metal_memory SOURCES runtime/test/test_memory.cpp EXTRA_LIBS
148+
metal_backend
149+
)
150+
endif()

‎backends/apple/metal/runtime/shims/memory.cpp‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -706,6 +706,13 @@ void cleanup_memory() {
706706
// tensors map should now be empty, but ensure it's cleared
707707
tensors.clear();
708708

709+
// Tensors created from a blob are tracked as NOT_OWN and
710+
// aoti_torch_delete_tensor_object leaves their address in the map, since
711+
// several of them may alias it. With every tensor gone nothing is tracked
712+
// anymore, and a stale entry would make the next model fail to load as soon
713+
// as its constants land on an address used before.
714+
memory_to_n_tensor.clear();
715+
709716
// Clean up Metal resources
710717
metal_cleanup_resources();
711718

Lines changed: 93 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,93 @@
1+
/*
2+
* Copyright (c) Meta Platforms, Inc. and affiliates.
3+
* All rights reserved.
4+
*
5+
* This source code is licensed under the BSD-style license found in the
6+
* LICENSE file in the root directory of this source tree.
7+
*/
8+
9+
#include <gtest/gtest.h>
10+
11+
#include <cstdint>
12+
#include <vector>
13+
14+
#include <executorch/backends/apple/metal/runtime/shims/memory.h>
15+
#include <executorch/runtime/core/error.h>
16+
#include <executorch/runtime/platform/platform.h>
17+
18+
using namespace executorch::backends::metal;
19+
using executorch::runtime::Error;
20+
21+
namespace {
22+
23+
// ScalarType::Float, as AOTInductor passes it to the shims.
24+
constexpr int32_t kFloat32 = 6;
25+
// DeviceType::MPS.
26+
constexpr int32_t kDeviceMps = 13;
27+
28+
} // namespace
29+
30+
class MetalMemoryTest : public ::testing::Test {
31+
protected:
32+
void SetUp() override {
33+
et_pal_init();
34+
}
35+
36+
void TearDown() override {
37+
cleanup_memory();
38+
}
39+
40+
// Wraps `data` the way a model's constants are wrapped: a tensor that
41+
// borrows memory it does not own.
42+
Error createFromBlob(void* data, AOTITensorHandle* tensor) {
43+
const std::vector<int64_t> sizes = {2, 2};
44+
const std::vector<int64_t> strides = {2, 1};
45+
return aoti_torch_create_tensor_from_blob_v2(
46+
data,
47+
static_cast<int64_t>(sizes.size()),
48+
sizes.data(),
49+
strides.data(),
50+
/*storage_offset=*/0,
51+
kFloat32,
52+
kDeviceMps,
53+
/*device_index=*/0,
54+
tensor,
55+
/*layout=*/0,
56+
/*opaque_metadata=*/nullptr,
57+
/*opaque_metadata_size=*/0);
58+
}
59+
60+
std::vector<float> blob_ = std::vector<float>(4, 1.0f);
61+
};
62+
63+
TEST_F(MetalMemoryTest, BlobAddressIsTrackedWhileTensorIsAlive) {
64+
AOTITensorHandle first = nullptr;
65+
ASSERT_EQ(createFromBlob(blob_.data(), &first), Error::Ok);
66+
67+
AOTITensorHandle second = nullptr;
68+
EXPECT_NE(createFromBlob(blob_.data(), &second), Error::Ok);
69+
}
70+
71+
// A model's constants blob can be mapped at an address a previously destroyed
72+
// model used. cleanup_memory() runs when a model is destroyed, so it must leave
73+
// nothing tracked behind.
74+
TEST_F(MetalMemoryTest, BlobAddressCanBeReusedAfterCleanup) {
75+
AOTITensorHandle first = nullptr;
76+
ASSERT_EQ(createFromBlob(blob_.data(), &first), Error::Ok);
77+
78+
cleanup_memory();
79+
80+
AOTITensorHandle second = nullptr;
81+
EXPECT_EQ(createFromBlob(blob_.data(), &second), Error::Ok);
82+
}
83+
84+
TEST_F(MetalMemoryTest, CleanupLeavesNoTrackedMemory) {
85+
AOTITensorHandle tensor = nullptr;
86+
ASSERT_EQ(createFromBlob(blob_.data(), &tensor), Error::Ok);
87+
ASSERT_FALSE(memory_to_n_tensor.empty());
88+
89+
cleanup_memory();
90+
91+
EXPECT_TRUE(tensors.empty());
92+
EXPECT_TRUE(memory_to_n_tensor.empty());
93+
}

0 commit comments

Comments
 (0)