Scala中如何为Systolic Array输入可计算的非方阵数据?
非方阵Systolic Array数据流控制方案
要扩展你的代码到非方阵场景,首先明确矩阵乘法的核心参数:
- 设输入矩阵 A为M行K列,B为K行N列(满足矩阵乘法的维度要求:A的列数=B的行数)
- Systolic Array的规模对应结果矩阵C的维度:M行N列,因此
h_in有M个端口(对应阵列每行的横向输入,即A的行元素),v_in有N个端口(对应阵列每列的纵向输入,即B的列元素)
核心逻辑调整
原方阵代码利用了M=N=K的对称性,非方阵下需要拆分h_in和v_in的处理逻辑,分别匹配A、B的维度:
- 横向输入(h_in):A的第i行元素从第
i个周期开始,每个周期输入1个元素,共输入K个周期,之后补0 - 纵向输入(v_in):B的第j列元素从第
j个周期开始,每个周期输入1个元素,共输入K个周期,之后补0 - 总周期数:取
M + K + N - 2,覆盖所有元素输入完成及最后一个PE的计算结束
修改后的Scala代码
// 定义非方阵参数:A(M×K),B(K×N) val M = 4 // A的行数(阵列行数) val K = 3 // A的列数 = B的行数 val N = 3 // B的列数(阵列列数) // 遍历所有需要的时钟周期 for (clk <- 0 until M + K + N - 2) { // 处理横向输入h_in:对应A的每一行 for (i <- 0 until M) { val colIdx = clk - i if (colIdx >= 0 && colIdx < K) { dut.io.h_in(i).poke(a(i)(colIdx)) } else { dut.io.h_in(i).poke(0) } } // 处理纵向输入v_in:对应B的每一列 for (j <- 0 until N) { val rowIdx = clk - j if (rowIdx >= 0 && rowIdx < K) { dut.io.v_in(j).poke(b(rowIdx)(j)) } else { dut.io.v_in(j).poke(0) } } // 触发时钟步进(根据你的测试框架调整,比如Chisel的clock.step(1)) dut.clock.step(1) }
关键说明
- 对于
h_in[i]:colIdx = clk - i确保第i行的元素从第i个周期开始依次输入A[i][0], A[i][1], ..., A[i][K-1] - 对于
v_in[j]:rowIdx = clk - j确保第j列的元素从第j个周期开始依次输入B[0][j], B[1][j], ..., B[K-1][j] - 总周期数
M+K+N-2是因为阵列右下角的PE(i=M-1, j=N-1)需要等待A的最后一个元素(clk=(M-1)+(K-1))和B的最后一个元素(clk=(N-1)+(K-1))到达,再完成K次累加计算,最终在(M-1)+(N-1)+K-1 = M+K+N-3周期完成,因此循环到M+K+N-2即可覆盖所有过程
内容的提问来源于stack exchange,提问作者IskandarZhang
相关产品推荐
相关产品推荐

