#include using namespace metal; // GPU deswizzle for RDNA2 tiled surfaces at 4/8/16 bytes/element — the MSL twin of // SpirvFixedShaders.CreateDetileCompute and a direct mirror of // GnmTiling.GetDetileParams. Handles both supported equation families and both // plain 2D textures and array textures (one grid-Z layer per slice; the caller // packs the tiled slices contiguously). width/height are ELEMENT dims (for // block-compressed formats a 4x4 block is one element); each element spans // uintsPerElement = bpp/4 words, and the X grid is widened by that factor so each // thread copies one word (elemX = gidX / upe, word = gidX % upe): // inBlock = equation == 1 // BlockTable, modes 1/4/8 // ? blockTable[(y % blockHeight) * blockWidth + (elemX % blockWidth)] // : xTerm[elemX & xMask] ^ yTerm[y & yMask]; // ExactXor 5/9/24/27 // srcElem = z * srcSliceElements // + (y / blockHeight * blocksPerRow + elemX / blockWidth) * blockElements // + inBlock; // dstElem = z * width * height + y * width + elemX; // out[dstElem * upe + word] = tiled[srcElem * upe + word]; // Term tables hold ELEMENT offsets. buffer(1) carries xTerm (ExactXor) OR the // block table (BlockTable); the two equations index different-sized buffers, so // exactly one branch runs. 1/2 bpp are sub-word and stay on the CPU. struct DetileParams { uint width; uint height; uint blockWidth; uint blockHeight; uint blockElements; uint blocksPerRow; uint xMask; uint yMask; uint srcSliceElements; uint equation; uint uintsPerElement; }; kernel void detile_cs( device const uint* tiled [[buffer(0)]], device const uint* xTermOrTable [[buffer(1)]], device const uint* yTerm [[buffer(2)]], device uint* outLinear [[buffer(3)]], constant DetileParams& params [[buffer(4)]], uint3 gid [[thread_position_in_grid]]) { uint elemX = gid.x / params.uintsPerElement; uint word = gid.x - elemX * params.uintsPerElement; if (elemX >= params.width || gid.y >= params.height) { return; } uint blockIndex = (gid.y / params.blockHeight) * params.blocksPerRow + (elemX / params.blockWidth); uint off; if (params.equation != 0u) { uint inX = elemX - (elemX / params.blockWidth) * params.blockWidth; uint inY = gid.y - (gid.y / params.blockHeight) * params.blockHeight; off = xTermOrTable[inY * params.blockWidth + inX]; } else { off = xTermOrTable[elemX & params.xMask] ^ yTerm[gid.y & params.yMask]; } uint srcElem = gid.z * params.srcSliceElements + blockIndex * params.blockElements + off; uint dstElem = gid.z * params.width * params.height + gid.y * params.width + elemX; outLinear[dstElem * params.uintsPerElement + word] = tiled[srcElem * params.uintsPerElement + word]; }