@@ -219,6 +219,7 @@ LogicalResult CudaBackend::makeASM(MLIRContext &context, ModuleOp module) {
219219 llvm::errs () << " Could not find kernel name in PTX output\n " ;
220220 return LogicalResult::failure ();
221221 }
222+ m_metadata.name = kernel_name;
222223
223224 // 4. Post-process version and target
224225 char ptx_major_minor[8 ];
@@ -233,6 +234,20 @@ LogicalResult CudaBackend::makeASM(MLIRContext &context, ModuleOp module) {
233234 // No 'knobs' defined; always remove for now or leave as TODO.
234235 ret = std::regex_replace (ret, std::regex (R"( ,\s*debug|debug,\s*)" ), " " );
235236
237+ // 6. Append kernel metadata as PTX line comments. PTX comments are ignored
238+ // by ptxas and the CUDA driver, so this is safe. The Rust side parses
239+ // these lines to recover launch parameters without a separate FFI channel.
240+ ret += " \n // --- triton-metadata ---\n " ;
241+ ret += " // meta:name=" + m_metadata.name + " \n " ;
242+ ret += " // meta:num_warps=" + std::to_string (m_metadata.num_warps ) + " \n " ;
243+ ret += " // meta:num_ctas=" + std::to_string (m_metadata.num_ctas ) + " \n " ;
244+ ret += " // meta:shared=" + std::to_string (m_metadata.shared ) + " \n " ;
245+ ret += " // meta:tmem_size=" + std::to_string (m_metadata.tmem_size ) + " \n " ;
246+ ret += " // meta:global_scratch_size=" + std::to_string (m_metadata.global_scratch_size ) + " \n " ;
247+ ret += " // meta:global_scratch_align=" + std::to_string (m_metadata.global_scratch_align ) + " \n " ;
248+ ret += " // meta:profile_scratch_size=" + std::to_string (m_metadata.profile_scratch_size ) + " \n " ;
249+ ret += " // meta:profile_scratch_align=" + std::to_string (m_metadata.profile_scratch_align ) + " \n " ;
250+
236251 // 7. Save PTX (exposed via getASM())
237252 m_asm = std::move (ret);
238253 return LogicalResult::success ();
@@ -474,45 +489,30 @@ LogicalResult CudaBackend::makeLLIR(MLIRContext &context, ModuleOp module) {
474489 // CUDABackend.instrumentation.patch("llvmir_to_llvm", pm, mod.context)
475490 }
476491
477- return pm.run (op);
492+ auto result = pm.run (op);
493+ if (succeeded (result)) {
494+ // Read resource metadata written by the allocation passes as MLIR module
495+ // attributes. These are later appended to the PTX as comments so Rust can
496+ // recover them without a separate FFI channel.
497+ auto getInt = [&](llvm::StringRef key, int32_t def = 0 ) -> int32_t {
498+ auto attr = op->getAttrOfType <mlir::IntegerAttr>(key);
499+ return attr ? static_cast <int32_t >(attr.getInt ()) : def;
500+ };
478501
479- // # LLVM-IR (MLIR) -> LLVM-IR (LLVM)
480- // llvm.init_targets()
481- // context = llvm.context()
482- // if knobs.compilation.enable_asan:
483- // raise RuntimeError(
484- // "Address Sanitizer Error: Address sanitizer is currently only
485- // supported on the AMD backend")
486- // llvm_mod = llvm.to_module(mod, context)
487- // proc = sm_arch_from_capability(capability)
488- // features = get_features(options, self.target.arch)
489- // triple = 'nvptx64-nvidia-cuda'
490- // nvidia.set_short_ptr()
491- // llvm.attach_datalayout(llvm_mod, triple, proc, features)
492- // nvidia.set_nvvm_reflect_ftz(llvm_mod)
493-
494- // if options.extern_libs and nvidia.has_extern_deps(llvm_mod):
495- // paths = [path for (name, path) in options.extern_libs]
496- // llvm.link_extern_libs(llvm_mod, paths)
497-
498- // llvm.optimize_module(llvm_mod, llvm.OPTIMIZE_O3)
499-
500- // # Get some metadata
501- // # warp-specialization mutates num_warps
502- // total_num_warps = src.get_int_attr("ttg.total-num-warps")
503- // if total_num_warps is not None:
504- // metadata["num_warps"] = total_num_warps
505- // metadata["shared"] = src.get_int_attr("ttg.shared")
506- // metadata["tmem_size"] = src.get_int_attr("ttg.tensor_memory_size")
507- // metadata["global_scratch_size"] =
508- // src.get_int_attr("ttg.global_scratch_memory_size")
509- // metadata["global_scratch_align"] =
510- // src.get_int_attr("ttg.global_scratch_memory_alignment")
511- // metadata["profile_scratch_size"] =
512- // src.get_int_attr("ttg.profile_scratch_memory_size") or 0
513- // metadata["profile_scratch_align"] =
514- // src.get_int_attr("ttg.profile_scratch_memory_alignment") or 1 ret =
515- // str(llvm_mod) del llvm_mod del context return ret
502+ // Warp-specialization may have mutated the warp count, so prefer the
503+ // post-pipeline attribute over the original options value.
504+ auto totalWarps = op->getAttrOfType <mlir::IntegerAttr>(" ttg.total-num-warps" );
505+ m_metadata.num_warps = totalWarps ? static_cast <int32_t >(totalWarps.getInt ())
506+ : m_options.num_warps ;
507+ m_metadata.num_ctas = m_options.num_ctas ;
508+ m_metadata.shared = getInt (" ttg.shared" );
509+ m_metadata.tmem_size = getInt (" ttg.tensor_memory_size" );
510+ m_metadata.global_scratch_size = getInt (" ttg.global_scratch_memory_size" );
511+ m_metadata.global_scratch_align = getInt (" ttg.global_scratch_memory_alignment" , 1 );
512+ m_metadata.profile_scratch_size = getInt (" ttg.profile_scratch_memory_size" );
513+ m_metadata.profile_scratch_align = getInt (" ttg.profile_scratch_memory_alignment" , 1 );
514+ }
515+ return result;
516516}
517517
518518std::unique_ptr<mlir::Pass>
0 commit comments