Ling-3.0-tiny-RKNN / tools /mla_npu_probe.cpp
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
14.5 kB
// Standalone MLA attention experiment, not enabled in the inference engine.
// Times dynamic BF16 KV staging, DMA sync, FP16 NPU QK/AV, CPU online softmax
// and FP32 merge against the engine's four-A76 NEON attention algorithm.
#include "core_workers.h"
#include "mla_npu.h"
#include <rknn_api.h>
#include <rknn_matmul_api.h>
#include <nlohmann/json.hpp>
#include <arm_neon.h>
#include <algorithm>
#include <array>
#include <bit>
#include <chrono>
#include <cmath>
#include <cstdint>
#include <cstring>
#include <cstdlib>
#include <fstream>
#include <iostream>
#include <limits>
#include <memory>
#include <random>
#include <stdexcept>
#include <vector>
using Clock = std::chrono::steady_clock;
using Json = nlohmann::json;
constexpr int H = 16, D = 192, V = 128;
double Ms(Clock::time_point t) { return std::chrono::duration<double, std::milli>(Clock::now()-t).count(); }
void Check(int rc, const char * op) { if (rc != 0) throw std::runtime_error(std::string(op)+": "+std::to_string(rc)); }
float Float(std::uint16_t x) { return std::bit_cast<float>(std::uint32_t(x) << 16); }
std::uint16_t Bf16(float x) {
auto u = std::bit_cast<std::uint32_t>(x);
return (u + 0x7fffU + ((u >> 16) & 1U)) >> 16;
}
struct Data {
int rows, history;
std::vector<float> q;
std::vector<std::uint16_t> k, v;
Data(int m, int n): rows(m), history(n), q(m*H*D), k(std::size_t(n)*H*D), v(std::size_t(n)*H*V) {}
};
void Randomize(Data & d, int seed) {
std::mt19937 rng(seed); std::normal_distribution<float> normal(0, 1);
for (auto & x : d.q) x = normal(rng);
for (auto & x : d.k) x = Bf16(normal(rng));
for (auto & x : d.v) x = Bf16(normal(rng));
}
template<class T> void Read(std::ifstream & f, std::vector<T> & x) {
f.read(reinterpret_cast<char *>(x.data()), x.size()*sizeof(T));
if (!f) throw std::runtime_error("truncated snapshot");
}
Data Load(const char * path) {
std::ifstream f(path, std::ios::binary);
std::array<std::uint32_t, 4> header {};
f.read(reinterpret_cast<char *>(header.data()), sizeof(header));
if (!f || header[0] != 0x31414c4d || header[1] < 1 || header[1] > 128 ||
header[2] < header[1] || header[2] > 65536 || header[3] != H)
throw std::runtime_error("invalid MLA snapshot");
Data d(header[1], header[2]); Read(f,d.q); Read(f,d.k); Read(f,d.v); return d;
}
float Dot(const float * q, const std::uint16_t * k) {
auto a = vdupq_n_f32(0), b = a;
for(int i=0;i<D;i+=8) {
auto p=vld1q_u16(k+i);
a=vfmaq_f32(a,vld1q_f32(q+i),vreinterpretq_f32_u32(vshll_n_u16(vget_low_u16(p),16)));
b=vfmaq_f32(b,vld1q_f32(q+i+4),vreinterpretq_f32_u32(vshll_n_u16(vget_high_u16(p),16)));
}
return vaddvq_f32(vaddq_f32(a,b));
}
void Acc(const std::uint16_t * v, float p, float * out) {
for(int i=0;i<V;i+=8) {
auto x=vld1q_u16(v+i);
vst1q_f32(out+i,vfmaq_n_f32(vld1q_f32(out+i),vreinterpretq_f32_u32(vshll_n_u16(vget_low_u16(x),16)),p));
vst1q_f32(out+i+4,vfmaq_n_f32(vld1q_f32(out+i+4),vreinterpretq_f32_u32(vshll_n_u16(vget_high_u16(x),16)),p));
}
}
struct Timing {
double stage=0, sync=0, bind=0, qk=0, av=0, softmax=0, merge=0;
Json json() const { return {{"stage_ms",stage},{"sync_ms",sync},{"bind_dynamic_b_ms",bind},{"qk_ms",qk},{"av_ms",av},{"softmax_ms",softmax},{"merge_ms",merge}}; }
};
struct Cpu {
std::array<std::vector<float>,4> scores;
explicit Cpu(int n) { for(auto & x:scores) x.resize(n); }
double Run(const Data & d, std::vector<float> & out) {
const auto start=Clock::now(); constexpr std::array<int,4> cores {0,1,2,3};
ling3::CoreWorkers::Instance().Run(cores,[&](int c){
auto & s=scores[c];
for(int h=c;h<H;h+=4) for(int r=0;r<d.rows;++r) {
const int n=d.history-d.rows+r+1;
auto * dst=out.data()+(r*H+h)*V;
float mx=-std::numeric_limits<float>::infinity(), sum=0;
for(int t=0;t<n;++t) { s[t]=Dot(d.q.data()+(r*H+h)*D,d.k.data()+(std::size_t(t)*H+h)*D)/std::sqrt(float(D)); mx=std::max(mx,s[t]); }
for(int t=0;t<n;++t) { s[t]=std::exp(s[t]-mx); sum+=s[t]; }
std::fill(dst,dst+V,0);
for(int t=0;t<n;++t) Acc(d.v.data()+(std::size_t(t)*H+h)*V,s[t]/sum,dst);
}
}); return Ms(start);
}
};
struct Matmul {
rknn_matmul_ctx ctx=0; rknn_matmul_io_attr attr {};
rknn_tensor_mem * a=nullptr, * b=nullptr, * c=nullptr;
rknn_matmul_info info {};
int m,k,n;
bool direct=std::getenv("LING3_MLA_PROBE_DIRECT") != nullptr;
std::vector<__fp16> normal_a,normal_b;
std::vector<float> normal_c;
Matmul(int rows,int inner,int columns,int core):m(rows),k(inner),n(columns),normal_a(m*k),normal_b(k*n),normal_c(m*n) {
try {
info.M=m;info.K=k;info.N=n;
info.type=RKNN_FLOAT16_MM_FLOAT16_TO_FLOAT32;
info.B_layout=RKNN_MM_LAYOUT_NATIVE;info.AC_layout=RKNN_MM_LAYOUT_NATIVE;info.iommu_domain_id=2;
Check(rknn_matmul_create(&ctx,&info,&attr),"create FP16 matmul");
Check(rknn_matmul_set_core_mask(ctx,static_cast<rknn_core_mask>(1U<<core)),"core mask");
if(attr.A.size != std::size_t(m)*k*2 || attr.B.size != std::size_t(k)*n*2 || attr.C.size != std::size_t(m)*n*4)
throw std::runtime_error("unexpected native-layout size");
if(direct && (attr.B.n_dims!=4 || attr.B.dims[2]!=16 || attr.B.dims[3]!=32))
throw std::runtime_error("direct B needs native [N/16,K/32,16,32]");
a=rknn_create_mem2(ctx,attr.A.size,RKNN_FLAG_MEMORY_CACHEABLE);
b=rknn_create_mem2(ctx,attr.B.size,RKNN_FLAG_MEMORY_CACHEABLE);
c=rknn_create_mem2(ctx,attr.C.size,RKNN_FLAG_MEMORY_CACHEABLE);
if(!a||!b||!c) throw std::runtime_error("RKNN allocation failed");
Check(rknn_matmul_set_io_mem(ctx,a,&attr.A),"bind A");
Check(rknn_matmul_set_io_mem(ctx,b,&attr.B),"bind B");
Check(rknn_matmul_set_io_mem(ctx,c,&attr.C),"bind C");
} catch(...) { Release(); throw; }
}
Matmul(const Matmul &)=delete;
void Release() { if(c)rknn_destroy_mem(ctx,c); if(b)rknn_destroy_mem(ctx,b); if(a)rknn_destroy_mem(ctx,a); if(ctx)rknn_matmul_destroy(ctx); }
~Matmul(){ Release(); }
void A(int row,int kk,float x) {
if(direct)static_cast<__fp16 *>(a->virt_addr)[(kk/8*m+row)*8+kk%8]=x;
else normal_a[row*k+kk]=x;
}
void B(int kk,int col,float x) {
if(direct)static_cast<__fp16 *>(b->virt_addr)[((col/16*(k/32)+kk/32)*16+col%16)*32+kk%32]=x;
else normal_b[kk*n+col]=x;
}
float C(int row,int col) const {
return direct?static_cast<const float *>(c->virt_addr)[(col/4*m+row)*4+col%4]:normal_c[row*n+col];
}
void Run(Timing & s, bool qk) {
auto t=Clock::now();
if(!direct) {
auto * pa=static_cast<__fp16 *>(a->virt_addr);
for(int kk=0;kk<k;kk+=8)for(int row=0;row<m;++row)
std::memcpy(pa+(kk/8*m+row)*8,normal_a.data()+row*k+kk,16);
Check(rknn_B_normal_layout_to_native_layout(normal_b.data(),b->virt_addr,k,n,&info),"pack native B");
}
s.stage+=Ms(t);t=Clock::now();
Check(rknn_mem_sync(ctx,a,RKNN_MEMORY_SYNC_TO_DEVICE),"sync A");
Check(rknn_mem_sync(ctx,b,RKNN_MEMORY_SYNC_TO_DEVICE),"sync dynamic B"); s.sync+=Ms(t);
t=Clock::now(); Check(rknn_matmul_run(ctx),"run matmul"); (qk?s.qk:s.av)+=Ms(t);
t=Clock::now(); Check(rknn_mem_sync(ctx,c,RKNN_MEMORY_SYNC_FROM_DEVICE),"sync C");s.sync+=Ms(t);
t=Clock::now();const auto * pc=static_cast<const float *>(c->virt_addr);
if(!direct) {
for(int nn=0;nn<n;nn+=4)for(int row=0;row<m;++row)
std::memcpy(normal_c.data()+row*n+nn,pc+(nn/4*m+row)*4,16);
}
s.stage+=Ms(t);
}
void Warmup() { std::fill(normal_a.begin(),normal_a.end(),0);std::fill(normal_b.begin(),normal_b.end(),0);std::memset(a->virt_addr,0,a->size);std::memset(b->virt_addr,0,b->size);Timing s;Run(s,true); }
};
struct Lane {
Matmul qk,av;
std::vector<float> maximum,denominator,factor,accumulator;
Timing timing;
Lane(int rows,int tile,int core):qk(rows,D,tile,core),av(rows,tile,V,core),maximum(rows),denominator(rows),factor(rows),accumulator(rows*V) {}
void Head(const Data & d,int h,int tile,std::vector<float> & out) {
std::fill(maximum.begin(),maximum.end(),-std::numeric_limits<float>::infinity());
std::fill(denominator.begin(),denominator.end(),0);std::fill(accumulator.begin(),accumulator.end(),0);
auto t=Clock::now();
for(int r=0;r<d.rows;++r) for(int i=0;i<D;++i) qk.A(r,i,d.q[(r*H+h)*D+i]);
timing.stage+=Ms(t);
for(int base=0;base<d.history;base+=tile) {
const int valid=std::min(tile,d.history-base);
t=Clock::now();
for(int j=0;j<tile;++j) for(int i=0;i<D;++i)
qk.B(i,j,j<valid?Float(d.k[(std::size_t(base+j)*H+h)*D+i]):0);
for(int j=0;j<tile;++j) for(int i=0;i<V;++i)
av.B(j,i,j<valid?Float(d.v[(std::size_t(base+j)*H+h)*V+i]):0);
timing.stage+=Ms(t);qk.Run(timing,true);
t=Clock::now();
for(int r=0;r<d.rows;++r) {
const int count=std::clamp(d.history-d.rows+r+1-base,0,valid);
if(count==0) { factor[r]=1;for(int j=0;j<tile;++j)av.A(r,j,0);continue; }
float mx=maximum[r];
for(int j=0;j<count;++j) mx=std::max(mx,qk.C(r,j)/std::sqrt(float(D)));
factor[r]=std::exp(maximum[r]-mx); float sum=0;
for(int j=0;j<count;++j) {
const float p=std::exp(qk.C(r,j)/std::sqrt(float(D))-mx);sum+=p; av.A(r,j,p);
}
for(int j=count;j<tile;++j)av.A(r,j,0);
denominator[r]=denominator[r]*factor[r]+sum;maximum[r]=mx;
}
timing.softmax+=Ms(t);av.Run(timing,false);
t=Clock::now();
for(int r=0;r<d.rows;++r)for(int i=0;i<V;++i) accumulator[r*V+i]=accumulator[r*V+i]*factor[r]+av.C(r,i);
timing.merge+=Ms(t);
}
t=Clock::now();
for(int r=0;r<d.rows;++r)for(int i=0;i<V;++i)out[(r*H+h)*V+i]=accumulator[r*V+i]/denominator[r];
timing.merge+=Ms(t);
}
};
struct Npu {
int tile;std::array<std::unique_ptr<Lane>,3> lanes;
Npu(int rows,int block):tile(block) { for(int c=0;c<3;++c)lanes[c]=std::make_unique<Lane>(rows,tile,c); }
double Run(const Data & d,std::vector<float> & out) {
auto t=Clock::now();constexpr std::array<int,3> cores {0,1,2};
ling3::CoreWorkers::Instance().Run(cores,[&](int c){ auto & l=*lanes[c];l.timing={};for(int h=c;h<H;h+=3)l.Head(d,h,tile,out); });
return Ms(t);
}
};
Json Error(const std::vector<float> & a,const std::vector<float> & b) {
double aa=0,bb=0,ab=0,ss=0,mx=0;
for(std::size_t i=0;i<a.size();++i) {
if(!std::isfinite(a[i])||!std::isfinite(b[i]))throw std::runtime_error("nonfinite output");
double x=a[i],y=b[i],e=x-y;aa+=x*x;bb+=y*y;ab+=x*y;ss+=e*e;mx=std::max(mx,std::abs(e));
}
if(aa==0 || bb==0)throw std::runtime_error("zero attention output");
return {{"cosine",ab/std::sqrt(aa*bb)},{"relative_l2",std::sqrt(ss/bb)},{"max_abs",mx}};
}
int main(int argc,char ** argv) {
try {
if(argc!=5 && argc!=6)throw std::invalid_argument("usage: ling3-mla-npu-probe HISTORY ROWS TILE REPEATS [SNAPSHOT]");
int n=std::stoi(argv[1]),m=std::stoi(argv[2]),tile=std::stoi(argv[3]),repeats=std::stoi(argv[4]);
if(m<1||m>128||n<m||n>65536||tile<32||tile>4096||tile%32||repeats<1||repeats>10)throw std::invalid_argument("invalid dimensions");
Data data=argc==6?Load(argv[5]):Data(m,n);
if(data.rows!=m||data.history!=n)throw std::invalid_argument("snapshot dimensions differ");
if(argc==5)Randomize(data,12345);
std::vector<float> cpu(m*H*V),npu(cpu.size());Cpu reference(n);
if(std::getenv("LING3_MLA_PROBE_INTEGRATED")) {
ling3::MlaNpu backend;backend.Prepare(m);
const double cpu_ms=reference.Run(data,cpu);const auto started=Clock::now();
if(!backend.Run(data.q,data.k,data.v,m,n,npu))throw std::runtime_error("integrated NPU path did not run");
const double npu_ms=Ms(started);const auto error=Error(npu,cpu);
std::cout<<Json({{"integrated",true},{"cpu_ms",cpu_ms},{"npu_ms",npu_ms},{"source",argc==6?argv[5]:"synthetic"},{"error",error}}).dump()<<std::endl;
return error["relative_l2"].get<double>()<0.001?0:1;
}
auto begin=Clock::now();Npu candidate(m,tile);double setup=Ms(begin);
for(auto & l:candidate.lanes){l->qk.Warmup();l->av.Warmup();}
// Also warms the persistent worker pool. Neither path's setup is in timings.
const double first=candidate.Run(data,npu);
const double cpu_warmup=reference.Run(data,cpu);
Json results=Json::array();
for(int r=0;r<repeats;++r) {
// The second repetition replaces ALL dynamic K/V. Exposes runtimes that
// incorrectly retain B. Real captured tensors stay intact on all repeats.
if(r==1 && argc==5)Randomize(data,67890);
double cpu_ms=reference.Run(data,cpu),npu_ms=candidate.Run(data,npu);
Json lanes=Json::array();for(auto & l:candidate.lanes)lanes.push_back(l->timing.json());
auto error=Error(npu,cpu);
results.push_back({{"cpu_ms",cpu_ms},{"npu_cpu_ms",npu_ms},{"speedup",cpu_ms/npu_ms},{"error",error},{"lane_work_ms",lanes}});
if(error["cosine"].get<double>()<0.999 || error["relative_l2"].get<double>()>0.02)
throw std::runtime_error("attention accuracy failed: "+results.dump());
}
std::cout<<Json({{"passed",true},{"scope","single MLA layer attention body; no projections/FFN/model TTFT"},
{"source",argc==6?argv[5]:"synthetic gaussian; BF16 KV; dynamic B replacement on repeat 2"},
{"history",n},{"rows",m},{"tile",tile},{"heads",H},{"npu_cores",3},{"cpu_baseline_cores",4},
{"direct_native_buffers",candidate.lanes[0]->qk.direct},
{"setup_ms",setup},{"warmup_attention_ms",first},{"cpu_warmup_ms",cpu_warmup},{"runs",results}}).dump()<<std::endl;
}catch(const std::exception & e){std::cerr<<"probe_failed="<<e.what()<<std::endl;return 1;}
}