已合并
[feat] support environment variable LD_PRELOAD #34371
liujunzhu创建于 4月25日
[feat] support environment variable LD_PRELOAD #34371
已合并
liujunzhu创建于 4月25日
共 2 个文件变更+226-9
@@ -0,0 +1,193 @@
1+# Integration test: verify that LD_PRELOAD-injected overrides of ACL symbols
2+# loaded via FunctionLoader (e.g. aclrtMallocAlign32) are picked up.
3+#
4+# Requires: an NPU environment with libascendcl available, a C compiler.
5+# Run: python test/npu/test_ld_preload_acl_hook.py
6+ 
7+import os
8+import shutil
9+import subprocess
10+import sys
11+import tempfile
12+import textwrap
13+import unittest
14+ 
15+import torch_npu
16+from torch_npu.testing.testcase import TestCase, run_tests
17+ 
18+ 
19+HOOK_SRC = r"""
20+#define _GNU_SOURCE
21+#include <dlfcn.h>
22+#include <stdio.h>
23+#include <stdlib.h>
24+ 
25+typedef int aclError;
26+typedef int aclrtMemMallocPolicy;
27+ 
28+__attribute__((visibility("default")))
29+aclError aclrtMallocAlign32(void **devPtr, size_t size, aclrtMemMallocPolicy policy)
30+{
31+ const char *flag = getenv("ACL_HOOK_FLAG_FILE");
32+ if (flag != NULL) {
33+ FILE *fp = fopen(flag, "w");
34+ if (fp != NULL) {
35+ fprintf(fp, "hit\n");
36+ fclose(fp);
37+ }
38+ }
39+ typedef aclError (*real_fn_t)(void **, size_t, aclrtMemMallocPolicy);
40+ static real_fn_t real = NULL;
41+ if (real == NULL) {
42+ real = (real_fn_t)dlsym(RTLD_NEXT, "aclrtMallocAlign32");
43+ }
44+ if (real == NULL) {
45+ return -1;
46+ }
47+ return real(devPtr, size, policy);
48+}
49+"""
50+ 
51+ 
52+# Minimal Python payload: allocate a tensor on NPU, which eventually exercises
53+# aclrtMallocAlign32 via the caching allocator's large-block path.
54+PAYLOAD_SRC = textwrap.dedent(
55+ """
56+ import sys
57+ import torch
58+ import torch_npu
59+ torch.npu.set_device(0)
60+ t = torch.empty(1024 * 1024, dtype=torch.float32, device='npu')
61+ del t
62+ torch.npu.empty_cache()
63+ sys.exit(0)
64+ """
65+)
66+ 
67+ 
68+def _compile_hook(tmpdir, cc):
69+ src_path = os.path.join(tmpdir, "hook.c")
70+ so_path = os.path.join(tmpdir, "libacl_hook_test.so")
71+ with open(src_path, "w") as f:
72+ f.write(HOOK_SRC)
73+ subprocess.check_call(
74+ [cc, "-shared", "-fPIC", "-O2", "-o", so_path, src_path, "-ldl"]
75+ )
76+ return so_path
77+ 
78+ 
79+def _compile_unrelated(tmpdir, cc):
80+ src_path = os.path.join(tmpdir, "other.c")
81+ so_path = os.path.join(tmpdir, "libunrelated_hook.so")
82+ with open(src_path, "w") as f:
83+ f.write("int unrelated_noop(void) { return 0; }\n")
84+ subprocess.check_call(
85+ [cc, "-shared", "-fPIC", "-O2", "-o", so_path, src_path]
86+ )
87+ return so_path
88+ 
89+ 
90+def _run_payload(env):
91+ with tempfile.NamedTemporaryFile("w", suffix=".py", delete=False) as f:
92+ f.write(PAYLOAD_SRC)
93+ script = f.name
94+ try:
95+ return subprocess.run(
96+ [sys.executable, script],
97+ env=env,
98+ capture_output=True,
99+ text=True,
100+ timeout=120,
101+ )
102+ finally:
103+ os.unlink(script)
104+ 
105+ 
106+@unittest.skipIf(
107+ not torch_npu.npu.is_available(), "NPU not available; skipping LD_PRELOAD test"
108+)
109+class TestLdPreloadAclHook(TestCase):
110+ @classmethod
111+ def setUpClass(cls):
112+ cc = os.environ.get("CC", "gcc")
113+ if shutil.which(cc) is None:
114+ raise unittest.SkipTest(
115+ "C compiler '{}' not found; skipping LD_PRELOAD test".format(cc)
116+ )
117+ cls._cc = cc
118+ cls._tmpdir = tempfile.mkdtemp(prefix="acl_hook_")
119+ try:
120+ cls._hook_so = _compile_hook(cls._tmpdir, cc)
121+ cls._unrelated_so = _compile_unrelated(cls._tmpdir, cc)
122+ except Exception:
123+ shutil.rmtree(cls._tmpdir, ignore_errors=True)
124+ raise
125+ 
126+ @classmethod
127+ def tearDownClass(cls):
128+ shutil.rmtree(cls._tmpdir, ignore_errors=True)
129+ 
130+ def _base_env(self):
131+ env = os.environ.copy()
132+ env.pop("LD_PRELOAD", None)
133+ return env
134+ 
135+ def test_without_preload_behaves_unchanged(self):
136+ """No LD_PRELOAD: allocation succeeds, hook never invoked."""
137+ flag = os.path.join(self._tmpdir, "flag_no_preload")
138+ env = self._base_env()
139+ env["ACL_HOOK_FLAG_FILE"] = flag
140+ result = _run_payload(env)
141+ if result.returncode != 0:
142+ sys.stderr.write(result.stderr)
143+ self.assertEqual(result.returncode, 0)
144+ self.assertFalse(
145+ os.path.exists(flag), "hook must not be called when LD_PRELOAD is unset"
146+ )
147+ 
148+ def test_with_preload_hook_is_invoked(self):
149+ """LD_PRELOAD set to hook .so: hook is called for aclrtMallocAlign32."""
150+ flag = os.path.join(self._tmpdir, "flag_preload")
151+ env = self._base_env()
152+ env["LD_PRELOAD"] = self._hook_so
153+ env["ACL_HOOK_FLAG_FILE"] = flag
154+ result = _run_payload(env)
155+ if result.returncode != 0:
156+ sys.stderr.write(result.stderr)
157+ self.assertEqual(result.returncode, 0)
158+ self.assertTrue(
159+ os.path.exists(flag),
160+ "hook must be called when LD_PRELOAD overrides aclrtMallocAlign32",
161+ )
162+ 
163+ def test_preload_without_symbol_falls_back(self):
164+ """LD_PRELOAD set to a .so that does NOT define ACL symbols:
165+ allocation must still succeed via fallback to libascendcl."""
166+ flag = os.path.join(self._tmpdir, "flag_unrelated")
167+ env = self._base_env()
168+ env["LD_PRELOAD"] = self._unrelated_so
169+ env["ACL_HOOK_FLAG_FILE"] = flag
170+ result = _run_payload(env)
171+ if result.returncode != 0:
172+ sys.stderr.write(result.stderr)
173+ self.assertEqual(result.returncode, 0)
174+ self.assertFalse(
175+ os.path.exists(flag),
176+ "unrelated preload must not cause the hook flag to be set",
177+ )
178+ 
179+ def test_multiple_preload_sos(self):
180+ """LD_PRELOAD with multiple .so (hook first): hook still wins."""
181+ flag = os.path.join(self._tmpdir, "flag_multi")
182+ env = self._base_env()
183+ env["LD_PRELOAD"] = "{}:{}".format(self._hook_so, self._unrelated_so)
184+ env["ACL_HOOK_FLAG_FILE"] = flag
185+ result = _run_payload(env)
186+ if result.returncode != 0:
187+ sys.stderr.write(result.stderr)
188+ self.assertEqual(result.returncode, 0)
189+ self.assertTrue(os.path.exists(flag), "first-loaded preload must win")
190+ 
191+ 
192+if __name__ == "__main__":
193+ run_tests()
@@ -1,3 +1,4 @@
1+#include <cstdlib>
1#include "torch_npu/csrc/core/npu/NPUException.h"2#include "torch_npu/csrc/core/npu/NPUException.h"
2#include "torch_npu/csrc/core/npu/register/FunctionLoader.h"3#include "torch_npu/csrc/core/npu/register/FunctionLoader.h"
3 4 
@@ -24,14 +25,7 @@ void FunctionLoader::Set(const std::string &name)
24 25 
25void *FunctionLoader::Get(const std::string &name)26void *FunctionLoader::Get(const std::string &name)
26{27{
27- if (this->handle == nullptr) {28+ std::lock_guard<std::mutex> lock(this->mu_);
28- auto handle = dlopen(this->fileName.c_str(), this->flags);
29- if (handle == nullptr) {
30- AT_ERROR(dlerror());
31- return nullptr;
32- }
33- this->handle = handle;
34- }
35 29 
36 auto itr = registry.find(name);30 auto itr = registry.find(name);
37 if (itr == registry.end()) {31 if (itr == registry.end()) {
@@ -43,7 +37,37 @@ void *FunctionLoader::Get(const std::string &name)
43 return itr->second;37 return itr->second;
44 }38 }
45 39 
46- auto func = dlsym(this->handle, name.c_str());40+ // When LD_PRELOAD is set, prefer RTLD_DEFAULT so that symbols overridden
41+ // by the user's preloaded .so take precedence over libascendcl.so's own
42+ // implementation. Opt-in by the presence of LD_PRELOAD itself; without
43+ // LD_PRELOAD the behavior is identical to the original path.
44+ static const bool preload_enabled = []() {
45+ const char *env = std::getenv("LD_PRELOAD");
46+ bool enabled = (env != nullptr && env[0] != '\0');
47+ if (enabled) {
48+ TORCH_NPU_WARN_ONCE("LD_PRELOAD detected, FunctionLoader prefers "
49+ "RTLD_DEFAULT for symbol resolution.");
50+ }
51+ return enabled;
52+ }();
53+ 
54+ void *func = nullptr;
55+ if (preload_enabled) {
56+ func = dlsym(RTLD_DEFAULT, name.c_str());
57+ }
58+ 
59+ if (func == nullptr) {
60+ if (this->handle == nullptr) {
61+ auto handle = dlopen(this->fileName.c_str(), this->flags);
62+ if (handle == nullptr) {
63+ AT_ERROR(dlerror());
64+ return nullptr;
65+ }
66+ this->handle = handle;
67+ }
68+ func = dlsym(this->handle, name.c_str());
69+ }
70+ 
47 if (func == nullptr) {71 if (func == nullptr) {
48 return nullptr;72 return nullptr;
49 }73 }