Keyboard shortcuts

Press or to navigate between chapters

Press S or / to search in the book

Press ? to show this help

Press Esc to hide this help

ndarray适配器进阶用法

除了普通的实数乘积,适配器还镜像了 gemmkit 的整个表面:

  • 整数 GEMM
  • 面向量化推理的重量化输出
  • 带可选共轭的复数乘积
  • 融合的偏置与激活
  • 用户自定义的逐元素映射
  • 在三维数组上的批量乘法
  • 预打包操作数

适配器让每个类型族都由与之同名的 Cargo feature 门控。每个类型族也都保持快速上手页里那些普通入口的形态:直接从数组读步长,转发给 gemmkit,仅在维度不匹配时 panic。fused 入口再多一种 panic 情形:偏置切片与 C 重叠。每个入口也都有一个携带调用方自有 Workspace_with 孪生。

整数 GEMM(int8

gemm_i8i8 输入乘进一个 i32 累加器:C(i32) <- alpha*A(i8)*B(i8) + beta*C,其中 alphabetaC 都是 i32。它是独立于 gemm 的入口,因为输入与输出的元素类型不同。算术在溢出时回绕,这是整数 GEMM 的惯例语义。dot_i8 是它的便捷孪生。它返回一个新建的 Array2<i32>

#![allow(unused)]
fn main() {
use gemmkit_ndarray::{Parallelism, dot_i8, gemm_i8};
use ndarray::Array2;

let a = Array2::<i8>::zeros((16, 12));
let b = Array2::<i8>::zeros((12, 10));

// i8 输入,i32 累加器
let c: Array2<i32> = dot_i8(&a, &b);

// 带 i32 alpha/beta 的通用形式,累加进已有累加器
let mut acc = Array2::<i32>::zeros((16, 10));
gemm_i8(2, &a, &b, 1, &mut acc, Parallelism::Serial);
}

重量化输出(int8 + epilogue

量化推理很少想要原始的 i32 累加器,它要的是一个 8 位张量。gemm_i8_requant 把乘法和重量化融进 1 趟。它直接把 i32 累加器折叠成 i8 输出,全程不物化完整的 m*n 中间结果。它没有 alpha 参数,因为其折进了 scale。它也没有 beta 参数,因为累加进量化输出没有良好定义。参数装在一个 Requantize 里:

#![allow(unused)]
fn main() {
use gemmkit_ndarray::{Parallelism, RequantScale, Requantize, gemm_i8_requant, gemm_i8_requant_u8};
use ndarray::Array2;

let a = Array2::<i8>::zeros((16, 12));
let b = Array2::<i8>::zeros((12, 10));

// i8 输出,范围 [-128, 127],单个 per-tensor scale,per-row 偏置(长度 A.rows)
let bias: Vec<i32> = vec![0; 16];
let mut c = Array2::<i8>::zeros((16, 10));
let req = Requantize {
    scale: RequantScale::PerTensor(0.05),
    zero_point: -7,
    bias: Some(&bias),
};
gemm_i8_requant(&a, &b, req, &mut c, Parallelism::default());

// u8 输出,范围 [0, 255],per-channel scale,无偏置
let scales: Vec<f32> = vec![0.02; 16]; // 每个输出行 / 通道一个
let mut cu = Array2::<u8>::zeros((16, 10));
gemm_i8_requant_u8(
    &a,
    &b,
    Requantize { scale: RequantScale::PerRow(&scales), zero_point: 128, bias: None },
    &mut cu,
    Parallelism::default(),
);
}

输出为 clamp(zero_point + round_ne(scale * (accumulator + bias[i])), LO, HI),采用四舍六入五成双。其中 scale 是 per-tensor 的那个值,或 per-row 的 scale_iu8 变体是 ONNX-QLinearMatMul 风格的激活:除了输出域 [0, 255]zero_point 的取值范围外,与 gemm_i8_requant 完全相同。两者都会拒绝:

  • 非有限或非正的 scale,per-tensor 或 per-row
  • 长度不等于 A.rows 的 per-row scale 或偏置
  • C 重叠的切片
  • 超出该入口取值域的 zero_point

复数 GEMM(complex

复数乘积有自己的入口,因为这 2 个共轭标志放不进同构的实数签名。gemm_cplx 计算 C <- alpha*op(A)*op(B) + beta*C,元素类型为 Complex<f32>Complex<f64>。设置 conj_a 时,op(A)conj(A)。设置 conj_b 时,op(B)conj(B)dot_cplx 是不做共轭的便捷入口。

#![allow(unused)]
fn main() {
use gemmkit_ndarray::{Complex, Parallelism, dot_cplx, gemm_cplx};
use ndarray::Array2;

type C = Complex<f64>;
let a = Array2::<C>::from_elem((8, 6), Complex::new(0.0, 0.0));
let b = Array2::<C>::from_elem((6, 5), Complex::new(0.0, 0.0));

// 普通 A*B
let c = dot_cplx(&a, &b);

// 对 A 取共轭,累加进已有 C
let mut acc = Array2::<C>::from_elem((8, 5), Complex::new(0.0, 0.0));
gemm_cplx(
    Complex::new(1.0, 0.0),
    &a,
    true,  // conj_a
    &b,
    false, // conj_b
    Complex::new(0.0, 0.0),
    &mut acc,
    Parallelism::Serial,
);
}

同时开启 complexepilogue 后,gemm_cplx_fused 会在同一趟里加上一个可选的 Bias(原样相加,绝不共轭)。它不接受激活参数。像 ReLU 这样带序关系的激活,在复数上没有定义。

融合偏置、激活与映射(epilogue

gemm_fused 在 1 趟里算出 C <- act(alpha*A*B + beta*C + bias)。偏置是可选的 Bias::PerRow(长度 A.rows)或 Bias::PerCol(长度 B.cols)。激活是可选的 ReluLeakyRelu(slope),最后施加。两者都设为 None 时,gemm_fused 就是 gemm

#![allow(unused)]
fn main() {
use gemmkit_ndarray::{Activation, Bias, Parallelism, gemm_fused, gemm_map};
use ndarray::Array2;

let a = Array2::<f32>::zeros((12, 9));
let b = Array2::<f32>::zeros((9, 7));

// 一趟算出 C <- ReLU(A*B + bias);PerRow 偏置长度为 A.rows
let bias: Vec<f32> = vec![0.0; 12];
let mut c = Array2::<f32>::zeros((12, 7));
gemm_fused(
    1.0, &a, &b, 0.0, &mut c,
    Some(Bias::PerRow(&bias)),
    Some(Activation::Relu),
    Parallelism::default(),
);

// 任意逐元素闭包 f(value, row, col);这里是一个 relu6
let f = |v: f32, _r: usize, _c: usize| v.max(0.0).min(6.0);
let mut c2 = Array2::<f32>::zeros((12, 7));
gemm_map(1.0, &a, &b, 0.0, &mut c2, &f, Parallelism::default());
}

BiasActivationgemmkit_ndarray 重新导出,所以你无需为它们再点名 gemmkit。对 f32/f64,无论什么形状,gemm_fused 都与“先 gemm 再做同样的标量映射”逐位相同。对 f16/bf16,尾部运算会先以 f32 进行,再做一次窄化。这比单独的窄化映射更精确,所以对这些类型而言,结果与“先 gemm 再映射”并不逐位相等。

gemm_map 是通用的逐元素扩展点。闭包 f(value, row, col) 看到的是每个输出元素的最终值,(row, col) 处在 C 的用户坐标系里。gemmkit 对每个元素恰好调用它一次。它的代价是每个元素 1 次间接调用。普通的偏置或激活优先用 gemm_fused,因为它会向量化。GELU、sigmoid、clamp,或依赖位置的变换,才用 gemm_map。这里的 T 只能是 f32/f64

批量 GEMM

这是唯一没有普通 gemm 对应、在同类适配器里也没有对手的运算。它是一叠彼此独立的乘积,承载在三维 Array3 上,批次维在 0 轴。a(batch, m, k)b(batch, k, n)c(batch, m, n)。0 轴是各操作数的批次步长。1、2 轴是元素步长。gemm_batched 在批次上并行。每个元素都在 1 个 worker 上运行,因此结果精确复现一个 gemm 调用循环。

#![allow(unused)]
fn main() {
use gemmkit_ndarray::{Parallelism, dot_batched, gemm_batched};
use ndarray::Array3;

let a = Array3::<f32>::zeros((32, 8, 5)); // (batch, m, k)
let b = Array3::<f32>::zeros((32, 5, 6)); // (batch, k, n)

// 一叠乘积
let c = dot_batched(&a, &b); // (32, 8, 6)

// 累加进已有累加器的通用形式
let mut acc = Array3::<f32>::zeros((32, 8, 6));
gemm_batched(0.7, &a, &b, 1.3, &mut acc, Parallelism::default());
}

适配器只读步长,所以一个换轴的、或本就一般步长的三维视图,都能无拷贝转发。比如,a.view().permuted_axes([0, 2, 1]) 把一块 (batch, k, m) 缓冲区变成 (batch, m, k) 视图,并直接批量转发。

epilogue 下,gemm_batched_fused 对这叠里的每个元素都施加同一个共享的 Bias/Activation。这正是批量线性层的情形。偏置按单个元素定尺寸(PerRow 长度 mPerCol 长度 n),而不是整批。

预打包操作数

当一个操作数固定、另一个成流而来时,把固定那一侧打包一次并复用。这样就省掉了每次调用的重复打包。prepack_rhs 为复用的 B 返回一个 PackedRhs<T>,由 gemm_packed_b 消费。prepack_lhs 为复用的 A 返回一个 PackedLhs<T>,由 gemm_packed_a 消费。打包函数直接读步长,所以 BA 可以是任意布局。

两者各有 1 条朝向约束。gemm_packed_b 需要一个偏列主序的 C|col stride| >= |row stride|)。gemm_packed_a 需要一个偏行主序的 C|col stride| <= |row stride|)。另一种朝向会交换操作数,并使打包句柄失效,gemmkit 会拒绝这种情况。不合适的布局请用普通 gemm

融合孪生 gemm_packed_b_fusedgemm_packed_a_fused 接受同样的句柄,再加上偏置和激活。这正是固定权重的推理层。下面的例子把一个权重矩阵作为 LHS 打包一次,并在各推理步之间复用。它融进了一个 per-output-channel 的偏置和一个 ReLU:

#![allow(unused)]
fn main() {
use gemmkit_ndarray::{Activation, Bias, Parallelism, gemm_packed_a_fused, prepack_lhs};
use ndarray::Array2;

let (out, in_features) = (256usize, 512usize);

// 把固定权重 W: (out, in) 打包一次
let w = Array2::<f32>::zeros((out, in_features));
let packed = prepack_lhs(&w);
let bias: Vec<f32> = vec![0.0; out]; // per-output-channel,长度 C.rows

// 每个推理步:激活 x (in, batch) -> y (out, batch)
let batch = 32;
let x = Array2::<f32>::zeros((in_features, batch));
let mut y = Array2::<f32>::zeros((out, batch)); // 行主序(packed_a 朝向)
gemm_packed_a_fused(
    1.0,
    &packed,
    &x,
    0.0,
    &mut y,
    Some(Bias::PerRow(&bias)),
    Some(Activation::Relu),
    Parallelism::default(),
);
}

把打包句柄与 _with 工作区变体 gemm_packed_a_fused_with 搭配。这样一来,一个稳定的推理循环在首次调用后就不再分配。你在用户坐标系里指定偏置的轴。打包路径对 gemm_packed_b_fused 原样转发它,而对 gemm_packed_a_fused 则让核心去翻转它。无论你打包了哪个操作数,PerRow 始终表示“每个输出行 1 个值”。

本适配器与 ndarray 自带乘积的取舍

ndarray 本就能做矩阵乘法:.dot() 给出普通乘积,general_mat_mul 给出就地的 alpha/beta 形式。对一次没有额外需求的 f32/f64 乘积,这些函数就够用,还能少拉 1 个依赖。没必要出于习惯就走 gemmkit。

当你需要 ndarray 内建路径给不了的东西时,再选本适配器。gemmkit 会在运行的机器上,于运行时挑选最快的指令集。它不会在编译期就把某个选择写死(见运行时ISA分发)。它带来本页覆盖的更宽表面:融合偏置与激活、i8 与重量化推理、带共轭的复数、批量乘积,以及预打包。它还暴露调优旋钮,外加一个把分块校准到部署机器的安装期自动调优器(见调优旋钮)。

这些都不会改变你传入的数组或拿回的结果。适配器自始至终是同一套零拷贝步长转接。它只是拓宽了你能提出的请求。