Skip to content
KernelIndex
Search⌘K

submission 709283

Elán Zainos Corona · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 142 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-709283?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, int32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD Instinct MI355X
170.8µs
#547 of 766
2026-04-03

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:0e98ad3c1cf4b3f92b8d0f1fb29a380375544035308a16b67fe6dcfd98e690d5
license declaredunknown
license concludedunknown
authorsElán Zainos Corona
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

persistent-kernel"""Contenedor estático para metadatos y descriptores persistentes."""

Kernel source

submission.py142 lines
# -*- coding: utf-8 -*-
"""
Fractal Core Research | Sentinel Omega Project
Versión: V24.1 (Logit Constraint Patch)
Hardware: AMD MI355X (gfx950) | Protocolo: Topology Singularity
"""

import os
import torch
import logging
import aiter
from aiter.mla import mla_decode_fwd
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1

# ─── CONFIGURACIÓN DE ENTORNO DE BAJO NIVEL ──────────────────────────────
os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
os.environ["AITER_DISABLE_TUNE"] = "1"
os.environ["TRITON_ALWAYS_COMPILE"] = "0"

logging.basicConfig(level=logging.INFO, format='[%(asctime)s] [%(levelname)s] %(message)s')
logger = logging.getLogger("SentinelCore")

_FP8_DTYPE = aiter_dtypes.fp8
_INT32 = torch.int32
_MAX_FP8 = 448.0
_LOGIT_CAP = 0.0  # Parche V24.1: Desactivado por falta de soporte en aiter.mla

class SentinelCache:
    """Contenedor estático para metadatos y descriptores persistentes."""
    meta = {}
    idx_buffer = None
    MAX_TOKENS = 4194304 

# ─── PROTOCOLO DE FORJA GLOBAL (WARMUP SÍNCRONO) ──────────────────────────
def _ignicion_global():
    if os.environ.get("LOCAL_RANK", "0") != "0":
        return
    try:
        w_dev = "cuda"
        w_qo = torch.tensor([0, 1], dtype=_INT32, device=w_dev)
        w_kv_p = torch.tensor([0, 1], dtype=_INT32, device=w_dev)
        w_l = torch.tensor([1], dtype=_INT32, device=w_dev)
        
        w_info = get_mla_metadata_info_v1(1, 1, 1, _FP8_DTYPE, _FP8_DTYPE, 
                                         is_sparse=False, fast_mode=True, 
                                         num_kv_splits=1, intra_batch_mode=True)
        w_work = [torch.empty(s, dtype=t, device=w_dev) for s, t in w_info]
        
        get_mla_metadata_v1(w_qo, w_kv_p, w_l, 1, 1, True, 
                            w_work[0], w_work[2], w_work[1], w_work[3], w_work[4], w_work[5], 
                            page_size=1, kv_granularity=16, max_seqlen_qo=1, uni_seqlen_qo=1, 
                            fast_mode=True, max_split_per_batch=1, intra_batch_mode=True, 
                            dtype_q=_FP8_DTYPE, dtype_kv=_FP8_DTYPE)
        
        mla_decode_fwd(
            torch.zeros((1, 1, 16), dtype=_FP8_DTYPE, device=w_dev),
            torch.zeros((1, 1, 1, 16), dtype=_FP8_DTYPE, device=w_dev),
            torch.zeros((1, 1, 16), dtype=torch.bfloat16, device=w_dev),
            w_qo, w_kv_p, torch.zeros(1, dtype=_INT32, device=w_dev), w_l, 1, 
            num_kv_splits=1, q_scale=torch.ones(1, device=w_dev), 
            kv_scale=torch.ones(1, device=w_dev), 
            logit_cap=_LOGIT_CAP,
            intra_batch_mode=True, 
            work_meta_data=w_work[0], work_indptr=w_work[1], work_info_set=w_work[2],
            reduce_indptr=w_work[3], reduce_final_map=w_work[4], reduce_partial_map=w_work[5]
        )
        torch.cuda.synchronize()
        torch.cuda.empty_cache()
    except Exception:
        pass

_ignicion_global()

def custom_kernel(data):
    q, kv_data, qo_ptr, kv_ptr, config = data
    bs = config["batch_size"]
    nq, nkv = config["num_heads"], config["num_kv_heads"]
    dq, dv = config["qk_head_dim"], config["v_head_dim"]
    sm_scale = config["sm_scale"]
    
    total_kv = int(kv_ptr[-1].item())
    num_kv_splits = max(1, 256 // bs)
    state_key = (bs, total_kv, num_kv_splits)
    
    if state_key not in SentinelCache.meta:
        if len(SentinelCache.meta) > 16:
            SentinelCache.meta.clear()
            
        if SentinelCache.idx_buffer is None:
            SentinelCache.idx_buffer = torch.arange(SentinelCache.MAX_TOKENS, dtype=_INT32, device=q.device)
            
        info = get_mla_metadata_info_v1(bs, 1, nq, _FP8_DTYPE, _FP8_DTYPE, 
                                         is_sparse=False, fast_mode=True, 
                                         num_kv_splits=num_kv_splits, intra_batch_mode=True)
        work = [torch.empty(s, dtype=t, device=q.device) for s, t in info]
        lens = (kv_ptr[1:] - kv_ptr[:-1]).to(_INT32)
        
        get_mla_metadata_v1(
            qo_ptr, kv_ptr, lens, nq // nkv, nkv, True, 
            work[0], work[2], work[1], work[3], work[4], work[5], 
            1, 16, 1, 1, fast_mode=True, max_split_per_batch=num_kv_splits, 
            intra_batch_mode=True, dtype_q=_FP8_DTYPE, dtype_kv=_FP8_DTYPE
        )
        
        SentinelCache.meta[state_key] = {
            "kwargs": {
                "work_meta_data": work[0], "work_indptr": work[1], "work_info_set": work[2],
                "reduce_indptr": work[3], "reduce_final_map": work[4], "reduce_partial_map": work[5]
            },
            "lens": lens
        }
    
    ctx = SentinelCache.meta[state_key]
    q_scale = (q.abs().amax() / _MAX_FP8).to(torch.float32).reshape(1)
    q_fp8 = (q / q_scale).clamp(min=-_MAX_FP8, max=_MAX_FP8).to(_FP8_DTYPE)
    
    kv_input, kv_scale = kv_data["fp8"]
    kv_buffer_4d = kv_input.view(kv_input.shape[0], 1, nkv, kv_input.shape[-1])
    out = torch.empty((bs, nq, dv), dtype=torch.bfloat16, device=q.device)
    
    mla_decode_fwd(
        q_fp8.view(-1, nq, dq),
        kv_buffer_4d,
        out,
        qo_ptr,
        kv_ptr,
        SentinelCache.idx_buffer[:total_kv],
        ctx["lens"],
        1,
        page_size=1,
        nhead_kv=nkv,
        sm_scale=sm_scale,
        logit_cap=_LOGIT_CAP,
        num_kv_splits=num_kv_splits,
        q_scale=q_scale,
        kv_scale=kv_scale,
        intra_batch_mode=True,
        **ctx["kwargs"]
    )
    
    return out
scrolls · 142 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON