makepad/libs/ai/models/trellis/tests/pixal_naf_oracle.py
Admin 60615db4ed libs: ai hub/models/cuda, speech, chat_ui
Squash of 54 work commits (Sep 1–12):
  6251f7c  ai-hub: body domain — live pose packets ride the realtime session
  ea50c77  chat_ui: the feed's session gets its profile brief back
  f51b5f3  ai-body: the crate for the native SAM 3D Body port, with its weights reader
  8211ae6  ai-body: the MHR rig and the pose head's parameter decoding, oracle-exact
  9e343a8  ai-body: the DINOv3 ViT-H+/16 backbone, crop and ray conditioning; Metal gains rope-half and affine layer norm
  69d842c  ai-body: the promptable pose decoder and its refinement loop, oracle-matched on Metal
  66e5e2f  ai-hub: SAM 3D Body runs natively — `sam3dbody` on the body domain, oracle-matched end to end
  a634198  ai-hub: the body-native commit carried a peer's in-flight hub hunks; put them back where they were
  9ff44e8  ai-hub: the body-native wiring, this time only the lane's hunks
  6a1c16b  ai-body: third-party notices — what the port is implemented after, and what it is not
  d78411a  ai-body: the per-step work moves to the GPU
  b22259b  ai-body: the context stays on the GPU; only the pose token leaves the loop
  346f31f  ai-body: flash attention for the head-dim-64 blocks
  45b5b98  ai-body: the crop size is a runtime knob, and the loop reports where its time goes
  4be6d19  ai-body: the test modules import the grid constants they still use
  7598346  ai-body: tensor-core GEMMs for the backbone, and the rig's correctives only where they count
  a9ce596  ai-body: the crop warp runs across cores
  8964ba6  ai-body: an FP8 backbone mode, off by default, measured against the oracle
  a2aaa8f  ai-body: the FP8 bias rides a column-broadcast add on the device
  d53c77d  metal: a device-resident ViT stack, and the body backbone rides it
  d006d0a  metal: resident f32 linears keep their weight on the device
  525ba1c  metal: a device-resident two-way decoder layer, and the body decoder rides it
  c9e6d88  ai-body: the hands pass — hand crops, the hand decoder, the hand-mode rig and the wrist fusion
  62dff26  ai-body: the mask prompt — a person's segmentation mask conditions the body pass
  a648cf8  ai-hub: body session options — hands, detect, persons=N
  8c568df  ai-hub: drop the SAM 3D Body reference worker backend
  7ff875a  ai-hub: keep a peer's in-flight beats/notes/local work out of the body commits
  31e5faa  ai-hub: local model runner, licence acknowledgements, a shared install panel; Beat This!, Basic Pitch and the Salamander drum-kit entries
  b94bc58  ai-services: the wire, the app port and the panel state — one conversation, many apps
  2acb798  ai-services: wire v2 — endpoints, receiver-side caps, result disposition
  8ae0ffb  ai-services: the engine core — registry, router and conversation, tested against a scripted model
  2308736  ai-services: the real models behind the engine feature — local through the hub, Claude, and none
  c3f631d  livepipe: one reusable pipe from a camera to a fleet node and back
  ff62db3  ai libs: the runtime env-var cleanup — precision is a per-caller policy, not an environment side channel
  04a94ef  realtime: one service-log line when a live session opens and one when it closes
  0ecb81c  ai models: the model-crates env-var cleanup — 172 research knobs gone, the unset default is the code
  4ca36c1  ai hub + services: the assistant's model comes from wherever it is resident — the fleet chat box, with tools, then the local weights
  432121e  aichat engine + wm: launch, then use — the assistant continues in the same turn once the app it started is on the bus
  7a5bf69  ai-hub registry: the Salamander drumkit samples come from the makepad.nl mirror — the GitHub repo only carries the .sfz files
  102ffc5  ai-services: messages on the bus — a manifest declares topics, the engine subscribes on a tool's behalf or by ToolResult.subscribe, a service publishes Message frames, an idle conversation wakes on a message as an event turn under rate laws; the WM bus forwards the new frames; every app that matches the wire gets its arm
  a837792  hub + flow: a whitespace-only chat completion is retried once and then fails instead of passing as an answer; a flow's model is a fleet model id unless it names a weight file on disk; chat models show under the text domain in /v1/models
  bc6c620  hub + flow: what the chat review found — the in-process route retries an empty completion too, a node says whether its prefill opened thinking so a brief-mode answer is never discarded, a preferred model falls back to normal election when no node has it, discovery keeps looking for the preferred model until patience runs out
  75c3441  hub: the PRO 6000 serves image as well as chat and text
  ad5e98b  hub registry: flux2-dev's VRAM estimate is its measured peak, 30 GB
  c7241e0  hub: a node that evicted every resident releases its cached allocator pool before refusing a load or publishing usable VRAM
  30575f0  flow: route generation by request workload
  1be1e21  ai-hub: gate downloads by disk capacity and recover fleet admission
  df6b394  filesystem_watcher, bounded_http, ai services: live and tool prerequisites
  79ebdb9  ai-hub: add a native Pixal3D image-to-3D backend
  0ba0d74  ai-hub: propagate typed refusals under reject queue policy
  cc6c872  Speed up H3 conditioning and video decoding
  e512059  Fix Qwen vision residency and generated material colors
  2864f68  ai-hub http client: bound every plain TCP connect to 3 s per address
  3d93229  ai: CUDA is a Linux/Windows-only dependency; the hub library defaults to llm + stt

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
2026-09-15 13:40:31 +02:00

