ó
    �~id  ã                   óŠ  • S SK r S SKrS SKrS SKrS SKrS SKJr  S SKJr  S SKJ	r	  S SK
Jr  S SKJr  S SKJr  \R                   R#                  \R                   R%                  \5      5      r\R                   R)                  \S5      /r\R                   R)                  \S	5      rS
/rSq\ R2                  " 5       S 5       r\ R2                  " 5       S 5       r " S S\5      rS rSSSSSS.rSSSSSS.r Sr!\"" \!5      r#S r$\%" S \&" S5       5       5      r'S\'S'   S\'S'   S\'S'   S r(S  r) " S! S"\5      r* " S# S$\5      r+g)%é    N)ÚPath)Úknobs)Úcompile_module_from_src)Ú_allocation)Ú	GPUTarget)Ú	GPUDriverÚincludeÚlibúlibcuda.so.1c            	      ó(  • [         R                  R                  =n (       a  U /$ [        R                  " SS/5      R                  SS9nUR                  5        Vs/ s H  nSU;   d  M  UR                  5       S   PM      nnU Vs/ s H"  n[        R                  R                  U5      PM$     nn[        R                  " S5      nU(       al  U(       de  UR                  S5       Vs/ s HI  n[        R                  R                  [        R                  R                  US5      5      (       d  MG  UPMK     nnS	nU(       a  US
[        U5      -  -  nUS-  nO
US-  nUS-  n[        S U 5       5      (       d   U5       eU$ s  snf s  snf s  snf )Nz/sbin/ldconfigz-pÚignore)Úerrorsr   éÿÿÿÿÚLD_LIBRARY_PATHÚ:zlibcuda.so cannot found!
z!Possible files are located at %s.z:Please create a symlink of libcuda.so to any of the files.z<Please make sure GPU is set up and then run "/sbin/ldconfig"z- (requires sudo) to refresh the linker cache.c              3   óœ   #   • U  HB  n[         R                  R                  [         R                  R                  US 5      5      v •  MD     g7f)r   N)ÚosÚpathÚexistsÚjoin)Ú.0r   s     ÚZ/home/mande/repo/quber/.venv/lib/python3.13/site-packages/triton/backends/nvidia/driver.pyÚ	<genexpr>Úlibcuda_dirs.<locals>.<genexpr>(   s/   é € ÐSÊdÀdŒr�w‰w�~‰~œbŸg™gŸl™l¨4°Ó@×AÐAÊdùs   ‚A
A)r   ÚnvidiaÚlibcuda_pathÚ
subprocessÚcheck_outputÚdecodeÚ
splitlinesÚsplitr   r   ÚdirnameÚgetenvr   r   ÚstrÚany)	Úenv_libcuda_pathÚlibsÚlineÚlocsÚlocÚdirsÚenv_ld_library_pathÚdirÚmsgs	            r   Úlibcuda_dirsr/      sd  € ä Ÿ<™<×4Ñ4Ð4ÐÕ4Ø Ð!Ð!ä×"Ò"Ð$4°dÐ#;Ó<×CÑCÈ8ÐCÐT€Dð *.¯©Ô):ÓUÒ): ¸nÐPTÑ>TÓˆD�J‰J‹L˜ÔÑ):€DÐUÙ,0Ó1ªD SŒB�G‰G�O‰O˜CÖ ©D€DÐ1ÜŸ)š)Ð$5Ó6ÐÞ¦4Ø2×8Ñ8¸Ô=ÓsÒ=˜ÄÇÁÇÁÔPR×PWÑPW×P\ÑP\Ð]`ÐbpÓPq×Ar—Ñ=ˆÐsØ
&€CÞØÐ2´S¸³YÑ>Ñ>ˆØÐKÑK‰àÐMÑMˆØÐ>Ñ>ˆÜÑSÉdÓS×SÑSÐXÐUXÓXÐSØ€Kùò VùÚ1ùò ts   Á
FÁ*FÂ)F
Ã)AFÄ3Fc                  ó$   • [         /[        5       Q$ ©N)Úlibdevice_dirr/   © ó    r   Úlibrary_dirsr5   ,   s   € äÐ+œL›NÐ+Ð+r4   c                   ó.   ^ • \ rS rSrU 4S jrS rSrU =r$ )Ú	CudaUtilsé6   c                 ón   >• [        U S5      (       d  [        [        U ]  U 5      U l        U R                  $ )NÚinstance)ÚhasattrÚsuperr7   Ú__new__r:   )ÚclsÚ	__class__s    €r   r=   ÚCudaUtils.__new__8   s-   ø€ Ü�s˜J×'Ñ'Ü ¤¨CÑ8¸Ó=ˆCŒLØ�|‰|Ðr4   c                 ór  • [        [        [        R                  R	                  [
        S5      5      R                  5       S[        5       [        [        S9nUR                  q
UR                  U l        UR                  U l        UR                  U l        UR                  U l        UR                  U l        g )Nzdriver.cÚ
cuda_utils©ÚsrcÚnamer5   Úinclude_dirsÚ	libraries)r   r   r   r   r   r"   Ú	read_textr5   rF   rG   ÚPyCUtensorMapÚload_binaryÚget_device_propertiesÚcuOccupancyMaxActiveClustersÚset_printf_fifo_sizeÚfill_tma_descriptor)ÚselfÚmods     r   Ú__init__ÚCudaUtils.__init__=   s‹   € Ü%Ü”R—W‘W—\‘\¤'¨:Ó6Ó7×AÑAÓCØÜ%›Ü%Üñ
ˆð ×)Ñ)ˆØŸ?™?ˆÔØ%(×%>Ñ%>ˆÔ"Ø,/×,LÑ,LˆÔ)Ø$'×$<Ñ$<ˆÔ!Ø#&×#:Ñ#:ˆÕ r4   )rL   rN   rK   rJ   rM   )Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__r=   rQ   Ú__static_attributes__Ú__classcell__©r?   s   @r   r7   r7   6   s   ø† õ÷
;ð ;r4   r7   c                 ó®   • U S   S:X  a  gU R                  S5      (       a  g0 SS_SS_S	S
_SS_SS_SS_SS_SS_SS_SS_SS_SS_SS_SS_SS_SS_U    $ )Nr   Ú*ÚCUdeviceptrÚ
tensordescÚCUtensorMapÚi1Úint8_tÚi8Úi16Úint16_tÚi32Úint32_tÚi64Úint64_tÚu1Úuint8_tÚu8Úu16Úuint16_tÚu32Úuint32_tÚu64Úuint64_tÚfp16ÚdoubleÚbf16Úfp32Úf32Úfp64Ú	nvTmaDesc)Ú
startswith)Útys    r   Ú	ty_to_cpprz   S   sò   € Ø	ˆ!�u�ƒ|ØØ	‡}�}�\×"Ñ"ØðØˆhðàˆhðð 	ˆyðð 	ˆyð	ð
 	ˆyðð 	ˆiðð 	ˆiðð 	ˆzðð 	ˆzðð 	ˆzðð 	�ðð 	�ðð 	�ðð 	ˆxðð 	�ðð  	�]ð!ð" 	ñ#
