Ling-3.0-tiny-RKNN / src /mla_npu.cpp
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
9.89 kB
#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;
}
}