Managed ONNX operator kernels and execution helpers for the editor inference runtime. Implements elementwise ops, broadcasting, shape ops, reductions, softmax/normalization, matmul/GEMM (including packed kernels), and utilities for tensor manipulation and execution context.
using System;
using System.Collections.Generic;
using System.Linq;
using System.Numerics;
using System.Runtime.InteropServices;
using System.Threading.Tasks;
namespace TextToAnimation.Editor.Inference.Onnx;
/// <summary>Per-run execution context (caches derived constants such as transposed weights).</summary>
public sealed class ExecContext
{
public readonly Dictionary<Tensor, Tensor> TransposedCache = new( ReferenceEqualityComparer.Instance );
/// <summary>Packed GEMM panels of constant matrices (weights), built on first use.</summary>
public readonly Dictionary<Tensor, FastKernels.PackedMatrix> Packed = new( ReferenceEqualityComparer.Instance );
/// <summary>Tensors that never change between runs (initializers).</summary>
public readonly HashSet<Tensor> Constants = new( ReferenceEqualityComparer.Instance );
/// <summary>Buffers of intermediate tensors, reused between nodes and runs.</summary>
public readonly TensorPool Pool = new();
/// <summary>The context of the run on this thread (kernels rent their outputs from its pool).</summary>
[ThreadStatic] public static ExecContext Current;
/// <summary>An output buffer with undefined contents: pooled while a session runs, fresh otherwise.</summary>
public static float[] Alloc( int n ) => Current?.Pool.Rent( n ) ?? new float[n];
/// <summary>A zero-filled output buffer.</summary>
public static float[] AllocZeroed( int n ) => Current?.Pool.RentZeroed( n ) ?? new float[n];
// two cores stay free: the editor's main and render threads keep running smoothly while a model runs on the CPU
public int MaxThreads = Math.Max( 1, Environment.ProcessorCount - 2 );
public ParallelOptions Parallel => new() { MaxDegreeOfParallelism = MaxThreads };
}
/// <summary>
/// The operator kernels of the managed ONNX interpreter (opset 13-18 semantics, NumPy broadcasting).
/// Floating math is fp32; heavy kernels (MatMul/Gemm, softmax, normalization) use <see cref="Vector{T}"/>
/// and run across cores.
/// </summary>
public static class OnnxOps
{
public delegate Tensor[] Kernel( OnnxNode node, Tensor[] inputs, ExecContext ctx );
static readonly Dictionary<string, Kernel> Kernels = new( StringComparer.Ordinal )
{
["Identity"] = ( n, i, c ) => new[] { i[0] },
["Dropout"] = ( n, i, c ) => new[] { i[0], Tensor.Bool( i[0].Shape ) },
["Add"] = ( n, i, c ) => new[] { FastBinary( i[0], i[1], BinOp.Add, FastKernels.Op.Add, c ) },
["Sub"] = ( n, i, c ) => new[] { FastBinary( i[0], i[1], BinOp.Sub, FastKernels.Op.Sub, c ) },
["Mul"] = ( n, i, c ) => new[] { FastBinary( i[0], i[1], BinOp.Mul, FastKernels.Op.Mul, c ) },
["Div"] = ( n, i, c ) => new[] { FastBinary( i[0], i[1], BinOp.Div, FastKernels.Op.Div, c ) },
["RMSNormalization"] = ( n, i, c ) => new[] { FastKernels.RmsNorm( i[0], i[1], (int)n.GetInt( "axis", -1 ), n.GetFloat( "epsilon", 1e-5f ), c ) },
["Attention"] = ( n, i, c ) => new[] { FastKernels.Attention( i[0], i[1], i[2], i.Length > 3 ? i[3] : null,
n.Attributes.ContainsKey( "scale" ) ? n.GetFloat( "scale", 1f ) : 1f / MathF.Sqrt( i[0].Shape[^1] ), c ) },
["Pow"] = ( n, i, c ) => new[] { Binary( i[0], i[1], BinOp.Pow ) },
["Max"] = ( n, i, c ) => new[] { i.Skip( 1 ).Aggregate( i[0], ( a, b ) => Binary( a, b, BinOp.Max ) ) },
["Min"] = ( n, i, c ) => new[] { i.Skip( 1 ).Aggregate( i[0], ( a, b ) => Binary( a, b, BinOp.Min ) ) },
["Sum"] = ( n, i, c ) => new[] { i.Skip( 1 ).Aggregate( i[0], ( a, b ) => Binary( a, b, BinOp.Add ) ) },
["Mod"] = ( n, i, c ) => new[] { Binary( i[0], i[1], n.GetInt( "fmod", 0 ) == 1 ? BinOp.FMod : BinOp.Mod ) },
["Equal"] = ( n, i, c ) => new[] { Compare( i[0], i[1], CmpOp.Eq ) },
["Less"] = ( n, i, c ) => new[] { Compare( i[0], i[1], CmpOp.Lt ) },
["LessOrEqual"] = ( n, i, c ) => new[] { Compare( i[0], i[1], CmpOp.Le ) },
["Greater"] = ( n, i, c ) => new[] { Compare( i[0], i[1], CmpOp.Gt ) },
["GreaterOrEqual"] = ( n, i, c ) => new[] { Compare( i[0], i[1], CmpOp.Ge ) },
["And"] = ( n, i, c ) => new[] { Logical( i[0], i[1], ( a, b ) => a && b ) },
["Or"] = ( n, i, c ) => new[] { Logical( i[0], i[1], ( a, b ) => a || b ) },
["Xor"] = ( n, i, c ) => new[] { Logical( i[0], i[1], ( a, b ) => a ^ b ) },
["Not"] = ( n, i, c ) => new[] { Tensor.Bool( i[0].Shape, i[0].B.Select( v => !v ).ToArray() ) },
["Where"] = ( n, i, c ) => new[] { Where( i[0], i[1], i[2] ) },
["Sqrt"] = ( n, i, c ) => new[] { Unary( i[0], MathF.Sqrt ) },
["Exp"] = ( n, i, c ) => new[] { Unary( i[0], MathF.Exp ) },
["Log"] = ( n, i, c ) => new[] { Unary( i[0], MathF.Log ) },
["Tanh"] = ( n, i, c ) => new[] { Unary( i[0], MathF.Tanh ) },
["Sin"] = ( n, i, c ) => new[] { Unary( i[0], MathF.Sin ) },
["Cos"] = ( n, i, c ) => new[] { Unary( i[0], MathF.Cos ) },
["Floor"] = ( n, i, c ) => new[] { Unary( i[0], MathF.Floor ) },
["Ceil"] = ( n, i, c ) => new[] { Unary( i[0], MathF.Ceiling ) },
["Round"] = ( n, i, c ) => new[] { Unary( i[0], v => MathF.Round( v, MidpointRounding.ToEven ) ) },
["Reciprocal"] = ( n, i, c ) => new[] { Unary( i[0], v => 1f / v ) },
["Sigmoid"] = ( n, i, c ) => new[] { FastKernels.Sigmoid( i[0], c ) },
["Relu"] = ( n, i, c ) => new[] { Unary( i[0], v => v > 0 ? v : 0 ) },
["Erf"] = ( n, i, c ) => new[] { Unary( i[0], Erf ) },
["Gelu"] = ( n, i, c ) => new[] { Unary( i[0], n.GetString( "approximate", "none" ) == "tanh"
? v => 0.5f * v * (1f + MathF.Tanh( 0.7978845608f * (v + 0.044715f * v * v * v) ))
: v => 0.5f * v * (1f + Erf( v * 0.70710678f )) ) },
["Neg"] = ( n, i, c ) => new[] { i[0].IsInt ? Tensor.Int64( i[0].Shape, i[0].L.Select( v => -v ).ToArray() ) : Unary( i[0], v => -v ) },
["Abs"] = ( n, i, c ) => new[] { i[0].IsInt ? Tensor.Int64( i[0].Shape, i[0].L.Select( Math.Abs ).ToArray() ) : Unary( i[0], MathF.Abs ) },
["Sign"] = ( n, i, c ) => new[] { Unary( i[0], v => MathF.Sign( v ) ) },
["IsNaN"] = ( n, i, c ) => new[] { Tensor.Bool( i[0].Shape, i[0].F.Select( float.IsNaN ).ToArray() ) },
["Clip"] = ( n, i, c ) => new[] { Clip( n, i ) },
["Cast"] = ( n, i, c ) => new[] { Cast( i[0], (OnnxType)n.GetInt( "to", 1 ) ) },
["CastLike"] = ( n, i, c ) => new[] { Cast( i[0], i[1].Type ) },
["Shape"] = ( n, i, c ) => new[] { ShapeOf( n, i[0] ) },
["Size"] = ( n, i, c ) => new[] { Tensor.ScalarInt( i[0].Length ) },
["Reshape"] = ( n, i, c ) => new[] { Reshape( i[0], i[1].AsLongs(), n.GetInt( "allowzero", 0 ) == 1 ) },
["Flatten"] = ( n, i, c ) => new[] { Flatten( i[0], (int)n.GetInt( "axis", 1 ) ) },
["Unsqueeze"] = ( n, i, c ) => new[] { Unsqueeze( i[0], i.Length > 1 ? i[1].AsLongs() : n.GetInts( "axes" ) ) },
["Squeeze"] = ( n, i, c ) => new[] { Squeeze( i[0], i.Length > 1 && i[1] is not null ? i[1].AsLongs() : n.GetInts( "axes" ) ) },
["Transpose"] = ( n, i, c ) => new[] { FastTranspose( i[0], n.GetInts( "perm" ) ) },
["Concat"] = ( n, i, c ) => new[] { Concat( i.Where( t => t is not null ).ToArray(), (int)n.GetInt( "axis", 0 ) ) },
["Split"] = ( n, i, c ) => Split( n, i ),
["Slice"] = ( n, i, c ) => new[] { FastSlice( i ) ?? Slice( i ) },
["Gather"] = ( n, i, c ) => new[] { Gather( i[0], i[1], (int)n.GetInt( "axis", 0 ) ) },
["GatherElements"] = ( n, i, c ) => new[] { GatherElements( i[0], i[1], (int)n.GetInt( "axis", 0 ) ) },
["Expand"] = ( n, i, c ) => new[] { Expand( i[0], i[1].AsLongs() ) },
["Tile"] = ( n, i, c ) => new[] { Tile( i[0], i[1].AsLongs() ) },
["Range"] = ( n, i, c ) => new[] { Range( i[0], i[1], i[2] ) },
["Constant"] = ( n, i, c ) => new[] { ConstantOf( n ) },
["ConstantOfShape"] = ( n, i, c ) => new[] { ConstantOfShape( n, i[0].AsLongs() ) },
["Trilu"] = ( n, i, c ) => new[] { Trilu( i[0], i.Length > 1 && i[1] is not null ? i[1].AsLongs()[0] : 0, n.GetInt( "upper", 1 ) == 1 ) },
["CumSum"] = ( n, i, c ) => new[] { CumSum( i[0], (int)i[1].AsLongs()[0], n.GetInt( "exclusive", 0 ) == 1, n.GetInt( "reverse", 0 ) == 1 ) },
["ReduceMean"] = ( n, i, c ) => new[] { Reduce( n, i, ReduceOp.Mean ) },
["ReduceSum"] = ( n, i, c ) => new[] { Reduce( n, i, ReduceOp.Sum ) },
["ReduceMax"] = ( n, i, c ) => new[] { Reduce( n, i, ReduceOp.Max ) },
["ReduceMin"] = ( n, i, c ) => new[] { Reduce( n, i, ReduceOp.Min ) },
["ReduceProd"] = ( n, i, c ) => new[] { Reduce( n, i, ReduceOp.Prod ) },
["ReduceL2"] = ( n, i, c ) => new[] { Reduce( n, i, ReduceOp.L2 ) },
["ReduceSumSquare"] = ( n, i, c ) => new[] { Reduce( n, i, ReduceOp.SumSquare ) },
["ArgMax"] = ( n, i, c ) => new[] { ArgMax( i[0], (int)n.GetInt( "axis", 0 ), n.GetInt( "keepdims", 1 ) == 1 ) },
["Softmax"] = ( n, i, c ) => new[] { Softmax( i[0], (int)n.GetInt( "axis", -1 ), c, log: false ) },
["LogSoftmax"] = ( n, i, c ) => new[] { Softmax( i[0], (int)n.GetInt( "axis", -1 ), c, log: true ) },
["LayerNormalization"] = ( n, i, c ) => new[] { LayerNorm( i[0], i[1], i.Length > 2 ? i[2] : null, (int)n.GetInt( "axis", -1 ), n.GetFloat( "epsilon", 1e-5f ), c ) },
["SimplifiedLayerNormalization"] = ( n, i, c ) => new[] { RmsNorm( i[0], i[1], (int)n.GetInt( "axis", -1 ), n.GetFloat( "epsilon", 1e-6f ), c ) },
["MatMul"] = ( n, i, c ) => new[] { MatMul( i[0], i[1], c ) },
// runtime fusions (see OnnxSession.Fuse)
["MatMulBias"] = ( n, i, c ) => new[] { MatMulBias( i[0], i[1], i[2], c ) },
["SiLU"] = ( n, i, c ) => new[] { FastKernels.SiLU( i[0], c ) },
["Rope"] = ( n, i, c ) => new[] { FastKernels.Rope( i[0], i[1], i[2], i[3], c ) },
["MulAdd"] = ( n, i, c ) => new[] { FastKernels.MulAdd( i[0], i[1], i[2], c )
?? FastBinary( FastBinary( i[0], i[1], BinOp.Mul, FastKernels.Op.Mul, c ), i[2], BinOp.Add, FastKernels.Op.Add, c ) },
["Gemm"] = ( n, i, c ) => new[] { Gemm( n, i, c ) },
};
public static bool Supports( string op ) => Kernels.ContainsKey( op );
public static IEnumerable<string> SupportedOps => Kernels.Keys;
public static Tensor[] Run( OnnxNode node, Tensor[] inputs, ExecContext ctx )
{
if ( !Kernels.TryGetValue( node.OpType, out var kernel ) )
throw new NotSupportedException( $"ONNX operator {node.OpType} is not supported by the managed runtime." );
return kernel( node, inputs, ctx );
}
// ======================================================================== elementwise
enum BinOp { Add, Sub, Mul, Div, Pow, Max, Min, Mod, FMod }
static Tensor FastBinary( Tensor a, Tensor b, BinOp op, FastKernels.Op fast, ExecContext ctx )
{
if ( a.IsFloat && b.IsFloat && (a.Length > 64 || b.Length > 64) )
{
var r = FastKernels.Binary( a, b, fast, ctx );
if ( r is not null ) return r;
}
return Binary( a, b, op );
}
static Tensor FastTranspose( Tensor t, long[] perm )
{
var p = perm is { Length: > 0 } ? perm.Select( x => (int)x ).ToArray() : Enumerable.Range( 0, t.Rank ).Reverse().ToArray();
if ( t.Rank >= 2 && p[^1] == t.Rank - 1 && !t.IsBool && t.Shape[^1] >= 4 ) return FastKernels.TransposeRuns( t, p );
return Transpose( t, perm );
}
static Tensor FastSlice( Tensor[] i )
{
var t = i[0];
var starts = i[1].AsLongs(); var ends = i[2].AsLongs();
if ( starts.Length != 1 ) return null;
var axis = i.Length > 3 && i[3] is not null ? (int)i[3].AsLongs()[0] : 0;
if ( axis < 0 ) axis += t.Rank;
if ( i.Length > 4 && i[4] is not null && i[4].AsLongs()[0] != 1 ) return null;
var dim = t.Shape[axis];
long s = starts[0], e = ends[0];
if ( s < 0 ) s += dim; if ( e < 0 ) e += dim;
s = Math.Clamp( s, 0, dim ); e = Math.Clamp( e, 0, dim );
return FastKernels.SliceAxis( t, axis, (int)s, (int)Math.Max( s, e ) );
}
enum CmpOp { Eq, Lt, Le, Gt, Ge }
public static int[] BroadcastShape( int[] a, int[] b )
{
var rank = Math.Max( a.Length, b.Length );
var r = new int[rank];
for ( var i = 0; i < rank; i++ )
{
var da = i - (rank - a.Length) >= 0 ? a[i - (rank - a.Length)] : 1;
var db = i - (rank - b.Length) >= 0 ? b[i - (rank - b.Length)] : 1;
if ( da != db && da != 1 && db != 1 )
throw new ArgumentException( $"Shapes [{string.Join( ",", a )}] and [{string.Join( ",", b )}] don't broadcast." );
r[i] = da == 1 ? db : da;
}
return r;
}
/// <summary>Element strides of <paramref name="shape"/> laid over a broadcast output shape (0 for broadcast dims).</summary>
static int[] BroadcastStrides( int[] shape, int[] outShape )
{
var s = new int[outShape.Length];
var strides = Tensor.Strides( shape );
var offset = outShape.Length - shape.Length;
for ( var i = 0; i < shape.Length; i++ ) s[offset + i] = shape[i] == 1 ? 0 : strides[i];
return s;
}
/// <summary>Calls <paramref name="body"/>(outIndex, aIndex, bIndex) for every output element.</summary>
static void ForEachBroadcast( int[] outShape, int[] sa, int[] sb, Action<int, int, int> body )
{
var n = Tensor.SizeOf( outShape );
var rank = outShape.Length;
var idx = new int[rank];
int ia = 0, ib = 0;
for ( var o = 0; o < n; o++ )
{
body( o, ia, ib );
for ( var d = rank - 1; d >= 0; d-- )
{
idx[d]++;
ia += sa[d]; ib += sb[d];
if ( idx[d] < outShape[d] ) break;
ia -= sa[d] * outShape[d]; ib -= sb[d] * outShape[d];
idx[d] = 0;
}
}
}
static Tensor Binary( Tensor a, Tensor b, BinOp op )
{
if ( a.IsInt && b.IsInt ) return BinaryInt( a, b, op );
var fa = a.AsFloats(); var fb = b.AsFloats();
var shape = BroadcastShape( a.Shape, b.Shape );
var result = ExecContext.AllocZeroed( Tensor.SizeOf( shape ) );
if ( a.Length == result.Length && b.Length == result.Length )
{
VecBinary( fa, 0, fb, 0, result, 0, result.Length, op, false );
}
else if ( a.Length == result.Length && IsSuffix( b.Shape, shape ) )
{
// b repeats every b.Length elements (bias-style broadcast)
var inner = b.Length;
for ( var o = 0; o < result.Length; o += inner ) VecBinary( fa, o, fb, 0, result, o, inner, op, false );
}
else if ( b.Length == 1 )
{
var s = fb[0];
for ( var o = 0; o < result.Length; o++ ) result[o] = Apply( fa[a.Length == 1 ? 0 : o], s, op );
}
else if ( a.Length == 1 )
{
var s = fa[0];
for ( var o = 0; o < result.Length; o++ ) result[o] = Apply( s, fb[o], op );
}
else
{
ForEachBroadcast( shape, BroadcastStrides( a.Shape, shape ), BroadcastStrides( b.Shape, shape ),
( o, ia, ib ) => result[o] = Apply( fa[ia], fb[ib], op ) );
}
return Tensor.Float( shape, result );
}
static bool IsSuffix( int[] small, int[] shape )
{
// small (ignoring leading 1s) equals the trailing dims of shape
var s = small.SkipWhile( d => d == 1 ).ToArray();
if ( s.Length > shape.Length ) return false;
for ( var i = 0; i < s.Length; i++ ) if ( s[^(i + 1)] != shape[^(i + 1)] ) return false;
return s.Length > 0;
}
static void VecBinary( float[] a, int ao, float[] b, int bo, float[] r, int ro, int count, BinOp op, bool _ )
{
var i = 0;
var w = Vector<float>.Count;
if ( op is BinOp.Add or BinOp.Sub or BinOp.Mul or BinOp.Div or BinOp.Max or BinOp.Min )
{
var va = MemoryMarshal.Cast<float, Vector<float>>( a.AsSpan( ao, count ) );
var vb = MemoryMarshal.Cast<float, Vector<float>>( b.AsSpan( bo, count ) );
var vr = MemoryMarshal.Cast<float, Vector<float>>( r.AsSpan( ro, count ) );
for ( var k = 0; k < vr.Length; k++ )
{
vr[k] = op switch
{
BinOp.Add => va[k] + vb[k],
BinOp.Sub => va[k] - vb[k],
BinOp.Mul => va[k] * vb[k],
BinOp.Div => va[k] / vb[k],
BinOp.Max => Vector.Max( va[k], vb[k] ),
_ => Vector.Min( va[k], vb[k] ),
};
}
i = vr.Length * w;
}
for ( ; i < count; i++ ) r[ro + i] = Apply( a[ao + i], b[bo + i], op );
}
static float Apply( float a, float b, BinOp op ) => op switch
{
BinOp.Add => a + b,
BinOp.Sub => a - b,
BinOp.Mul => a * b,
BinOp.Div => a / b,
BinOp.Pow => b == 2f ? a * a : b == 0.5f ? MathF.Sqrt( a ) : MathF.Pow( a, b ),
BinOp.Max => MathF.Max( a, b ),
BinOp.Min => MathF.Min( a, b ),
BinOp.FMod => a % b,
_ => a - MathF.Floor( a / b ) * b,
};
static Tensor BinaryInt( Tensor a, Tensor b, BinOp op )
{
var shape = BroadcastShape( a.Shape, b.Shape );
var r = new long[Tensor.SizeOf( shape )];
ForEachBroadcast( shape, BroadcastStrides( a.Shape, shape ), BroadcastStrides( b.Shape, shape ), ( o, ia, ib ) =>
{
long x = a.L[ia], y = b.L[ib];
r[o] = op switch
{
BinOp.Add => x + y,
BinOp.Sub => x - y,
BinOp.Mul => x * y,
BinOp.Div => y == 0 ? 0 : x / y,
BinOp.Pow => (long)Math.Pow( x, y ),
BinOp.Max => Math.Max( x, y ),
BinOp.Min => Math.Min( x, y ),
BinOp.FMod => y == 0 ? 0 : x % y,
_ => y == 0 ? 0 : ((x % y) + y) % y,
};
} );
return Tensor.Int64( shape, r );
}
static Tensor Compare( Tensor a, Tensor b, CmpOp op )
{
var shape = BroadcastShape( a.Shape, b.Shape );
var r = new bool[Tensor.SizeOf( shape )];
var sa = BroadcastStrides( a.Shape, shape ); var sb = BroadcastStrides( b.Shape, shape );
if ( a.IsInt && b.IsInt )
ForEachBroadcast( shape, sa, sb, ( o, ia, ib ) => r[o] = Cmp( a.L[ia].CompareTo( b.L[ib] ), op ) );
else if ( a.IsBool && b.IsBool )
ForEachBroadcast( shape, sa, sb, ( o, ia, ib ) => r[o] = op == CmpOp.Eq && a.B[ia] == b.B[ib] );
else
{
var fa = a.AsFloats(); var fb = b.AsFloats();
ForEachBroadcast( shape, sa, sb, ( o, ia, ib ) => r[o] = Cmp( fa[ia].CompareTo( fb[ib] ), op ) && !(op == CmpOp.Eq && float.IsNaN( fa[ia] )) );
}
return Tensor.Bool( shape, r );
}
static bool Cmp( int c, CmpOp op ) => op switch
{
CmpOp.Eq => c == 0,
CmpOp.Lt => c < 0,
CmpOp.Le => c <= 0,
CmpOp.Gt => c > 0,
_ => c >= 0,
};
static Tensor Logical( Tensor a, Tensor b, Func<bool, bool, bool> f )
{
var shape = BroadcastShape( a.Shape, b.Shape );
var r = new bool[Tensor.SizeOf( shape )];
ForEachBroadcast( shape, BroadcastStrides( a.Shape, shape ), BroadcastStrides( b.Shape, shape ), ( o, ia, ib ) => r[o] = f( a.B[ia], b.B[ib] ) );
return Tensor.Bool( shape, r );
}
static Tensor Where( Tensor cond, Tensor x, Tensor y )
{
var shape = BroadcastShape( BroadcastShape( cond.Shape, x.Shape ), y.Shape );
var n = Tensor.SizeOf( shape );
var sc = BroadcastStrides( cond.Shape, shape );
var sx = BroadcastStrides( x.Shape, shape );
var sy = BroadcastStrides( y.Shape, shape );
var rank = shape.Length;
var idx = new int[rank];
int ic = 0, ix = 0, iy = 0;
if ( x.IsInt && y.IsInt )
{
var r = new long[n];
for ( var o = 0; o < n; o++ ) { r[o] = cond.B[ic] ? x.L[ix] : y.L[iy]; Step( ref ic, ref ix, ref iy ); }
return Tensor.Int64( shape, r );
}
var fx = x.AsFloats(); var fy = y.AsFloats();
var rf = new float[n];
for ( var o = 0; o < n; o++ ) { rf[o] = cond.B[ic] ? fx[ix] : fy[iy]; Step( ref ic, ref ix, ref iy ); }
return Tensor.Float( shape, rf );
void Step( ref int a, ref int b, ref int c )
{
for ( var d = rank - 1; d >= 0; d-- )
{
idx[d]++; a += sc[d]; b += sx[d]; c += sy[d];
if ( idx[d] < shape[d] ) return;
a -= sc[d] * shape[d]; b -= sx[d] * shape[d]; c -= sy[d] * shape[d];
idx[d] = 0;
}
}
}
static Tensor Unary( Tensor t, Func<float, float> f )
{
var src = t.AsFloats();
var r = ExecContext.Alloc( src.Length );
if ( r.Length > 65536 ) Parallel.For( 0, (r.Length + 8191) / 8192, chunk =>
{
var end = Math.Min( r.Length, (chunk + 1) * 8192 );
for ( var i = chunk * 8192; i < end; i++ ) r[i] = f( src[i] );
} );
else for ( var i = 0; i < r.Length; i++ ) r[i] = f( src[i] );
return Tensor.Float( t.Shape, r );
}
/// <summary>Error function, |error| < 1.2e-7 (Numerical Recipes erfc).</summary>
public static float Erf( float xf )
{
double x = xf;
var z = Math.Abs( x );
var t = 1.0 / (1.0 + 0.5 * z);
var r = t * Math.Exp( -z * z - 1.26551223 + t * (1.00002368 + t * (0.37409196 + t * (0.09678418 + t * (-0.18628806 +
t * (0.27886807 + t * (-1.13520398 + t * (1.48851587 + t * (-0.82215223 + t * 0.17087277)))))))) );
return (float)(x >= 0 ? 1.0 - r : r - 1.0);
}
static Tensor Clip( OnnxNode n, Tensor[] i )
{
var lo = i.Length > 1 && i[1] is not null ? i[1].AsFloats()[0] : n.GetFloat( "min", float.NegativeInfinity );
var hi = i.Length > 2 && i[2] is not null ? i[2].AsFloats()[0] : n.GetFloat( "max", float.PositiveInfinity );
return Unary( i[0], v => Math.Clamp( v, lo, hi ) );
}
static Tensor Cast( Tensor t, OnnxType to )
{
switch ( to )
{
case OnnxType.Float: case OnnxType.Double: case OnnxType.Float16: case OnnxType.BFloat16:
if ( to == OnnxType.Float16 && t.IsFloat ) return Tensor.Float( t.Shape, t.F.Select( v => (float)(Half)v ).ToArray() );
return t.IsFloat ? t : Tensor.Float( t.Shape, t.AsFloats() );
case OnnxType.Bool:
return t.IsBool ? t : Tensor.Bool( t.Shape, t.IsInt ? t.L.Select( v => v != 0 ).ToArray() : t.F.Select( v => v != 0 ).ToArray() );
default:
return t.IsInt ? t : Tensor.Int64( t.Shape, t.IsFloat ? t.F.Select( v => (long)v ).ToArray() : t.B.Select( v => v ? 1L : 0L ).ToArray() );
}
}
// ======================================================================== shape ops
static Tensor ShapeOf( OnnxNode n, Tensor t )
{
var rank = t.Rank;
var start = (int)n.GetInt( "start", 0 );
var end = (int)n.GetInt( "end", rank );
if ( start < 0 ) start += rank;
if ( end < 0 ) end += rank;
start = Math.Clamp( start, 0, rank ); end = Math.Clamp( end, 0, rank );
return Tensor.Vector( t.Shape.Skip( start ).Take( Math.Max( 0, end - start ) ).Select( d => (long)d ).ToArray() );
}
public static Tensor Reshape( Tensor t, long[] target, bool allowZero )
{
var shape = new int[target.Length];
var infer = -1;
long known = 1;
for ( var i = 0; i < target.Length; i++ )
{
var d = target[i];
if ( d == 0 && !allowZero ) d = t.Shape[i];
if ( d == -1 ) { infer = i; shape[i] = 1; continue; }
shape[i] = (int)d;
known *= d;
}
if ( infer >= 0 ) shape[infer] = known == 0 ? 0 : (int)(t.Length / known);
return t.WithShape( shape );
}
static Tensor Flatten( Tensor t, int axis )
{
if ( axis < 0 ) axis += t.Rank;
var a = t.Shape.Take( axis ).Aggregate( 1, ( x, y ) => x * y );
return t.WithShape( new[] { a, t.Length / Math.Max( 1, a ) } );
}
static Tensor Unsqueeze( Tensor t, long[] axes )
{
var rank = t.Rank + axes.Length;
var norm = axes.Select( a => (int)(a < 0 ? a + rank : a) ).OrderBy( a => a ).ToArray();
var shape = new List<int>( t.Shape );
foreach ( var a in norm ) shape.Insert( a, 1 );
return t.WithShape( shape.ToArray() );
}
static Tensor Squeeze( Tensor t, long[] axes )
{
if ( axes is null || axes.Length == 0 ) return t.WithShape( t.Shape.Where( d => d != 1 ).ToArray() );
var set = axes.Select( a => (int)(a < 0 ? a + t.Rank : a) ).ToHashSet();
return t.WithShape( t.Shape.Where( ( d, i ) => !set.Contains( i ) ).ToArray() );
}
public static Tensor Transpose( Tensor t, long[] perm )
{
var rank = t.Rank;
var p = perm is { Length: > 0 } ? perm.Select( x => (int)x ).ToArray() : Enumerable.Range( 0, rank ).Reverse().ToArray();
var outShape = p.Select( i => t.Shape[i] ).ToArray();
var inStrides = Tensor.Strides( t.Shape );
var src = p.Select( i => inStrides[i] ).ToArray();
var n = t.Length;
int[] map = new int[n];
var idx = new int[rank];
var off = 0;
for ( var o = 0; o < n; o++ )
{
map[o] = off;
for ( var d = rank - 1; d >= 0; d-- )
{
idx[d]++; off += src[d];
if ( idx[d] < outShape[d] ) break;
off -= src[d] * outShape[d]; idx[d] = 0;
}
}
return Gathered( t, outShape, map );
}
/// <summary>New tensor whose element o is the source element map[o].</summary>
static Tensor Gathered( Tensor t, int[] shape, int[] map )
{
if ( t.IsFloat ) { var r = new float[map.Length]; for ( var i = 0; i < r.Length; i++ ) r[i] = t.F[map[i]]; return Tensor.Float( shape, r ); }
if ( t.IsInt ) { var r = new long[map.Length]; for ( var i = 0; i < r.Length; i++ ) r[i] = t.L[map[i]]; return Tensor.Int64( shape, r ); }
var b = new bool[map.Length]; for ( var i = 0; i < b.Length; i++ ) b[i] = t.B[map[i]]; return Tensor.Bool( shape, b );
}
public static Tensor Concat( Tensor[] parts, int axis )
{
var rank = parts[0].Rank;
if ( axis < 0 ) axis += rank;
var shape = (int[])parts[0].Shape.Clone();
shape[axis] = parts.Sum( p => p.Shape[axis] );
var outer = shape.Take( axis ).Aggregate( 1, ( a, b ) => a * b );
var innerOut = Tensor.SizeOf( shape ) / Math.Max( 1, outer );
var isFloat = parts.Any( p => p.IsFloat );
var isInt = !isFloat && parts.Any( p => p.IsInt );
var f = isFloat ? ExecContext.AllocZeroed( Tensor.SizeOf( shape ) ) : null;
var l = isInt ? new long[Tensor.SizeOf( shape )] : null;
var bo = !isFloat && !isInt ? new bool[Tensor.SizeOf( shape )] : null;
var offset = 0;
foreach ( var p in parts )
{
var inner = p.Length / Math.Max( 1, outer );
for ( var o = 0; o < outer; o++ )
{
if ( f is not null ) Array.Copy( p.AsFloats(), o * inner, f, o * innerOut + offset, inner );
else if ( l is not null ) Array.Copy( p.AsLongs(), o * inner, l, o * innerOut + offset, inner );
else Array.Copy( p.B, o * inner, bo, o * innerOut + offset, inner );
}
offset += inner;
}
return f is not null ? Tensor.Float( shape, f ) : l is not null ? Tensor.Int64( shape, l ) : Tensor.Bool( shape, bo );
}
static Tensor[] Split( OnnxNode n, Tensor[] i )
{
var t = i[0];
var axis = (int)n.GetInt( "axis", 0 );
if ( axis < 0 ) axis += t.Rank;
long[] sizes = i.Length > 1 && i[1] is not null ? i[1].AsLongs() : n.GetInts( "split" );
var count = sizes?.Length ?? (int)n.GetInt( "num_outputs", n.Outputs.Length );
if ( sizes is null || sizes.Length == 0 )
{
var dim = t.Shape[axis];
var each = (dim + count - 1) / count;
sizes = Enumerable.Range( 0, count ).Select( k => (long)Math.Min( each, dim - k * each ) ).ToArray();
}
var results = new Tensor[sizes.Length];
long start = 0;
for ( var k = 0; k < sizes.Length; k++ )
{
results[k] = FastKernels.SliceAxis( t, axis, (int)start, (int)(start + sizes[k]) );
start += sizes[k];
}
return results;
}
static Tensor SliceAxis( Tensor t, int axis, int start, int end )
{
var starts = new long[t.Rank]; var ends = t.Shape.Select( d => (long)d ).ToArray();
starts[axis] = start; ends[axis] = end;
return SliceCore( t, starts, ends, Enumerable.Repeat( 1L, t.Rank ).ToArray() );
}
static Tensor Slice( Tensor[] i )
{
var t = i[0];
var startsIn = i[1].AsLongs(); var endsIn = i[2].AsLongs();
var axes = i.Length > 3 && i[3] is not null ? i[3].AsLongs() : Enumerable.Range( 0, startsIn.Length ).Select( x => (long)x ).ToArray();
var stepsIn = i.Length > 4 && i[4] is not null ? i[4].AsLongs() : Enumerable.Repeat( 1L, startsIn.Length ).ToArray();
var starts = new long[t.Rank]; var ends = t.Shape.Select( d => (long)d ).ToArray(); var steps = Enumerable.Repeat( 1L, t.Rank ).ToArray();
for ( var k = 0; k < axes.Length; k++ )
{
var a = (int)(axes[k] < 0 ? axes[k] + t.Rank : axes[k]);
var dim = t.Shape[a];
var step = stepsIn[k];
long s = startsIn[k], e = endsIn[k];
if ( s < 0 ) s += dim;
if ( e < 0 ) e += dim;
if ( step > 0 ) { s = Math.Clamp( s, 0, dim ); e = Math.Clamp( e, 0, dim ); }
else { s = Math.Clamp( s, 0, dim - 1 ); e = Math.Clamp( e, -1, dim - 1 ); }
starts[a] = s; ends[a] = e; steps[a] = step;
}
return SliceCore( t, starts, ends, steps );
}
static Tensor SliceCore( Tensor t, long[] starts, long[] ends, long[] steps )
{
var rank = t.Rank;
var shape = new int[rank];
for ( var d = 0; d < rank; d++ )
{
var len = steps[d] > 0 ? (ends[d] - starts[d] + steps[d] - 1) / steps[d] : (starts[d] - ends[d] - steps[d] - 1) / -steps[d];
shape[d] = (int)Math.Max( 0, len );
}
var n = Tensor.SizeOf( shape );
var strides = Tensor.Strides( t.Shape );
var map = new int[n];
var idx = new int[rank];
for ( var o = 0; o < n; o++ )
{
var off = 0;
for ( var d = 0; d < rank; d++ ) off += (int)(starts[d] + idx[d] * steps[d]) * strides[d];
map[o] = off;
for ( var d = rank - 1; d >= 0; d-- ) { if ( ++idx[d] < shape[d] ) break; idx[d] = 0; }
}
return Gathered( t, shape, map );
}
static Tensor Gather( Tensor data, Tensor indices, int axis )
{
if ( axis < 0 ) axis += data.Rank;
var idx = indices.AsLongs();
var dim = data.Shape[axis];
var outer = data.Shape.Take( axis ).Aggregate( 1, ( a, b ) => a * b );
var inner = data.Shape.Skip( axis + 1 ).Aggregate( 1, ( a, b ) => a * b );
var shape = data.Shape.Take( axis ).Concat( indices.Shape ).Concat( data.Shape.Skip( axis + 1 ) ).ToArray();
var map = new int[Tensor.SizeOf( shape )];
var o = 0;
for ( var a = 0; a < outer; a++ )
foreach ( var raw in idx )
{
var k = raw < 0 ? raw + dim : raw;
if ( k < 0 || k >= dim ) throw new IndexOutOfRangeException( $"Gather index {raw} out of range {dim}." );
var src = (a * dim + (int)k) * inner;
for ( var i = 0; i < inner; i++ ) map[o++] = src + i;
}
return Gathered( data, shape, map );
}
static Tensor GatherElements( Tensor data, Tensor indices, int axis )
{
if ( axis < 0 ) axis += data.Rank;
var idx = indices.AsLongs();
var ds = Tensor.Strides( data.Shape );
var map = new int[idx.Length];
var pos = new int[indices.Rank];
for ( var o = 0; o < idx.Length; o++ )
{
var off = 0;
for ( var d = 0; d < indices.Rank; d++ )
{
var v = d == axis ? (int)(idx[o] < 0 ? idx[o] + data.Shape[d] : idx[o]) : pos[d];
off += v * ds[d];
}
map[o] = off;
for ( var d = indices.Rank - 1; d >= 0; d-- ) { if ( ++pos[d] < indices.Shape[d] ) break; pos[d] = 0; }
}
return Gathered( data, indices.Shape, map );
}
static Tensor Expand( Tensor t, long[] target )
{
var shape = BroadcastShape( t.Shape, target.Select( d => (int)d ).ToArray() );
var strides = BroadcastStrides( t.Shape, shape );
var map = new int[Tensor.SizeOf( shape )];
ForEachBroadcast( shape, strides, strides, ( o, ia, _ ) => map[o] = ia );
return Gathered( t, shape, map );
}
static Tensor Tile( Tensor t, long[] repeats )
{
var shape = t.Shape.Select( ( d, i ) => d * (int)repeats[i] ).ToArray();
var strides = Tensor.Strides( t.Shape );
var map = new int[Tensor.SizeOf( shape )];
var idx = new int[shape.Length];
for ( var o = 0; o < map.Length; o++ )
{
var off = 0;
for ( var d = 0; d < shape.Length; d++ ) off += (idx[d] % t.Shape[d]) * strides[d];
map[o] = off;
for ( var d = shape.Length - 1; d >= 0; d-- ) { if ( ++idx[d] < shape[d] ) break; idx[d] = 0; }
}
return Gathered( t, shape, map );
}
static Tensor Range( Tensor start, Tensor limit, Tensor delta )
{
if ( start.IsInt && limit.IsInt && delta.IsInt )
{
long s = start.L[0], e = limit.L[0], d = delta.L[0];
var n = (int)Math.Max( 0, (long)Math.Ceiling( (e - s) / (double)d ) );
return Tensor.Int64( new[] { n }, Enumerable.Range( 0, n ).Select( i => s + i * d ).ToArray() );
}
float fs = start.AsFloats()[0], fe = limit.AsFloats()[0], fd = delta.AsFloats()[0];
var count = (int)Math.Max( 0, Math.Ceiling( (fe - fs) / fd ) );
return Tensor.Float( new[] { count }, Enumerable.Range( 0, count ).Select( i => fs + i * fd ).ToArray() );
}
static Tensor ConstantOf( OnnxNode n )
{
if ( n.Attributes.TryGetValue( "value", out var v ) && v.Tensor is { } init ) return Tensor.FromInitializer( init, "" );
if ( n.Attributes.TryGetValue( "value_float", out var vf ) ) return Tensor.Scalar( vf.Float ?? 0 );
if ( n.Attributes.TryGetValue( "value_int", out var vi ) ) return Tensor.ScalarInt( vi.Int ?? 0 );
if ( n.Attributes.TryGetValue( "value_ints", out var vis ) ) return Tensor.Vector( vis.Ints );
if ( n.Attributes.TryGetValue( "value_floats", out var vfs ) ) return Tensor.Float( new[] { vfs.Floats.Length }, vfs.Floats );
throw new NotSupportedException( "Constant node without a supported value." );
}
static Tensor ConstantOfShape( OnnxNode n, long[] dims )
{
var shape = dims.Select( d => (int)d ).ToArray();
if ( n.Attributes.TryGetValue( "value", out var v ) && v.Tensor is { } init )
{
var value = Tensor.FromInitializer( init, "" );
if ( value.IsInt ) return Tensor.Int64( shape, Enumerable.Repeat( value.L[0], Tensor.SizeOf( shape ) ).ToArray() );
if ( value.IsBool ) return Tensor.Bool( shape, Enumerable.Repeat( value.B[0], Tensor.SizeOf( shape ) ).ToArray() );
return Tensor.Float( shape, Enumerable.Repeat( value.F[0], Tensor.SizeOf( shape ) ).ToArray() );
}
return Tensor.Float( shape );
}
static Tensor Trilu( Tensor t, long k, bool upper )
{
var rows = t.Shape[^2]; var cols = t.Shape[^1];
var keep = new bool[t.Length];
for ( var o = 0; o < t.Length; o++ )
{
var r = (o / cols) % rows; var c = o % cols;
keep[o] = upper ? c - r >= k : c - r <= k;
}
if ( t.IsFloat ) return Tensor.Float( t.Shape, t.F.Select( ( v, o ) => keep[o] ? v : 0f ).ToArray() );
if ( t.IsInt ) return Tensor.Int64( t.Shape, t.L.Select( ( v, o ) => keep[o] ? v : 0L ).ToArray() );
return Tensor.Bool( t.Shape, t.B.Select( ( v, o ) => keep[o] && v ).ToArray() );
}
static Tensor CumSum( Tensor t, int axis, bool exclusive, bool reverse )
{
if ( axis < 0 ) axis += t.Rank;
var dim = t.Shape[axis];
var inner = t.Shape.Skip( axis + 1 ).Aggregate( 1, ( a, b ) => a * b );
var outer = t.Length / Math.Max( 1, dim * inner );
var src = t.AsFloats();
var r = ExecContext.AllocZeroed( t.Length );
for ( var o = 0; o < outer; o++ )
for ( var i = 0; i < inner; i++ )
{
var acc = 0f;
for ( var s = 0; s < dim; s++ )
{
var k = reverse ? dim - 1 - s : s;
var at = (o * dim + k) * inner + i;
if ( exclusive ) { r[at] = acc; acc += src[at]; }
else { acc += src[at]; r[at] = acc; }
}
}
return t.IsInt ? Tensor.Int64( t.Shape, r.Select( v => (long)v ).ToArray() ) : Tensor.Float( t.Shape, r );
}
// ======================================================================== reductions
enum ReduceOp { Mean, Sum, Max, Min, Prod, L2, SumSquare }
static Tensor Reduce( OnnxNode n, Tensor[] i, ReduceOp op )
{
var t = i[0];
long[] axes = i.Length > 1 && i[1] is not null ? i[1].AsLongs() : n.GetInts( "axes" );
var keep = n.GetInt( "keepdims", 1 ) == 1;
var noopEmpty = n.GetInt( "noop_with_empty_axes", 0 ) == 1;
if ( (axes is null || axes.Length == 0) && noopEmpty ) return t;
var set = axes is null || axes.Length == 0
? Enumerable.Range( 0, t.Rank ).ToHashSet()
: axes.Select( a => (int)(a < 0 ? a + t.Rank : a) ).ToHashSet();
var outShapeKeep = t.Shape.Select( ( d, k ) => set.Contains( k ) ? 1 : d ).ToArray();
var outN = Tensor.SizeOf( outShapeKeep );
var src = t.AsFloats();
var acc = new double[outN];
var init = op switch { ReduceOp.Max => double.NegativeInfinity, ReduceOp.Min => double.PositiveInfinity, ReduceOp.Prod => 1.0, _ => 0.0 };
Array.Fill( acc, init );
// fast path: reduce over the trailing axes only
var trailing = set.All( a => a >= t.Rank - set.Count );
if ( trailing )
{
var inner = t.Length / Math.Max( 1, outN );
for ( var o = 0; o < outN; o++ )
{
double a = init;
var baseIdx = o * inner;
for ( var k = 0; k < inner; k++ ) a = Accumulate( a, src[baseIdx + k], op );
acc[o] = a;
}
}
else
{
var outStrides = BroadcastStrides( outShapeKeep, t.Shape );
ForEachBroadcast( t.Shape, outStrides, outStrides, ( o, io, _ ) => acc[io] = Accumulate( acc[io], src[o], op ) );
}
var count = t.Length / Math.Max( 1, outN );
var r = new float[outN];
for ( var o = 0; o < outN; o++ )
r[o] = (float)(op switch { ReduceOp.Mean => acc[o] / count, ReduceOp.L2 => Math.Sqrt( acc[o] ), _ => acc[o] });
var shape = keep ? outShapeKeep : t.Shape.Where( ( d, k ) => !set.Contains( k ) ).ToArray();
return Tensor.Float( shape, r );
}
static double Accumulate( double a, float v, ReduceOp op ) => op switch
{
ReduceOp.Max => Math.Max( a, v ),
ReduceOp.Min => Math.Min( a, v ),
ReduceOp.Prod => a * v,
ReduceOp.L2 or ReduceOp.SumSquare => a + (double)v * v,
_ => a + v,
};
static Tensor ArgMax( Tensor t, int axis, bool keep )
{
if ( axis < 0 ) axis += t.Rank;
var dim = t.Shape[axis];
var inner = t.Shape.Skip( axis + 1 ).Aggregate( 1, ( a, b ) => a * b );
var outer = t.Length / Math.Max( 1, dim * inner );
var src = t.AsFloats();
var r = new long[outer * inner];
for ( var o = 0; o < outer; o++ )
for ( var i = 0; i < inner; i++ )
{
var best = 0; var bv = float.NegativeInfinity;
for ( var k = 0; k < dim; k++ ) { var v = src[(o * dim + k) * inner + i]; if ( v > bv ) { bv = v; best = k; } }
r[o * inner + i] = best;
}
var shape = keep ? t.Shape.Select( ( d, k ) => k == axis ? 1 : d ).ToArray() : t.Shape.Where( ( d, k ) => k != axis ).ToArray();
return Tensor.Int64( shape, r );
}
// ======================================================================== normalization / softmax
static Tensor Softmax( Tensor t, int axis, ExecContext ctx, bool log )
{
if ( axis < 0 ) axis += t.Rank;
var dim = t.Shape[axis];
var inner = t.Shape.Skip( axis + 1 ).Aggregate( 1, ( a, b ) => a * b );
var outer = t.Length / Math.Max( 1, dim * inner );
var src = t.F;
var r = ExecContext.AllocZeroed( t.Length );
void Row( int o, int i )
{
var max = float.NegativeInfinity;
for ( var k = 0; k < dim; k++ ) max = MathF.Max( max, src[(o * dim + k) * inner + i] );
if ( float.IsNegativeInfinity( max ) ) max = 0;
var sum = 0f;
for ( var k = 0; k < dim; k++ ) { var at = (o * dim + k) * inner + i; var e = MathF.Exp( src[at] - max ); r[at] = e; sum += e; }
if ( log ) { var ls = MathF.Log( sum ); for ( var k = 0; k < dim; k++ ) { var at = (o * dim + k) * inner + i; r[at] = src[at] - max - ls; } }
else { var inv = sum > 0 ? 1f / sum : 0f; for ( var k = 0; k < dim; k++ ) r[(o * dim + k) * inner + i] *= inv; }
}
var rows = outer * inner;
if ( (long)rows * dim > 32768 ) Parallel.For( 0, rows, ctx.Parallel, x => Row( x / inner, x % inner ) );
else for ( var x = 0; x < rows; x++ ) Row( x / inner, x % inner );
return Tensor.Float( t.Shape, r );
}
static Tensor LayerNorm( Tensor x, Tensor scale, Tensor bias, int axis, float eps, ExecContext ctx )
{
if ( axis < 0 ) axis += x.Rank;
var norm = x.Shape.Skip( axis ).Aggregate( 1, ( a, b ) => a * b );
var rows = x.Length / Math.Max( 1, norm );
var src = x.F; var r = ExecContext.Alloc( x.Length );
var g = scale.F; var b = bias?.F;
void Row( int row )
{
var o = row * norm;
double mean = 0; for ( var k = 0; k < norm; k++ ) mean += src[o + k]; mean /= norm;
double var = 0; for ( var k = 0; k < norm; k++ ) { var d = src[o + k] - mean; var += d * d; } var /= norm;
var inv = (float)(1.0 / Math.Sqrt( var + eps ));
var m = (float)mean;
for ( var k = 0; k < norm; k++ )
r[o + k] = (src[o + k] - m) * inv * g[g.Length == 1 ? 0 : k] + (b is null ? 0f : b[b.Length == 1 ? 0 : k]);
}
if ( rows > 64 ) Parallel.For( 0, rows, ctx.Parallel, Row ); else for ( var i = 0; i < rows; i++ ) Row( i );
return Tensor.Float( x.Shape, r );
}
static Tensor RmsNorm( Tensor x, Tensor scale, int axis, float eps, ExecContext ctx )
{
if ( axis < 0 ) axis += x.Rank;
var norm = x.Shape.Skip( axis ).Aggregate( 1, ( a, b ) => a * b );
var rows = x.Length / Math.Max( 1, norm );
var src = x.F; var r = ExecContext.Alloc( x.Length ); var g = scale.F;
Parallel.For( 0, rows, ctx.Parallel, row =>
{
var o = row * norm;
double ss = 0; for ( var k = 0; k < norm; k++ ) ss += (double)src[o + k] * src[o + k];
var inv = (float)(1.0 / Math.Sqrt( ss / norm + eps ));
for ( var k = 0; k < norm; k++ ) r[o + k] = src[o + k] * inv * g[k];
} );
return Tensor.Float( x.Shape, r );
}
// ======================================================================== matrix multiply
static Tensor MatMulBias( Tensor a, Tensor b, Tensor bias, ExecContext ctx )
{
if ( a.Rank < 2 || (long)b.Shape[0] * b.Shape[1] <= 64 * 64 )
return FastBinary( MatMul( a, b, ctx ), bias, BinOp.Add, FastKernels.Op.Add, ctx );
return MatMulPacked( a, b, ctx, bias.F );
}
/// <summary>A[..., K] x B[K, N] with B packed for the register-blocked GEMM (cached when B is a constant).</summary>
static Tensor MatMulPacked( Tensor a, Tensor b, ExecContext ctx, float[] bias = null )
{
int K = b.Shape[0], N = b.Shape[1];
var M = a.Length / Math.Max( 1, K );
FastKernels.PackedMatrix packed;
if ( ctx.Constants.Contains( b ) )
{
lock ( ctx.Packed )
if ( !ctx.Packed.TryGetValue( b, out packed ) ) ctx.Packed[b] = packed = FastKernels.Pack( b.F, K, N );
}
else packed = FastKernels.Pack( b.F, K, N );
var c = ExecContext.Alloc( M * N );
FastKernels.Gemm( a.F, 0, M, packed, c, 0, ctx, bias );
return Tensor.Float( a.Shape.Take( a.Rank - 1 ).Append( N ).ToArray(), c );
}
/// <summary>NumPy matmul with batch broadcasting.</summary>
public static Tensor MatMul( Tensor a, Tensor b, ExecContext ctx )
{
if ( b.Rank == 2 && a.Rank >= 2 && a.Shape[^1] == b.Shape[0] && a.IsFloat && b.IsFloat && (long)b.Shape[0] * b.Shape[1] > 64 * 64 ) return MatMulPacked( a, b, ctx );
var aShape = a.Rank == 1 ? new[] { 1, a.Shape[0] } : a.Shape;
var bShape = b.Rank == 1 ? new[] { b.Shape[0], 1 } : b.Shape;
int M = aShape[^2], K = aShape[^1], N = bShape[^1];
if ( bShape[^2] != K ) throw new ArgumentException( $"MatMul shapes [{string.Join( ",", a.Shape )}] x [{string.Join( ",", b.Shape )}] don't match." );
var batchA = aShape.Take( aShape.Length - 2 ).ToArray();
var batchB = bShape.Take( bShape.Length - 2 ).ToArray();
var batch = BroadcastShape( batchA, batchB );
var batches = Tensor.SizeOf( batch );
var outShape = batch.Concat( new[] { M, N } ).ToArray();
var c = ExecContext.AllocZeroed( batches * M * N );
var sa = BroadcastStrides( batchA, batch );
var sb = BroadcastStrides( batchB, batch );
var aOff = new int[batches]; var bOff = new int[batches];
{
var idx = new int[batch.Length];
for ( var p = 0; p < batches; p++ )
{
int oa = 0, ob = 0;
for ( var d = 0; d < batch.Length; d++ ) { oa += idx[d] * sa[d]; ob += idx[d] * sb[d]; }
aOff[p] = oa * M * K; bOff[p] = ob * K * N;
for ( var d = batch.Length - 1; d >= 0; d-- ) { if ( ++idx[d] < batch[d] ) break; idx[d] = 0; }
}
}
var af = a.AsFloats(); var bf = b.AsFloats();
var rows = batches * M;
var work = (long)rows * K * N;
// When B is shared by all batches, treat the batch as extra rows of A (one big GEMM).
void Row( int r )
{
var p = r / M; var i = r % M;
RowAxpy( af, aOff[p] + i * K, bf, bOff[p], c, r * N, K, N );
}
if ( work > 1 << 16 ) Parallel.For( 0, rows, ctx.Parallel, Row );
else for ( var r = 0; r < rows; r++ ) Row( r );
var shape = outShape;
if ( a.Rank == 1 ) shape = shape.Take( shape.Length - 2 ).Concat( new[] { N } ).ToArray();
if ( b.Rank == 1 ) shape = shape.Take( shape.Length - 1 ).ToArray();
return Tensor.Float( shape, c );
}
/// <summary>c[co..co+N] = sum_k a[ao+k] * b[bo + k*N .. +N].</summary>
static void RowAxpy( float[] a, int ao, float[] b, int bo, float[] c, int co, int K, int N )
{
var cs = c.AsSpan( co, N );
var cv = MemoryMarshal.Cast<float, Vector<float>>( cs );
var tail = cv.Length * Vector<float>.Count;
for ( var k = 0; k < K; k++ )
{
var s = a[ao + k];
if ( s == 0f ) continue;
var bs = b.AsSpan( bo + k * N, N );
var bv = MemoryMarshal.Cast<float, Vector<float>>( bs );
var sv = new Vector<float>( s );
for ( var j = 0; j < cv.Length; j++ ) cv[j] += bv[j] * sv;
for ( var j = tail; j < N; j++ ) cs[j] += s * bs[j];
}
}
static Tensor Gemm( OnnxNode n, Tensor[] i, ExecContext ctx )
{
var a = i[0]; var b = i[1];
var alpha = n.GetFloat( "alpha", 1f ); var beta = n.GetFloat( "beta", 1f );
if ( n.GetInt( "transA", 0 ) == 1 ) a = Transpose( a, new long[] { 1, 0 } );
if ( n.GetInt( "transB", 0 ) == 1 )
{
if ( !ctx.TransposedCache.TryGetValue( b, out var bt ) )
{
{
ctx.TransposedCache[b] = bt = Transpose( b, new long[] { 1, 0 } );
ctx.Pool.Release( bt.F ); // cached across runs: never recycled
}
if ( ctx.Constants.Contains( b ) ) ctx.Constants.Add( bt );
}
b = bt;
}
var y = MatMul( a, b, ctx );
if ( alpha != 1f ) y = Unary( y, v => v * alpha );
if ( i.Length > 2 && i[2] is not null )
{
var cterm = beta == 1f ? i[2] : Unary( i[2], v => v * beta );
y = Binary( y, cterm, BinOp.Add );
}
return y;
}
}