ð 
r4   rl   rn   rp   )rq   rs   rt   ru   rv   Ú	pack_fp16Ú	pack_bf16Ú	pack_fp32Ú	pack_fp64ÚiiiKKppOOOOOOc                 ó|  ^^^^• U4S jnU4S jmU4S jmU4S jmU" UR                  5       5      n[        U5       VVs0 s H  u  pVXV_M	     nnnSR                  UR                  5        Vs/ s H  nT" U5      PM     sn5      n[        U-   n	/ n
UR                  5        H  nT" Xº5        M     [        U
5       VVs0 s H  u  pVXV_M	     nnn[	        U5      S:”  a)  SSR                  S UR                  5        5       5      -   OSn/ nUR                  5        HU  u  pWUS	:X  a  M  U[        ;   a  UR                  [        U    S
U 35        M6  UR                  [        U5       S
U 35        MW     SR                  U5      n/ nUR                  5        H~  u  pWUS   S:X  a  UR                  SU S35        M%  U[        ;   a  UR                  SU S35        MF  US:X  a  UR                  SU 35        Mb  US	:w  d  Mj  UR                  SU 35        M€     [        [	        U5      5      nSnUR                  5        VVs/ s H  u  pWUS   S:X  d  M  SU SU SU SU S3	PM!     nnnUR                  5        VVs/ s H  u  pWUS:X  d  M  SU SU SU S3PM     nnnUR                  5        VVs/ s H-  u  pWU[        ;   d  M  [        U    SU S[        U    SU S3PM/     nnnUR                  5        VVs/ s H  u  pWUS	:w  d  M  SU 3PM     nnnUR                  S 5        UR                  S!5        S"[	        U5      S:”  a  SU-   OS S#SR                  U5       S$UR                  UR                  5        VVs/ s H  u  pWT" U5       SU S%3PM     snn5       S&U	 S'U S(UR                  U5       SUR                  U5       SUR                  U5       S)[	        U5      S:”  a  SSR                  U5      -   OS S*3nU$ s  snnf s  snf s  snnf s  snnf s  snnf s  snnf s  snnf s  snnf )+Nc                 óä  >• / nSnU  GHK  n[        U[        5      (       Ga   UR                  S5      (       Ga	  T
(       a  T
U   OS nUS-  n[        R                  " SU5      nUR                  S5      nUR                  S5      nUR                  S5      S-   nUcL  UR                  SU-   5        [        SU-  5       H  n	UR                  S5        M     UR                  S	5        OUR                  S
5        [        U5       H  n	UR                  S5        M     [        U5       H  n	UR                  S5        M     GM:  UR                  U5        GMN     T
(       a  U[        T
5      :X  d   eU$ )Nr   r]   é   ztensordesc<([^[>]*)\[([^]]*)\]é   Ú,r[   rf   r_   rw   rd   )
Ú
isinstancer$   rx   ÚreÚmatchÚgroupÚcountÚappendÚrangeÚlen)Ú	signatureÚoutputÚtensordesc_idxÚsigÚmetar‡   ÚdtypeÚshapeÚndimÚ_Útensordesc_metas             €r   Ú_expand_signatureÚ(make_launcher.<locals>._expand_signature�   s5  ø€ ØˆØˆô ˆCÜ˜#œs×#Ò#¨¯©°|×(DÒ(DÞ:I� ~Ò6Èt�Ø !Ñ#�äŸšÐ!CÀSÓI�ØŸ™ A›�ØŸ™ A›�Ø—{‘{ 3Ó'¨!Ñ+�à‘<Ø—M‘M #¨¡+Ô.ô # 1 t¡8ž_˜ØŸ™ eÖ,ñ -à—M‘M $Õ'à—M‘M +Ô.ä˜tž�AØ—M‘M %Ö(ñ %ä˜tž�AØ—M‘M %Ö(ô %ð —‘˜c×"ñ9 ö< # n¼¸OÓ8LÓ&LÐLÐLØˆr4   c                 óv   >• [        U [        5      (       a  U  H  nT" X!5        M     g UR                  U 5        g r1   )r…   ÚtuplerŠ   )r�   rŽ   ÚxÚ_flatten_signatures      €r   rœ   Ú)make_launcher.<locals>._flatten_signature¨   s0   ø€ ä�cœ5×!Ñ!Û�Ù" 1Ö-ò ð �M‰M˜#Õr4   c                 ó¨   >• [        U [        5      (       a!  SR                  [        TU 5      5      nSU S3$ U S   S:X  a  gU S;   a  g[	        U 5      $ )Nr„   Ú[Ú]r   r[   z	PyObject*©Ú	constexprrw   )r…   rš   r   Úmaprz   )ry   ÚvalÚ_extracted_types     €r   r¥   Ú&make_launcher.<locals>._extracted_type°   sW   ø€ Ü�bœ%× Ñ Ø—(‘(œ3˜°Ó3Ó4ˆCØ�s�e˜1�:ÐØˆa‰5�C‹<ØØÐ+Ó+ØÜ˜‹}Ðr4   c                 óò   >• [        U [        5      (       a!  SR                  [        TU 5      5      nSU S3$ U S   S:X  a  gU S;   a  gU R	                  S5      (       a  gS	S
SSSSSSSSS.
[        U 5         $ )NÚ Ú(Ú)r   r[   ÚOr¡   r]   ÚdÚlÚbÚhÚiÚLÚBÚHÚIÚK)
rr   Úlongr`   rc   re   rg   ri   rl   rn   rp   )r…   rš   r   r£   rx   rz   )ry   r¤   Ú	format_ofs     €r   r·   Ú make_launcher.<locals>.format_ofº   s•   ø€ Ü�bœ%× Ñ Ø—'‘'œ#˜i¨Ó,Ó-ˆCØ�s�e˜1�:ÐØˆa‰5�C‹<ØØÐ+Ó+ØØ�=‰=˜×&Ñ&ØàØØØØØØØØØñ
ô �B‹-ñð 	r4   r¨   r   z, c              3   ó0   #   • U  H  u  pS U 3v •  M     g7f)z&_argNr3   )r   r°   ry   s      r   r   Ú make_launcher.<locals>.<genexpr>Û   s   é € Ð LÒ:K±° 5¨¨¥Ò:Kùs   ‚r¢   z argr[   Úptr_infoz.dev_ptrÚ_argÚ_storagerw   z*tma_ptrz
  zDevicePtrInfo ptr_infoz = getPointer(_argz); if (!ptr_infoz.valid) return NULL;zCUtensorMap* tma_ptrz = getTmaDesc(_argz); if (!tma_ptrz) return NULL;z _argz_storage = z(_argz);z&argz&global_scratchz&profile_scratchaÊ  
#include "cuda.h"
#include <dlfcn.h>
#include <stdbool.h>
#include <stdlib.h>
#define PY_SSIZE_T_CLEAN
#include <Python.h>

typedef struct {
  PyObject_HEAD;
  _Alignas(128) CUtensorMap tensorMap;
} PyCUtensorMapObject;

static inline void gpuAssert(CUresult code, const char *file, int line)
{
   if (code != CUDA_SUCCESS)
   {
      const char* prefix = "Triton Error [CUDA]: ";
      const char* str;
      cuGetErrorString(code, &str);
      char err[1024] = {0};
      strcat(err, prefix);
      strcat(err, str);
      PyGILState_STATE gil_state;
      gil_state = PyGILState_Ensure();
      PyErr_SetString(PyExc_RuntimeError, err);
      PyGILState_Release(gil_state);
   }
}

#define CUDA_CHECK(ans) { gpuAssert((ans), __FILE__, __LINE__); }

typedef CUresult (*cuLaunchKernelEx_t)(const CUlaunchConfig* config, CUfunction f, void** kernelParams, void** extra);

static cuLaunchKernelEx_t getLaunchKernelExHandle() {
  // Open the shared library
  void* handle = dlopen("libcuda.so.1", RTLD_LAZY);
  if (!handle) {
    PyErr_SetString(PyExc_RuntimeError, "Failed to open libcuda.so.1");
    return NULL;
  }
  // Clear any existing error
  dlerror();
  cuLaunchKernelEx_t cuLaunchKernelExHandle = (cuLaunchKernelEx_t)dlsym(handle, "cuLaunchKernelEx");
  // Check for errors
  const char *dlsym_error = dlerror();
  if (dlsym_error) {
    PyErr_SetString(PyExc_RuntimeError, "Failed to retrieve cuLaunchKernelEx from libcuda.so.1");
    return NULL;
  }
  return cuLaunchKernelExHandle;
}

static void _launch(int gridX, int gridY, int gridZ, int num_warps, int num_ctas, int launch_cooperative_grid, int launch_pdl, int shared_memory, CUstream stream, CUfunction function, CUdeviceptr global_scratch, CUdeviceptr profile_scratchz) {
  void *params[] = { au   };
  if (gridX*gridY*gridZ > 0) {
    // 4 attributes that we can currently pass maximum
    CUlaunchAttribute launchAttr[4];
    static cuLaunchKernelEx_t cuLaunchKernelExHandle = NULL;
    if (cuLaunchKernelExHandle == NULL) {
      cuLaunchKernelExHandle = getLaunchKernelExHandle();
    }
    CUlaunchConfig config;
    config.gridDimX = gridX * num_ctas;
    config.gridDimY = gridY;
    config.gridDimZ = gridZ;

    config.blockDimX = 32 * num_warps;
    config.blockDimY = 1;
    config.blockDimZ = 1;
    config.sharedMemBytes = shared_memory;
    config.hStream = stream;
    config.attrs = launchAttr;
    int num_attrs = 0;

    if (launch_pdl != 0) {
      CUlaunchAttribute pdlAttr = { .id = CU_LAUNCH_ATTRIBUTE_PROGRAMMATIC_STREAM_SERIALIZATION, .value = 1};
      launchAttr[num_attrs] = pdlAttr;
      ++num_attrs;
    }

    if (launch_cooperative_grid != 0) {
      CUlaunchAttribute coopAttr = { .id = CU_LAUNCH_ATTRIBUTE_COOPERATIVE, .value = 1};
      launchAttr[num_attrs] = coopAttr;
      ++num_attrs;
    }

    if (num_ctas != 1) {
      CUlaunchAttribute clusterAttr = {};
      clusterAttr.id = CU_LAUNCH_ATTRIBUTE_CLUSTER_DIMENSION;
      clusterAttr.value.clusterDim.x = num_ctas;
      clusterAttr.value.clusterDim.y = 1;
      clusterAttr.value.clusterDim.z = 1;
      launchAttr[num_attrs] = clusterAttr;
      ++num_attrs;

      CUlaunchAttribute clusterSchedulingAttr = {};
      clusterSchedulingAttr.id = CU_LAUNCH_ATTRIBUTE_CLUSTER_SCHEDULING_POLICY_PREFERENCE;
      clusterSchedulingAttr.value.clusterSchedulingPolicyPreference = CU_CLUSTER_SCHEDULING_POLICY_SPREAD;
      launchAttr[num_attrs] = clusterSchedulingAttr;
      ++num_attrs;
    }

    // num_ctas == 16 is non-portable. Does work for H100 and B200 tho
    config.numAttrs = num_attrs;
    if (num_ctas == 16) {
      CUDA_CHECK(cuFuncSetAttribute(
          function,
          CU_FUNC_ATTRIBUTE_NON_PORTABLE_CLUSTER_SIZE_ALLOWED,
          1
      ));
    }

    CUDA_CHECK(cuLaunchKernelExHandle(&config, function, params, 0));
  }
}

typedef struct _DevicePtrInfo {
    CUdeviceptr dev_ptr;
    bool valid;
} DevicePtrInfo;

static PyObject* data_ptr_str = NULL;
static PyObject* py_tensor_map_type = NULL;

static inline DevicePtrInfo getPointer(PyObject *obj, int idx) {
  DevicePtrInfo ptr_info;
  ptr_info.dev_ptr = 0;
  ptr_info.valid = true;
  if (PyLong_Check(obj)) {
    ptr_info.dev_ptr = PyLong_AsUnsignedLongLong(obj);
    return ptr_info;
  }
  if (obj == Py_None) {
    // valid nullptr
    return ptr_info;
  }
  PyObject *ret = PyObject_CallMethodNoArgs(obj, data_ptr_str);
  if (!ret) {
    PyErr_SetString(PyExc_TypeError, "Pointer argument must be either uint64 or have data_ptr method");
    ptr_info.valid = false;
    goto cleanup;
  }
  if (!PyLong_Check(ret)) {
    PyErr_SetString(PyExc_TypeError, "data_ptr method of Pointer object must return 64-bit int");
    ptr_info.valid = false;
    goto cleanup;
  }
  ptr_info.dev_ptr = PyLong_AsUnsignedLongLong(ret);
  if(!ptr_info.dev_ptr)
    return ptr_info;
  uint64_t dev_ptr;
  int status = cuPointerGetAttribute(&dev_ptr, CU_POINTER_ATTRIBUTE_DEVICE_POINTER, ptr_info.dev_ptr);
  if (status == CUDA_ERROR_INVALID_VALUE) {
      PyErr_Format(PyExc_ValueError,
                   "Pointer argument (at %d) cannot be accessed from Triton (cpu tensor?)", idx);
      ptr_info.valid = false;
  } else if (status != CUDA_SUCCESS) {
      CUDA_CHECK(status);  // Catch any other cuda API errors
      ptr_info.valid = false;
  }
  ptr_info.dev_ptr = dev_ptr;
cleanup:
  Py_XDECREF(ret);
  return ptr_info;

}

static inline CUtensorMap* getTmaDesc(PyObject *obj) {
  if (sizeof(CUtensorMap*) != 8) {
    PyErr_SetString(PyExc_SystemError, "getTmaDesc() requires 64-bit compilation");
    return NULL;
  }

if (Py_TYPE(obj) != (PyTypeObject*)py_tensor_map_type) {
    PyErr_Format(PyExc_TypeError, "object must be of type PyCUtensorMap, got %s", Py_TYPE(obj)->tp_name);
    return NULL;
}

  CUtensorMap* map = &((PyCUtensorMapObject*)obj)->tensorMap;
  uintptr_t align_128 = (uintptr_t)map & (128 - 1);
  if (align_128 != 0) {
    PyErr_Format(PyExc_ValueError, "CUtensorMap must be aligned to 128B, but got (&map) mod 128 = %ld", align_128);
    return NULL;
  }
  return map;
}

static void ensureCudaContext() {
  CUcontext pctx;
  CUDA_CHECK(cuCtxGetCurrent(&pctx));
  if (!pctx) {
    // Ensure device context.
    CUdevice device;
    CUDA_CHECK(cuDeviceGet(&device, 0));
    CUDA_CHECK(cuDevicePrimaryCtxRetain(&pctx, device));
    CUDA_CHECK(cuCtxSetCurrent(pctx));
  }
}

static uint16_t pack_fp16(double f) {
    uint16_t result;
    // from https://github.com/python/pythoncapi-compat
#if 0x030600B1 <= PY_VERSION_HEX && PY_VERSION_HEX <= 0x030B00A1 && !defined(PYPY_VERSION)
    _PyFloat_Pack2(f, (unsigned char*)&result, 1);
#else
    PyFloat_Pack2(f, (unsigned char*)&result, 1);
#endif
    return result;
}

static uint16_t pack_bf16(double f) {
    float f32 = (float)f;
    uint32_t u32 = *(uint32_t*)&f32;
    return (uint16_t)(u32 >> 16);
}

static uint32_t pack_fp32(double f) {
    float f32 = (float)f;
    return *(uint32_t*)&f32;
}

static uint64_t pack_fp64(double f) {
    return *(uint64_t*)&f;
}

static PyObject* launch(PyObject* self, PyObject* args) {
  // ensure cuda context is valid before calling any CUDA APIs, e.g. before getPointer calls cuPointerGetAttributes
  ensureCudaContext();

  int gridX, gridY, gridZ;
  uint64_t _stream;
  uint64_t _function;
  int launch_cooperative_grid;
  int launch_pdl;
  PyObject *launch_enter_hook = NULL;
  PyObject *launch_exit_hook = NULL;
  PyObject *kernel_metadata = NULL;
  PyObject *launch_metadata = NULL;
  PyObject *global_scratch_obj = NULL;
  PyObject *profile_scratch_obj = NULL;
  Ú;z
  if(!PyArg_ParseTuple(args, "aM  ", &gridX, &gridY, &gridZ,
                                           &_stream, &_function, &launch_cooperative_grid, &launch_pdl, &global_scratch_obj, &profile_scratch_obj,
                                           &kernel_metadata, &launch_metadata,
                                           &launch_enter_hook, &launch_exit_hooka   )) {
    return NULL;
  }

  int num_warps, num_ctas, shared_memory;
  if (!PyArg_ParseTuple(kernel_metadata, "iii", &num_warps, &num_ctas, &shared_memory)) {
    PyErr_SetString(PyExc_TypeError, "kernel_metadata must be a tuple");
    return NULL;
  }

  // extract launch metadata
  if (launch_enter_hook != Py_None){
    PyObject* ret = PyObject_CallOneArg(launch_enter_hook, launch_metadata);
    if (!ret)
      return NULL;
    Py_DECREF(ret);
  }

  CUdeviceptr global_scratch = 0;
  if (global_scratch_obj != Py_None) {
    DevicePtrInfo global_scratch_info = getPointer(global_scratch_obj, -1);
    if (!global_scratch_info.valid) {
      return NULL;
    }
    global_scratch = global_scratch_info.dev_ptr;
  }

  CUdeviceptr profile_scratch = 0;
  if (profile_scratch_obj != Py_None) {
    DevicePtrInfo profile_scratch_info = getPointer(profile_scratch_obj, -1);
    if (!profile_scratch_info.valid) {
      return NULL;
    }
    profile_scratch = profile_scratch_info.dev_ptr;
  }

  // raise exception asap
  zÌ
  Py_BEGIN_ALLOW_THREADS;
  _launch(gridX, gridY, gridZ, num_warps, num_ctas, launch_cooperative_grid, launch_pdl, shared_memory, (CUstream)_stream, (CUfunction)_function, global_scratch, profile_scratchap  );
  Py_END_ALLOW_THREADS;
  if (PyErr_Occurred()) {
    return NULL;
  }

  if(launch_exit_hook != Py_None){
    PyObject* ret = PyObject_CallOneArg(launch_exit_hook, launch_metadata);
    if (!ret)
      return NULL;
    Py_DECREF(ret);
  }

  Py_RETURN_NONE;
}

static PyMethodDef ModuleMethods[] = {
  {"launch", launch, METH_VARARGS, "Entry point for all kernels with this signature"},
  {NULL, NULL, 0, NULL} // sentinel
};

static struct PyModuleDef ModuleDef = {
  PyModuleDef_HEAD_INIT,
  "__triton_launcher",
  NULL, //documentation
  -1, //size
  ModuleMethods
};

PyMODINIT_FUNC PyInit___triton_launcher(void) {
  data_ptr_str = PyUnicode_InternFromString("data_ptr");
  if(data_ptr_str == NULL) {
    return NULL;
  }
  PyObject* driver_mod = PyImport_ImportModule("triton.backends.nvidia.driver");
  if (driver_mod == NULL) {
    return NULL;
  }
  py_tensor_map_type = PyObject_GetAttrString(driver_mod, "PyCUtensorMap");
  if (py_tensor_map_type == NULL) {
    return NULL;
  }

  PyObject *m = PyModule_Create(&ModuleDef);
  if(m == NULL) {
    return NULL;
  }
  PyModule_AddFunctions(m, ModuleMethods);
  return m;
}
)ÚvaluesÚ	enumerater   Ú_BASE_ARGS_FORMATrŒ   ÚitemsÚFLOAT_STORAGE_TYPErŠ   rz   r‹   ÚFLOAT_PACK_FUNCTION)Ú	constantsr�   r–   r—   Úexpand_signaturer°   Úsry   Úargs_formatÚformatÚflat_signaturer�   Ú	args_listÚarg_decl_listÚ	arg_declsÚinternal_args_listÚparamsÚnewlineÚ	ptr_declsÚ	tma_declsÚfloat_storage_declsrD   r¥   rœ   r·   s     `                   @@@r   Úmake_launcherrÔ      s  û€ õ%õNõõñ. )¨×)9Ñ)9Ó);Ó<ÐÜ"+Ð,<Ô"=Ô>Ò"=™$˜!�’Ñ"=€IÑ>à—'‘'°9×3CÑ3CÔ3EÓFÒ3E¨R™9 Rž=Ñ3EÑFÓG€KÜ Ñ,€Fà€NØ×ÑÖ!ˆÙ˜3Ö/ñ "ä"+¨NÔ";Ô<Ò";™$˜!�’Ñ";€IÑ<ÜPSÐT]ÓP^ÐabÓPb��t—y‘yÑ L¸)¿/¹/Ô:KÓ LÓLÒLÐhj€Ið €MØ—‘Ö"‰ˆØ�ÓÙØÔ#Ó#Ø× Ñ Ô$6°rÑ$:Ð#;¸4À¸sÐ!CÖDà× Ñ ¤I¨b£M ?°$°q°cÐ!:Ö;ñ #ð —	‘	˜-Ó(€IØÐØ—‘Ö"‰ˆØˆa‰5�C‹<Ø×%Ñ%¨°°°8Ð&<Ö=ØÔ%Ó%Ø×%Ñ%¨¨Q¨C¨xÐ&8Ö9Ø�;Óà×%Ñ%¨°° nÖ5Ø�;ÕØ×%Ñ%¨¨Q¨C jÖ1ñ #ô ”3�y“>Ó"€Fð €Gð —_‘_Ô&ôâ&‰EˆAØˆa‰5�C‰<ó 	fÐ
   Ð#5°a°S¸¸1¸#Ð=MÈaÈSÐPdÓeÙ&ð ñ ð fo×etÑetÔevôÚevÑ\aÐ\]Ø�Ñó 	XÐ
˜q˜cÐ!3°A°3°oÀaÀSÈÓWÑevð ñ ð —_‘_Ô&ôâ&‰EˆAØÔ#Ñ#ó 	ZÔ˜bÑ!Ð
" %¨ s¨+Ô6IÈ"Ñ6MÐ5NÈeÐTUÐSVÐVXÓYÙ&ð ñ ð
 '0§o¡oÔ&7ÔMÒ&7™U˜Q¸2ÀÑ;L‹j��Q�C‹jÑ&7€FÑMØ
‡M�MÐ#Ô$Ø
‡M�MÐ$Ô%ð5pôj EHð  IRó  ESð  VWó  EWð  quð  xAò  qAð  ]_ð  p`ð `Ø—y‘y Ó(Ð)ð {*ðv ‡<�<À	ÇÁÔ@QÔRÒ@Q±u°q‘O BÓ'Ð(¨¨a¨S°Ó2Ñ@QÒRÓSÐTð U Ø &˜xð (Qð R[ÐP[ð %\ðJ ‡<�<�	ÓÐð Ø
‡<�<�	ÓÐð Ø
‡<�<Ð#Ó$Ð%ð &rô [^ð  _qó  [rð  uvó  [vð  swð  z~÷  zCñ  zCð  DVó  zWò  sWð  |~ð  rð 2ð}P€Cðb
 €JùóM ?ùâFùó =ùó8ùó
ùóùó
 Nùóh SsH   ÁPÁ/PÃ PÉP ÉP Ê P&ÊP&Ê6P,Ë
 P,Ì P2Ì	P2ÎP8c              #   ó(   #   • U  H  oU4v •  M
     g 7fr1   r3   )r   r°   s     r   r   r   \  s   é € Ð:²	¨1 A¥²	ùs   ‚é   é
   é   é	   c           
      óR  • UcL  U R                   /U R                  QU R                  QU R                  S:H  PU R                  QU R                  Q$ US   nUS   nUS   nUS   nUS   nU R                  nU R                  nUS   S:X  d   eU R                  S:X  a  SOS	n	U(       a  [	        U5      nUS==   S
-  ss'   [
        R                  R                  R                  R                  R                  U R                   R                  5       UU[        U   UUUU	5      n
U
/UQUQ$ )NÚnanÚswizzleÚ	elem_sizeÚ	elem_typeÚ
block_sizeÚ
fp4_paddedr   r‚   r   rƒ   )Úbaser“   ÚstridesÚpaddingÚlistÚtritonÚruntimeÚdriverÚactiveÚutilsrN   Údata_ptrÚTMA_DTYPE_DEVICE_TO_HOST)ÚargÚmetadatarÜ   rÝ   rÞ   rß   rà   r“   râ   rã   Úcu_tensor_maps              r   Úmake_tensordesc_argrï   b  s2  € ØÑð —‘Ðc˜3Ÿ9™9Ðc s§{¡{Ðc°C·K±KÀ5Ñ4HÐcÈ3Ï9É9ÐcÐWZ×WbÑWbÐcÐcà�yÑ!€GØ˜Ñ%€IØ˜Ñ%€IØ˜,Ñ'€JØ˜,Ñ'€Jà�I‰I€EØ�k‰k€GØ�2‰;˜!ÓÐÐØ—;‘; %Ó'‰a¨Q€GæÜ�U“ˆØˆb‹	�Q‰‹	ä—N‘N×)Ñ)×0Ñ0×6Ñ6×JÑJØ�‰×ÑÓØØÜ  Ñ+ØØØØó	€Mð Ð,˜EÐ, GÐ,Ð,r4   c           
      ó´  ^ ^^• [        S UR                  5        5       5      nU(       d  T $ [        [        UR                  5       5       VVs/ s H6  u  pE[	        U[
        5      (       d  M  UR                  S5      (       d  M4  UPM8     snn5      mT(       a  [        T5      [        T5      :X  d   eT(       d  S /[        T5      -  mU UU4S jnU$ s  snnf )Nc              3   ór   #   • U  H-  n[        U[        5      =(       a    UR                  S 5      v •  M/     g7f)r]   N)r…   r$   rx   )r   r�   s     r   r   Ú)wrap_handle_tensordesc.<locals>.<genexpr>Š  s+   é € ÐrÒ_qÐX[œj¨¬cÓ2×S°s·~±~ÀlÓ7SÔSÒ_qùs   ‚57r]   c                  óä   >• [        U S [         5      nSn[        U [        S  5       HA  u  p4UT;   a%  UR                  [	        UTU   5      5        US-  nM0  UR                  U5        MC     T" U6 $ )Nr   r‚   )rä   Ú_BASE_ARGS_FORMAT_LENrÀ   Úextendrï   rŠ   )ÚargsÚ
final_argsr�   r°   rì   ÚlauncherÚtensordesc_indicesr–   s        €€€r   ÚinnerÚ%wrap_handle_tensordesc.<locals>.inner”  s~   ø€ Ü˜$Ð5Ô 5Ð6Ó7ˆ
ØˆÜ Ô%:Ð%;Ð <Ö=‰FˆAØÐ&Ó&Ø×!Ñ!Ô"5°c¸?È>Ñ;ZÓ"[Ô\Ø !Ñ#’à×!Ñ! #Ö&ñ >ñ ˜Ð$Ð$r4   )r%   r¿   ÚsetrÀ   r…   r$   rx   rŒ   )rø   r�   r–   Úhas_tensor_desc_argr°   r�   rú   rù   s   ` `    @r   Úwrap_handle_tensordescrþ   ‰  s²   ú€ ÜÑrÐ_h×_oÑ_oÔ_qÓrÓrÐÞØˆäÜ" 9×#3Ñ#3Ó#5Ô6ÔpÒ6‰vˆq¼*ÀSÌ#×:N‹ÐSV×SaÑSaÐbn×So�Ñ6ÒpórÐæ¤# oÓ"6¼#Ð>PÓ:QÓ"QÐQÐQÞØ˜&¤3Ð'9Ó#:Ñ:ˆ÷	%ð €Lùó! 	qs   ÁC
Á-C
ÂC
c                   ó    • \ rS rSrS rS rSrg)ÚCudaLauncheri¢  c                 ó¼  ^• [        TS5      (       a  TR                  O	[        5       nU4S jnUR                  5        VVs0 s H  u  pVU" U5      U_M     nnnTR                  R                  5        VVs0 s H  u  pVXV_M	     nnn[        USS 5      n[        X7U5      m[        TS[        5       [        [        S9n	[        USS5      U l        [        U	R                  Xx5      U l        UR                  U l        UR                  U l        UR                   U l        UR"                  U l        UR$                  U l        UR&                  U l        g s  snnf s  snnf )NrÅ   c                 ó~   >• [        U [        5      (       a&  TR                  R                  R	                  U 5      4$ U $ r1   )r…   r$   ÚfnÚ	arg_namesÚindex)r›   rD   s    €r   Ú<lambda>Ú'CudaLauncher.__init__.<locals>.<lambda>¦  s2   ø€ ¼ZÈÌ3×=OÑ=O˜SŸV™V×-Ñ-×3Ñ3°AÓ6Ð9ÐVÐUVÐVr4   r–   Ú__triton_launcherrC   Únum_ctasr‚   )r;   rÅ   ÚdictrÂ   r�   ÚgetattrrÔ   r   r5   rF   rG   r	  rþ   ÚlaunchÚglobal_scratch_sizeÚglobal_scratch_alignÚprofile_scratch_sizeÚprofile_scratch_alignÚlaunch_cooperative_gridÚ
launch_pdl)
rO   rD   rí   rÅ   Úarg_idxÚidxÚvaluer�   r–   rP   s
    `        r   rQ   ÚCudaLauncher.__init__¤  s&  ø€ Ü%,¨S°+×%>Ñ%>�C—M’MÄDÃFˆ	ÜVˆØ;D¿?¹?Ô;LÔMÒ;L©Z¨S‘W˜S“\ 5Ò(Ñ;Lˆ	ÑMØ25·-±-×2EÑ2EÔ2GÔHÒ2G¡J C�S’ZÑ2Gˆ	ÑHÜ! (Ð,=¸tÓDˆÜ˜I°/ÓBˆÜ%ØØ$Ü%›Ü%Üñ
ˆô   ¨*°aÓ8ˆŒÜ,¨S¯Z©Z¸ÓTˆŒØ#+×#?Ñ#?ˆÔ Ø$,×$AÑ$AˆÔ!Ø$,×$AÑ$AˆÔ!Ø%-×%CÑ%CˆÔ"Ø'/×'GÑ'GˆÔ$Ø"×-Ñ-ˆ�ùó' NùÛHs   ÁEÁ7Ec                 ó.  ^ ^^^^• UUUU U4S jnU" T R                   T R                  [        R                  5      nU" T R                  T R
                  [        R                  5      n	T R                  " TTTTUT R                  T R                  X‰/	UQ76   g )Nc                 óx   >• U S:”  a3  TT-  T-  nUT	R                   -  U -  nUR                  5       nU" XAT
5      $ g ©Nr   )r	  Úget)ÚsizeÚalignÚ	allocatorÚ	grid_sizeÚ
alloc_sizeÚalloc_fnÚgridXÚgridYÚgridZrO   Ústreams         €€€€€r   Úallocate_scratchÚ/CudaLauncher.__call__.<locals>.allocate_scratch¾  sF   ø€ Ø�a‹xØ! E™M¨EÑ1�	Ø&¨¯©Ñ6¸Ñ=�
Ø$Ÿ=™=›?�Ù 
°6Ó:Ð:Ør4   )
r  r  r   Ú
_allocatorr  r  Ú_profile_allocatorr  r  r  )
rO   r!  r"  r#  r$  Úfunctionrö   r%  Úglobal_scratchÚprofile_scratchs
   `````     r   Ú__call__ÚCudaLauncher.__call__¼  s†   ü€ ÷	ñ 	ñ *¨$×*BÑ*BÀD×D]ÑD]Ô_j×_uÑ_uÓvˆÙ*¨4×+DÑ+DÀd×F`ÑF`Ü+6×+IÑ+IóKˆà�Š�E˜5 %¨°¸4×;WÑ;WÐY]×YhÑYhØ"ð	<Ø6:ô	<r4   )r  r  r  r  r  r	  r  r  N)rS   rT   rU   rV   rQ   r,  rW   r3   r4   r   r   r   ¢  s   † ò.õ0<r4   r   c                   ón   ^ • \ rS rSrU 4S jrS rS rS r\S 5       r	S\
S\
4S	 jrS
 rS rS rSrU =r$ )Ú
CudaDriveriÍ  c                 óV   >• [        5       U l        [        U l        [        TU ]  5         g r1   )r7   ré   r   Úlauncher_clsr<   rQ   )rO   r?   s    €r   rQ   ÚCudaDriver.__init__Ï  s   ø€ Ü“[ˆŒ
Ü(ˆÔÜ‰ÑÕr4   c                 ó|   • U R                  5       nU R                  U5      nUS   S-  US   -   nSn[        SX#5      $ )Nr   r×   r‚   é    Úcuda)Úget_current_deviceÚget_device_capabilityr   )rO   ÚdeviceÚ
capabilityÚ	warp_sizes       r   Úget_current_targetÚCudaDriver.get_current_targetÔ  sI   € Ø×(Ñ(Ó*ˆØ×/Ñ/°Ó7ˆ
Ø ‘] RÑ'¨*°Q©-Ñ7ˆ
Øˆ	Ü˜ Ó7Ð7r4   c                 óJ   • SS K nUR                  SU R                  5       5      $ )Nr   r5  )Útorchr8  r6  ©rO   r>  s     r   Úget_active_torch_deviceÚ"CudaDriver.get_active_torch_deviceÛ  s   € ÛØ�|‰|˜F D×$;Ñ$;Ó$=Ó>Ð>r4   c                 ó"   • SS K nUR                  $ r  )r>  r5  r?  s     r   Úget_device_interfaceÚCudaDriver.get_device_interfaceß  s   € ÛØ�z‰zÐr4   c                  óž   •  SS K n U R                  R                  5       =(       a    U R                  R                  S L $ ! [
         a     gf = f)Nr   F)r>  r5  Úis_availableÚversionÚhipÚImportError)r>  s    r   Ú	is_activeÚCudaDriver.is_activeã  sC   € ð	ÛØ—:‘:×*Ñ*Ó,×L°%·-±-×2CÑ2CÀtÐ2KÐLøÜó 	Ùð	ús   ‚<? ¿
AÁAry   Úreturnc                 ó   • [        U5      $ r1   )rz   )rO   ry   s     r   Úmap_python_to_cpp_typeÚ!CudaDriver.map_python_to_cpp_typeë  s   € Ü˜‹}Ðr4   c                 ó   • SSK Jn  U$ )Nr   )Údo_bench)Útriton.testingrQ  )rO   rQ  s     r   Úget_benchmarkerÚCudaDriver.get_benchmarkerî  s
   € Ý+Øˆr4   c                 ó\   • SS K nSnUR                  [        US-  5      UR                  SS9$ )Nr   i   é   r5  )r’   r8  )r>  ÚemptyÚint)rO   r>  Ú
cache_sizes      r   Úget_empty_cache_for_benchmarkÚ(CudaDriver.get_empty_cache_for_benchmarkò  s.   € Ûð
 'ˆ
Ø�{‰{œ3˜z¨Q™Ó/°u·y±yÈˆ{ÐPÐPr4   c                 ó$   • UR                  5         g r1   )Úzero_)rO   Úcaches     r   Úclear_cacheÚCudaDriver.clear_cacheû  s   € Ø�‰�r4   )r1  ré   )rS   rT   rU   rV   rQ   r;  r@  rC  ÚstaticmethodrJ  r$   rN  rS  rZ  r_  rW   rX   rY   s   @r   r/  r/  Í  sS   ø† õò
8ò?òð ñó ðð¨ð °ô òòQ÷ð r4   r/  ),Ú	functoolsr   r   rå   r†   Úpathlibr   r   Útriton.runtime.buildr   Útriton.runtimer   Útriton.backends.compilerr   Útriton.backends.driverr   r   r"   ÚrealpathÚ__file__r   rF   r2   rG   rI   Ú	lru_cacher/   r5   Úobjectr7   rz   rÃ   rÄ   rÁ   rŒ   rô   rÔ   r
  r‹   rë   rï   rþ   r   r/  r3   r4   r   Ú<module>rl     ss  ðÛ Û 	Û Û Û 	Ý Ý Ý 8Ý &Ý .Ý ,à
�'‰'�/‰/˜"Ÿ'™'×*Ñ*¨8Ó4Ó
5€Ø—‘—‘˜W iÓ0Ð1€Ø—‘—‘˜W eÓ,€ØÐ€	Ø€ð ×ÒÓñó ðð. ×ÒÓñ,ó ð,ô;�ô ;ò:
ð4 ØØØØñÐ ð ØØØØñÐ ð $Ð ÙÐ-Ó.Ð òYñz  Ñ:±°b´	Ó:Ó:Ð Ø Ð ˜Ñ ØÐ ˜Ñ Ø Ð ˜Ñ ò$-òNô2(<�6ô (<ôV/�õ /r4   