3131from tvm .target import Target
3232from tvm .contrib import ndk
3333from tvm import tir , DataType
34+ from tvm .rpc import connect_tracker
3435
3536
36- class RemoteConnection :
37- def __init__ (self ):
38- self .RPC_TRACKER_HOST = os .getenv ("TVM_TRACKER_HOST" , "localhost" )
39- self .RPC_TRACKER_PORT = int (os .getenv ("TVM_TRACKER_PORT" , 7979 ))
40- self .RPC_KEY = os .getenv ("RPC_DEVICE_KEY" , "android" )
41- self .tracker = tvm .rpc .connect_tracker (self .RPC_TRACKER_HOST , self .RPC_TRACKER_PORT )
37+ def get_rpc ():
38+ """
39+ Establish an RPC connection to the remote device.
4240
43- def __enter__ (self ):
44- self .remote = self .tracker .request (self .RPC_KEY , priority = 0 , session_timeout = 600 )
45- return self .remote
46-
47- def __exit__ (self , exc_type , exc_value , traceback ):
48- self .remote .get_function ("CloseRPCConnection" )()
41+ Returns
42+ -------
43+ tvm.rpc.RPCSession or None
44+ The RPC session object if RPC_TARGET is set; otherwise, None.
45+ """
46+ rpc_target = os .getenv ("RPC_TARGET" , None )
47+ if rpc_target :
48+ host = os .getenv ("TVM_TRACKER_HOST" , "localhost" )
49+ port = int (os .getenv ("TVM_TRACKER_PORT" , 9090 ))
50+ device_key = os .getenv ("RPC_DEVICE_KEY" , "android" )
51+ tracker = connect_tracker (host , port )
52+ return tracker .request (device_key , priority = 1 , session_timeout = 1000 )
53+ else :
54+ return None
4955
5056
5157def preprocess_pipeline (mod : IRModule ) -> IRModule :
@@ -96,14 +102,13 @@ def postprocess_pipeline(mod: IRModule) -> IRModule:
96102
97103
98104@tvm .testing .requires_rpc
99- @tvm .testing .requires_opencl
100- @pytest .mark .parametrize (
101- "target" , [Target ("opencl -device=adreno" , "llvm -mtriple=aarch64-linux-android" )]
102- )
105+ @tvm .testing .requires_adreno_opencl
106+ @pytest .mark .parametrize ("backend" , ["opencl" ])
103107@pytest .mark .parametrize ("dtype" , ["int8" , "float16" , "int16" , "float32" , "int32" ])
104108@pytest .mark .parametrize ("channel_size" , [64 , 128 ])
105109@pytest .mark .parametrize ("read_width" , [1 , 2 , 4 , 8 , 16 ])
106- def test_texture_copy (target , dtype , channel_size , read_width ):
110+ def test_texture_copy (backend , dtype , channel_size , read_width ):
111+ remote = get_rpc ()
107112 M , N , K = (256 , 1024 , 128 )
108113 lanes = channel_size // DataType (dtype ).bits
109114 if read_width > lanes :
@@ -139,6 +144,12 @@ def schedule_default(blk, lanes):
139144 schedule_default (B_blk , read_width )
140145
141146 mod = TextureCopy
147+
148+ if remote is None :
149+ target = Target (backend + " -device=adreno" )
150+ else :
151+ target = Target (backend + " -device=adreno" , "llvm -mtriple=aarch64-linux-android" )
152+
142153 with target :
143154 mod = preprocess_pipeline (mod )
144155 sch = tir .Schedule (mod )
@@ -148,20 +159,43 @@ def schedule_default(blk, lanes):
148159 ex = relax .build (mod , target )
149160 load_path = "vm_library.so"
150161 inputs = [np .random .randint (0 , 128 , (M , N )).astype (dtype ), np .zeros ((M , N ), dtype )]
151- with RemoteConnection () as remote :
152- with tempfile . TemporaryDirectory () as temp_dir :
162+ with tempfile . TemporaryDirectory () as temp_dir :
163+ if remote is not None :
153164 path = temp_dir + "/" + load_path
154165 ex .export_library (path , fcompile = ndk .create_shared , options = ["-shared" , "-fPIC" , "-lm" ])
155-
156166 remote .upload (path )
157167 rexec = remote .load_module (load_path )
158168 dev = remote .cl ()
159-
160- vm = relax .VirtualMachine (rexec , [dev , dev , dev ])
161- inps = [tvm .runtime .tensor (inp , dev ) for inp in inputs ]
162- vm ["main" ](* inps )
163-
164- np .testing .assert_equal (inps [- 1 ].numpy (), inps [0 ].numpy ())
169+ if "vdevice" in mod .global_infos :
170+ device_arr = [dev for ii in range (len (mod .global_infos ["vdevice" ]))]
171+ else :
172+ device_arr = [dev ]
173+ vm = relax .VirtualMachine (rexec , device_arr )
174+ else :
175+ # local execution
176+ if "opencl" in backend :
177+ dev = tvm .opencl (0 )
178+ elif "vulkan" in backend :
179+ dev = tvm .vulkan (0 )
180+ else :
181+ raise RuntimeError ("Unsupported backend" )
182+
183+ if "vdevice" in mod .global_infos :
184+ device_arr = [dev for ii in range (len (mod .global_infos ["vdevice" ]))]
185+ else :
186+ device_arr = [dev ]
187+ vm = relax .VirtualMachine (ex , device_arr )
188+
189+ inps = [tvm .runtime .tensor (inp , dev ) for inp in inputs ]
190+ vm ["main" ](* inps )
191+
192+ out1 = inps [- 1 ].numpy ()
193+ out2 = inps [0 ].numpy ()
194+
195+ if remote :
196+ remote .get_function ("CloseRPCConnection" )()
197+
198+ np .testing .assert_equal (out1 , out2 )
165199
166200
167201if __name__ == "__main__" :
0 commit comments