File size: 9,892 Bytes
3fd1a35 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 | #include "mla_npu.h"
#include "core_workers.h"
#include <algorithm>
#include <array>
#include <bit>
#include <chrono>
#include <cmath>
#include <cstdlib>
#include <cstring>
#include <iostream>
#include <limits>
#include <map>
#include <stdexcept>
#include <vector>
#if LING3_WITH_RKNN && defined(__aarch64__)
#include <rknn_api.h>
#include <rknn_matmul_api.h>
#endif
namespace ling3 {
namespace {
using Clock=std::chrono::steady_clock;
constexpr int H=16,D=192,V=128,T=512;
[[maybe_unused]] std::size_t Bucket(std::size_t rows) {
for(std::size_t n:{16,32,64,128}) if(rows<=n)return n;
throw std::invalid_argument("MLA NPU rows exceed 128");
}
#if LING3_WITH_RKNN && defined(__aarch64__)
void Check(int rc,const char * op) {
if(rc)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);}
__fp16 Half(float x) {
if(!std::isfinite(x)||std::abs(x)>65504.0F)
throw std::runtime_error("MLA activation outside finite FP16 range");
return static_cast<__fp16>(x);
}
struct Matrix {
rknn_matmul_ctx ctx=0;
rknn_tensor_mem *a=nullptr,*b=nullptr,*c=nullptr;
int m,k,n;
Matrix(int rows,int inner,int cols,int core):m(rows),k(inner),n(cols) {
try {
rknn_matmul_info info{};rknn_matmul_io_attr io{};
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,&io),"MLA create");
Check(rknn_matmul_set_core_mask(ctx,static_cast<rknn_core_mask>(1U<<core)),"MLA core mask");
if(io.A.size!=std::size_t(m)*k*2||io.B.size!=std::size_t(k)*n*2||io.C.size!=std::size_t(m)*n*4||
io.B.n_dims!=4||io.B.dims[2]!=16||io.B.dims[3]!=32)
throw std::runtime_error("MLA unsupported native layout");
a=rknn_create_mem2(ctx,io.A.size,RKNN_FLAG_MEMORY_CACHEABLE);
b=rknn_create_mem2(ctx,io.B.size,RKNN_FLAG_MEMORY_CACHEABLE);
c=rknn_create_mem2(ctx,io.C.size,RKNN_FLAG_MEMORY_CACHEABLE);
if(!a||!b||!c)throw std::runtime_error("MLA workspace allocation failed");
std::memset(a->virt_addr,0,a->size);std::memset(b->virt_addr,0,b->size);
Check(rknn_matmul_set_io_mem(ctx,a,&io.A),"MLA bind A");
Check(rknn_matmul_set_io_mem(ctx,b,&io.B),"MLA bind B");
Check(rknn_matmul_set_io_mem(ctx,c,&io.C),"MLA bind C");
// Warm each native matmul without modifying decoder KV or recurrent state.
Run();
for(int r=0;r<m;++r)for(int i=0;i<k;++i)A(r,i,0.125F);
for(int i=0;i<k;++i)for(int col=0;col<n;++col)B(i,col,0.25F);
Run();
for(int r=0;r<m;++r)for(int col=0;col<n;++col)
if(!std::isfinite(C(r,col))||std::abs(C(r,col)-k/32.0F)>0.001F)
throw std::runtime_error("MLA native FP16 capability probe failed");
} catch(...) {Release();throw;}
}
Matrix(const Matrix &)=delete;
~Matrix(){Release();}
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);}
void A(int row,int kk,float x){static_cast<__fp16 *>(a->virt_addr)[(kk/8*m+row)*8+kk%8]=Half(x);}
void B(int kk,int col,float x){static_cast<__fp16 *>(b->virt_addr)[((col/16*(k/32)+kk/32)*16+col%16)*32+kk%32]=Half(x);}
float C(int row,int col)const{return static_cast<const float *>(c->virt_addr)[(col/4*m+row)*4+col%4];}
void Run(bool sync_b=true){
Check(rknn_mem_sync(ctx,a,RKNN_MEMORY_SYNC_TO_DEVICE),"MLA sync A");
if(sync_b)Check(rknn_mem_sync(ctx,b,RKNN_MEMORY_SYNC_TO_DEVICE),"MLA sync B");
Check(rknn_matmul_run(ctx),"MLA matmul");
Check(rknn_mem_sync(ctx,c,RKNN_MEMORY_SYNC_FROM_DEVICE),"MLA sync C");
}
};
struct Lane {
Matrix qk,av;
std::vector<float> maximum,denominator,factor,accumulator,scores,av_high;
Lane(int m,int core):qk(m,D,T,core),av(m,T,V,core),maximum(m),denominator(m),factor(m),accumulator(m*V),scores(m*T),av_high(m*V){}
void Head(std::span<const float> q,std::span<const std::uint16_t> k,std::span<const std::uint16_t> v,
int rows,std::size_t history,int h,std::span<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);
// Clear padded rows as workspaces are reused across partial requests.
const float scale=1.0F/std::sqrt(float(D));
constexpr float residual_scale=4096.0F;
for(std::size_t base=0;base<history;base+=T){
const int valid=std::min<std::size_t>(T,history-base);
for(int j=0;j<T;++j){
for(int i=0;i<D;++i)qk.B(i,j,j<valid?Float(k[((base+j)*H+h)*D+i]):0);
for(int i=0;i<V;++i)av.B(j,i,j<valid?Float(v[((base+j)*H+h)*V+i]):0);
}
// BF16 K/V are exact in FP16 over their representable range. Preserve
// FP32 Q and probabilities with a scaled residual pass; plain FP16 Q/P
// errors can otherwise amplify through quantization and MoE routing.
for(int r=0;r<qk.m;++r)for(int i=0;i<D;++i)qk.A(r,i,r<rows?q[(r*H+h)*D+i]:0);
qk.Run();
for(int r=0;r<rows;++r)for(int j=0;j<T;++j)scores[r*T+j]=qk.C(r,j);
for(int r=0;r<qk.m;++r)for(int i=0;i<D;++i){
const float x=r<rows?q[(r*H+h)*D+i]:0;
qk.A(r,i,(x-static_cast<float>(Half(x)))*residual_scale);
}
qk.Run(false);
for(int r=0;r<rows;++r)for(int j=0;j<T;++j)
scores[r*T+j]=(scores[r*T+j]+qk.C(r,j)/residual_scale)*scale;
for(int r=0;r<qk.m;++r){
const auto causal_end=history-rows+std::min(r,rows-1)+1;
const int count=r<rows&&causal_end>base?std::min<std::size_t>(valid,causal_end-base):0;
if(!count){factor[r]=1;for(int j=0;j<T;++j){av.A(r,j,0);scores[r*T+j]=0;}continue;}
float mx=maximum[r];
for(int j=0;j<count;++j)mx=std::max(mx,scores[r*T+j]);
factor[r]=std::exp(maximum[r]-mx);float sum=0;
for(int j=0;j<count;++j){const float p=std::exp(scores[r*T+j]-mx);sum+=p;av.A(r,j,p);scores[r*T+j]=p;}
for(int j=count;j<T;++j){av.A(r,j,0);scores[r*T+j]=0;}
denominator[r]=denominator[r]*factor[r]+sum;maximum[r]=mx;
}
av.Run();
for(int r=0;r<rows;++r)for(int i=0;i<V;++i)av_high[r*V+i]=av.C(r,i);
for(int r=0;r<av.m;++r)for(int j=0;j<T;++j){
const float p=scores[r*T+j];av.A(r,j,(p-static_cast<float>(Half(p)))*residual_scale);
}
av.Run(false);
for(int r=0;r<rows;++r)for(int i=0;i<V;++i)
accumulator[r*V+i]=accumulator[r*V+i]*factor[r]+(av_high[r*V+i]+av.C(r,i)/residual_scale);
}
for(int r=0;r<rows;++r)for(int i=0;i<V;++i){
const float x=accumulator[r*V+i]/denominator[r];
if(!std::isfinite(x))throw std::runtime_error("non-finite MLA NPU output");
out[(r*H+h)*V+i]=x;
}
}
};
struct Workspace {
std::array<std::unique_ptr<Lane>,3> lanes;
explicit Workspace(int rows){for(int i=0;i<3;++i)lanes[i]=std::make_unique<Lane>(rows,i);}
};
#endif
}
struct MlaNpu::Impl {
std::string mode=std::getenv("LING3_MLA_BACKEND")?std::getenv("LING3_MLA_BACKEND"):"auto";
bool disabled=false;
MlaBackendStats stats;
#if LING3_WITH_RKNN && defined(__aarch64__)
std::map<std::size_t,std::unique_ptr<Workspace>> workspaces;
#endif
Impl(){if(mode!="auto"&&mode!="cpu"&&mode!="npu")throw std::invalid_argument("MLA backend must be auto, cpu or npu");}
void Failure(const std::exception & e){disabled=true;++stats.fallbacks;std::cerr<<"mla_warning: NPU attention unavailable, using CPU: "<<e.what()<<'\n';}
void Prepare(std::size_t rows){
if(mode=="cpu"||disabled||rows<16)return;
#if LING3_WITH_RKNN && defined(__aarch64__)
const auto bucket=Bucket(rows);
try{if(!workspaces.contains(bucket))workspaces.emplace(bucket,std::make_unique<Workspace>(bucket));}
catch(const std::exception & e){Failure(e);}
#else
disabled=true;
#endif
}
};
MlaNpu::MlaNpu():impl_(std::make_unique<Impl>()){}
MlaNpu::~MlaNpu()=default;
void MlaNpu::Prepare(std::size_t rows){impl_->Prepare(rows);}
void MlaNpu::CpuCall(){++impl_->stats.cpu_calls;}
MlaBackendStats MlaNpu::Stats()const{return impl_->stats;}
bool MlaNpu::Run(std::span<const float> q,std::span<const std::uint16_t> k,std::span<const std::uint16_t> v,
std::size_t rows,std::size_t history,std::span<float> output){
if(rows<1||rows>128||history<rows||q.size()!=rows*H*D||k.size()!=history*H*D||v.size()!=history*H*V||output.size()!=rows*H*V)
throw std::invalid_argument("MLA backend shape mismatch");
if(impl_->mode=="cpu"||impl_->disabled||rows<16||(impl_->mode=="auto"&&history<256))return false;
impl_->Prepare(rows);
if(impl_->disabled)return false;
#if LING3_WITH_RKNN && defined(__aarch64__)
auto & work=*impl_->workspaces.at(Bucket(rows));const auto start=Clock::now();
try{
constexpr std::array<int,3> cores{0,1,2};
CoreWorkers::Instance().Run(cores,[&](int c){for(int h=c;h<H;h+=3)work.lanes[c]->Head(q,k,v,rows,history,h,output);});
++impl_->stats.npu_calls;
impl_->stats.npu_ms+=std::chrono::duration<double,std::milli>(Clock::now()-start).count();
return true;
}catch(const std::exception & e){impl_->Failure(e);}
#endif
return false;
}
}
|