MCPcopy Create free account
hub / github.com/NVIDIA/cuda-samples / main

Function main

python/2_CoreConcepts/tmaTensorMap/tmaTensorMap.py:179–273  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

177
178
179def main() -> int:
180 import argparse
181
182 parser = argparse.ArgumentParser(
183 description="Use a TMA tensor map to bulk-copy data on Hopper+ GPUs"
184 )
185 parser.add_argument(
186 "--elements",
187 type=int,
188 default=1024,
189 help="Total number of float32 elements (must be a multiple of 128)",
190 )
191 parser.add_argument("--device", type=int, default=0, help="CUDA device id")
192 args = parser.parse_args()
193
194 if args.elements % TILE_SIZE != 0:
195 print(f"--elements must be a multiple of TILE_SIZE={TILE_SIZE}")
196 return 1
197
198 dev = Device(args.device)
199 print_gpu_info(dev)
200
201 arch = dev.compute_capability
202 if arch < (9, 0):
203 print(
204 f"\nTMA requires compute capability >= 9.0 (Hopper or later); "
205 f"this device is {arch.major}.{arch.minor}. Exiting cleanly."
206 )
207 return 0
208
209 dev.set_current()
210 include_path = _get_cccl_include_paths()
211
212 # Compile with the CUBIN code type to target the exact device arch.
213 prog = Program(
214 KERNEL_SRC,
215 code_type="c++",
216 options=ProgramOptions(
217 std="c++17",
218 arch=f"sm_{dev.arch}",
219 include_path=include_path,
220 ),
221 )
222 mod = prog.compile("cubin")
223 kernel = mod.get_kernel("tma_copy")
224
225 # (1) Prepare input data and verify the initial TMA copy.
226 n = args.elements
227 src = cp.arange(n, dtype=cp.float32)
228 output = cp.zeros(n, dtype=cp.float32)
229 dev.sync() # CuPy uses its own stream
230
231 tensor_map = StridedMemoryView.from_any_interface(src, stream_ptr=-1).as_tensor_map(
232 box_dim=(TILE_SIZE,)
233 )
234
235 n_tiles = n // TILE_SIZE
236 config = LaunchConfig(grid=n_tiles, block=TILE_SIZE)

Callers 1

tmaTensorMap.pyFile · 0.70

Calls 3

print_gpu_infoFunction · 0.90
_get_cccl_include_pathsFunction · 0.85
fullMethod · 0.45

Tested by

no test coverage detected