You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

能否通过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)不直接支持双索引张量收缩,因为它仅能处理二维矩阵的单索引乘法。要实现双索引收缩,正确的做法是:

  1. 将张量的两个收缩维度合并为一个虚拟维度,把三维张量重塑为二维矩阵
  2. 调用dgemm完成矩阵乘法
  3. 如有需要,再将结果重塑回目标张量形状

需要注意:

  • 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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.12 06:20:26