86 lines
4 KiB
Python

"""Validate the native CUDA NAF sampler against dense PyTorch operations.
Build kernels/pixal.cu with nvcc -shared -O3, then pass the resulting library.
This is an offline test oracle; production inference never imports Python.
"""
import argparse
import ctypes
import json
import time
import torch
import torch.nn.functional as F
def dense_reference(q, k, v, width, low_width, heads, kernel):
qc, vc = q.shape[-1], v.shape[-1]
qd, vd = qc // heads, vc // heads
dilation = width // low_width
q = q.reshape(width, width, heads, qd).permute(2, 0, 1, 3).reshape(heads, width*width, qd)
def unfold(x, channels):
x = x.reshape(low_width, low_width, heads, channels).permute(2, 3, 0, 1)
x = F.interpolate(x, size=(width, width), mode="nearest-exact")
x = F.unfold(x, kernel, dilation=dilation, padding=(kernel//2)*dilation)
return x.reshape(heads, channels, kernel*kernel, width*width).permute(0, 3, 2, 1)
keys = unfold(k, qd)
scores = (q.unsqueeze(2)*keys).sum(-1) * qd**-0.5
attn = scores.softmax(-1)
del keys, scores
chunks = []
# Bound the oracle's unfold workspace without changing the math.
values = v.reshape(low_width*low_width, heads, vd)
for start in range(0, vd, 16):
part = values[:, :, start:start+16].reshape(low_width*low_width, -1)
val = unfold(part, part.shape[-1]//heads)
chunks.append((attn.unsqueeze(-1)*val).sum(-2))
out = torch.cat(chunks, dim=-1).permute(1,0,2).reshape(width,width,vc)
return out.permute(2,0,1).unsqueeze(0)
def main():
parser=argparse.ArgumentParser()
parser.add_argument("library")
parser.add_argument("--benchmark", action="store_true")
args=parser.parse_args()
lib=ctypes.CDLL(args.library)
op=lib.makepad_cuda_pixal_naf_sample_f32
op.argtypes=[ctypes.c_void_p]*5+[ctypes.c_uint32]*9+[ctypes.c_void_p]
op.restype=ctypes.c_int
torch.manual_seed(19)
torch.backends.cuda.matmul.allow_tf32=False
def native(q,k,v,uv,width,low,heads,kernel):
out=torch.empty((len(uv),v.shape[-1]),device="cuda",dtype=torch.float32)
status=op(q.data_ptr(),k.data_ptr(),v.data_ptr(),uv.data_ptr(),out.data_ptr(),
len(uv),width,width,low,low,heads,q.shape[-1],v.shape[-1],kernel,
torch.cuda.current_stream().cuda_stream)
if status: raise RuntimeError(f"CUDA launch status {status}")
return out
for width,low,qc,vc,heads,kernel in [(8,2,8,16,2,3),(32,4,256,1024,4,9),(48,12,32,64,4,5)]:
q=torch.randn(width*width,qc,device="cuda")
k=torch.randn(low*low,qc,device="cuda")
v=torch.randn(low*low,vc,device="cuda")
uv=torch.rand(73,2,device="cuda")*1.4-0.2
uv[:4]=torch.tensor([[0,0],[1,1],[0.5,0.5],[-1,2]],device="cuda")
dense=dense_reference(q,k,v,width,low,heads,kernel)
expected=F.grid_sample(dense,(uv*2-1).view(1,-1,1,2),padding_mode="border",align_corners=False)
expected=expected[0,:,:,0].T.contiguous()
actual=native(q,k,v,uv,width,low,heads,kernel)
torch.testing.assert_close(actual,expected,rtol=3e-5,atol=3e-6)
print(json.dumps({"width":width,"max_error":float((actual-expected).abs().max())}),flush=True)
if args.benchmark:
for width,low,count in [(512,32,12000),(512,64,30000),(1024,64,30000)]:
q=torch.randn(width*width,256,device="cuda")
k=torch.randn(low*low,256,device="cuda")
v=torch.randn(low*low,1024,device="cuda")
uv=torch.rand(count,2,device="cuda")
for _ in range(2): actual=native(q,k,v,uv,width,low,4,9)
torch.cuda.synchronize()
start=time.perf_counter()
for _ in range(5): actual=native(q,k,v,uv,width,low,4,9)
torch.cuda.synchronize()
print(json.dumps({"width":width,"low":low,"voxels":count,
"sample_ms":(time.perf_counter()-start)*200,
"output_bytes":actual.numel()*4,"dense_output_bytes":width*width*1024*4}),flush=True)
if __name__=="__main__": main()