能否通过BLAS实现双索引张量收缩?Fortran测试案例解析
如何用BLAS实现双索引张量收缩?
问题背景
需要计算张量收缩 D[a,d] = A[a,b,c] * B[b,c,d],尝试两种实现方式:
- 方法一:将
A[a,b,c]重塑为C1[a,e](e = nb*nc),B[b,c,d]重塑为C2[e,d],再调用dgemm计算,该方法可行,测试耗时2.4476×10⁻²秒。 - 方法二:直接调用
dgemm传入三维数组,运行时触发Intel MKL错误:"Parameter 10 was incorrect on entry to DGEMM",测试耗时1.8388×10⁻²秒,且结果与方法一的最大差值为5.46978468774136。
现咨询:是否可以通过BLAS直接实现双索引的张量收缩?已知现有BLAS方案仅支持单索引收缩。
错误原因分析
BLAS的dgemm是为二维矩阵乘法设计的,它要求输入数组严格遵循矩阵的内存存储顺序(Fortran默认列优先),并正确指定矩阵的主维度(leading dimension)。
方法二中直接传入三维数组a和b时:
- 对于
a,你指定的主维度是Size(a, Dim=1)(即na),但三维数组a的内存布局是a(1,1,1), a(2,1,1), ..., a(na,1,1), a(1,2,1), ..., a(na,nb,nc),这与重塑后的C1[a,e](布局为a(1,1,1), a(1,1,2), ..., a(1,nb,nc), a(2,1,1), ...)完全不匹配。 dgemm会错误地将三维数组当成二维矩阵读取,导致参数校验失败(错误提示的Parameter 10对应LDB,即b的主维度),同时读取的数据完全错误,最终结果偏差。
解决方案与结论
BLAS基础接口(如dgemm)不直接支持双索引张量收缩,因为它仅能处理二维矩阵的单索引乘法。要实现双索引收缩,正确的做法是:
- 将张量的两个收缩维度合并为一个虚拟维度,把三维张量重塑为二维矩阵
- 调用
dgemm完成矩阵乘法 - 如有需要,再将结果重塑回目标张量形状
需要注意:
- Fortran的
reshape操作在内存布局兼容时可以避免数据复制,但如果原张量的维度顺序与目标矩阵不匹配,会触发数据复制。你可以通过指定reshape的order参数,或者调整数组的声明顺序,来优化内存布局,减少复制开销。例如,若提前将B声明为B(c,b,d),则reshape(B, [nb*nc, nd])可以直接匹配C2[e,d]的布局,无需复制。
测试代码
Program reshape_for_blas Use, Intrinsic :: iso_fortran_env, Only : wp => real64, li => int64 Implicit None Real( wp ), Dimension( :, :, : ), Allocatable :: a Real( wp ), Dimension( :, :, : ), Allocatable :: b Real( wp ), Dimension( :, : ), Allocatable :: c1, c2 Real( wp ), Dimension( :, : ), Allocatable :: d Real( wp ), Dimension( :, : ), Allocatable :: e Integer :: na, nb, nc, nd, ne Integer( li ) :: start, finish, rate Write( *, * ) 'na, nb, nc, nd ?' Read( *, * ) na, nb, nc, nd ne = nb * nc Allocate( a ( 1:na, 1:nb, 1:nc ) ) Allocate( b ( 1:nb, 1:nc, 1:nd ) ) Allocate( c1( 1:na, 1:ne ) ) Allocate( c2( 1:ne, 1:nd ) ) Allocate( d ( 1:na, 1:nd ) ) Allocate( e ( 1:na, 1:nd ) ) ! Set up some data Call Random_number( a ) Call Random_number( b ) ! With reshapes Call System_clock( start, rate ) c1 = Reshape( a, Shape( c1 ) ) c2 = Reshape( b, Shape( c2 ) ) Call dgemm( 'N', 'N', na, nd, ne, 1.0_wp, c1, Size( c1, Dim = 1 ), & c2, Size( c2, Dim = 1 ), & 0.0_wp, e, Size( e, Dim = 1 ) ) Call System_clock( finish, rate ) Write( *, * ) 'Time for reshaping method ', Real( finish - start, wp ) / rate ! Direct Call System_clock( start, rate ) Call dgemm( 'N', 'N', na, nd, ne, 1.0_wp, a , Size( a , Dim = 1 ), & b , Size( b , Dim = 1 ), & 0.0_wp, d, Size( d, Dim = 1 ) ) Call System_clock( finish, rate ) Write( *, * ) 'Time for straight method ', Real( finish - start, wp ) / rate Write( *, * ) 'Difference between result matrices ', Maxval( Abs( d - e ) ) End Program reshape_for_blas
内容的提问来源于stack exchange,提问作者Mao Yang
相关产品推荐
相关产品推荐

