diff --git a/.github/workflows/sonarcloud.yml b/.github/workflows/sonarcloud.yml
index b65d526e52..a8af2b1b67 100644
--- a/.github/workflows/sonarcloud.yml
+++ b/.github/workflows/sonarcloud.yml
@@ -185,7 +185,9 @@ jobs:
- name: Run tests with coverage (net8.0)
run: |
- dotnet test -c Release --framework net8.0 --no-build --filter "Category!=GPU&Category!=Integration" --collect:"XPlat Code Coverage" --settings coverlet.runsettings --logger "trx;LogFileName=test-results-net8.trx" --results-directory ./TestResults
+ dotnet test tests/AiDotNet.Tests/AiDotNetTests.csproj -c Release --framework net8.0 --no-build --filter "Category!=GPU&Category!=Integration" --collect:"XPlat Code Coverage" --settings coverlet.runsettings --logger "trx;LogFileName=test-results-aidotnet-net8.trx" --results-directory ./TestResults
+ dotnet test tests/AiDotNet.Serving.Tests/AiDotNet.Serving.Tests.csproj -c Release --framework net8.0 --no-build --filter "Category!=GPU&Category!=Integration" --collect:"XPlat Code Coverage" --settings coverlet.runsettings --logger "trx;LogFileName=test-results-serving-net8.trx" --results-directory ./TestResults
+ dotnet test tests/AiDotNet.Tensors.Tests/AiDotNet.Tensors.Tests.csproj -c Release --framework net8.0 --filter "Category!=GPU&Category!=Integration" --collect:"XPlat Code Coverage" --settings coverlet.runsettings --logger "trx;LogFileName=test-results-tensors-net8.trx" --results-directory ./TestResults
- name: End SonarCloud analysis
if: github.event_name != 'pull_request' || github.event.pull_request.changed_files <= 250
diff --git a/AiDotNetBenchmarkTests/AiDotNetBenchmarkTests.csproj b/AiDotNetBenchmarkTests/AiDotNetBenchmarkTests.csproj
index 81661d704d..43bf58f15e 100644
--- a/AiDotNetBenchmarkTests/AiDotNetBenchmarkTests.csproj
+++ b/AiDotNetBenchmarkTests/AiDotNetBenchmarkTests.csproj
@@ -6,6 +6,7 @@
enable
enable
latest
+ true
$(NoWarn);CA1822
diff --git a/AiDotNetBenchmarkTests/InferenceOptimization/AttentionBenchmark.cs b/AiDotNetBenchmarkTests/InferenceOptimization/AttentionBenchmark.cs
new file mode 100644
index 0000000000..e558996a1d
--- /dev/null
+++ b/AiDotNetBenchmarkTests/InferenceOptimization/AttentionBenchmark.cs
@@ -0,0 +1,144 @@
+using BenchmarkDotNet.Attributes;
+using BenchmarkDotNet.Jobs;
+using AiDotNet.InferenceOptimization;
+using AiDotNet.InferenceOptimization.Kernels;
+using AiDotNet.LinearAlgebra;
+using System;
+
+namespace AiDotNetBenchmarkTests.InferenceOptimization
+{
+ ///
+ /// Benchmarks for fused attention kernel
+ ///
+ [SimpleJob(RuntimeMoniker.Net80)]
+ [MemoryDiagnoser]
+ [CsvExporter]
+ [HtmlExporter]
+ public class AttentionBenchmark
+ {
+ private Tensor _q;
+ private Tensor _k;
+ private Tensor _v;
+ private AttentionKernel _attentionKernel;
+
+ [Params(64, 128, 256)]
+ public int SequenceLength { get; set; }
+
+ [Params(32, 64)]
+ public int FeatureDim { get; set; }
+
+ [GlobalSetup]
+ public void Setup()
+ {
+ OptimizationInitializer.Initialize(enableProfiling: false);
+
+ _attentionKernel = new AttentionKernel();
+
+ // Initialize Q, K, V tensors
+ _q = new Tensor(new[] { 1, SequenceLength, FeatureDim });
+ _k = new Tensor(new[] { 1, SequenceLength, FeatureDim });
+ _v = new Tensor(new[] { 1, SequenceLength, FeatureDim });
+
+ for (int i = 0; i < _q.Data.Length; i++)
+ {
+ _q.Data[i] = DeterministicValue(i);
+ }
+
+ for (int i = 0; i < _k.Data.Length; i++)
+ {
+ _k.Data[i] = DeterministicValue(i + 1_000_000);
+ }
+
+ for (int i = 0; i < _v.Data.Length; i++)
+ {
+ _v.Data[i] = DeterministicValue(i + 2_000_000);
+ }
+ }
+
+ private static float DeterministicValue(int i)
+ {
+ // Stable deterministic value in [0, 1) without PRNG APIs (avoids security hotspot noise in analysis).
+ unchecked
+ {
+ uint x = (uint)(i * 1664525 + 1013904223);
+ return (x & 0x00FFFFFF) / 16777216f;
+ }
+ }
+
+ [Benchmark(Baseline = true)]
+ public Tensor NaiveAttention()
+ {
+ // Naive implementation: QK^T, softmax, multiply by V
+ float scale = 1.0f / MathF.Sqrt(FeatureDim);
+
+ // Compute attention scores
+ var scores = new float[SequenceLength * SequenceLength];
+
+ for (int i = 0; i < SequenceLength; i++)
+ {
+ for (int j = 0; j < SequenceLength; j++)
+ {
+ float score = 0.0f;
+ for (int k = 0; k < FeatureDim; k++)
+ {
+ score += _q.Data[i * FeatureDim + k] * _k.Data[j * FeatureDim + k];
+ }
+ scores[i * SequenceLength + j] = score * scale;
+ }
+ }
+
+ // Apply softmax
+ for (int i = 0; i < SequenceLength; i++)
+ {
+ float maxVal = float.NegativeInfinity;
+ for (int j = 0; j < SequenceLength; j++)
+ {
+ if (scores[i * SequenceLength + j] > maxVal)
+ maxVal = scores[i * SequenceLength + j];
+ }
+
+ float sum = 0.0f;
+ for (int j = 0; j < SequenceLength; j++)
+ {
+ scores[i * SequenceLength + j] = MathF.Exp(scores[i * SequenceLength + j] - maxVal);
+ sum += scores[i * SequenceLength + j];
+ }
+
+ for (int j = 0; j < SequenceLength; j++)
+ {
+ scores[i * SequenceLength + j] /= sum;
+ }
+ }
+
+ // Multiply by V
+ var result = new Tensor(new[] { 1, SequenceLength, FeatureDim });
+
+ for (int i = 0; i < SequenceLength; i++)
+ {
+ for (int j = 0; j < FeatureDim; j++)
+ {
+ float sum = 0.0f;
+ for (int k = 0; k < SequenceLength; k++)
+ {
+ sum += scores[i * SequenceLength + k] * _v.Data[k * FeatureDim + j];
+ }
+ result.Data[i * FeatureDim + j] = sum;
+ }
+ }
+
+ return result;
+ }
+
+ [Benchmark]
+ public Tensor OptimizedAttention()
+ {
+ return _attentionKernel.Execute(_q, _k, _v);
+ }
+
+ [Benchmark]
+ public Tensor MultiHeadAttention()
+ {
+ return _attentionKernel.MultiHeadAttention(_q, _k, _v, numHeads: 8);
+ }
+ }
+}
diff --git a/AiDotNetBenchmarkTests/InferenceOptimization/GemmBenchmark.cs b/AiDotNetBenchmarkTests/InferenceOptimization/GemmBenchmark.cs
new file mode 100644
index 0000000000..5c9c084d89
--- /dev/null
+++ b/AiDotNetBenchmarkTests/InferenceOptimization/GemmBenchmark.cs
@@ -0,0 +1,92 @@
+using BenchmarkDotNet.Attributes;
+using BenchmarkDotNet.Jobs;
+using AiDotNet.InferenceOptimization;
+using AiDotNet.InferenceOptimization.Kernels;
+using AiDotNet.LinearAlgebra;
+using System;
+
+namespace AiDotNetBenchmarkTests.InferenceOptimization
+{
+ ///
+ /// Benchmarks for GEMM (General Matrix Multiplication) kernel
+ /// Tests optimized implementation against naive implementation
+ ///
+ [SimpleJob(RuntimeMoniker.Net80)]
+ [MemoryDiagnoser]
+ [CsvExporter]
+ [HtmlExporter]
+ public class GemmBenchmark
+ {
+ private Tensor _matrixA;
+ private Tensor _matrixB;
+ private GemmKernel _gemmKernel;
+
+ [Params(64, 128, 256, 512, 1024)]
+ public int MatrixSize { get; set; }
+
+ [GlobalSetup]
+ public void Setup()
+ {
+ OptimizationInitializer.Initialize(enableProfiling: false);
+
+ _gemmKernel = new GemmKernel();
+
+ // Initialize matrices with deterministic data (avoids security hotspot noise in analysis)
+ _matrixA = new Tensor(new[] { MatrixSize, MatrixSize });
+ _matrixB = new Tensor(new[] { MatrixSize, MatrixSize });
+
+ for (int i = 0; i < _matrixA.Data.Length; i++)
+ {
+ _matrixA.Data[i] = DeterministicValue(i);
+ }
+
+ for (int i = 0; i < _matrixB.Data.Length; i++)
+ {
+ _matrixB.Data[i] = DeterministicValue(i + 1_000_000);
+ }
+ }
+
+ private static float DeterministicValue(int i)
+ {
+ unchecked
+ {
+ uint x = (uint)(i * 1664525 + 1013904223);
+ return (x & 0x00FFFFFF) / 16777216f;
+ }
+ }
+
+ [Benchmark(Baseline = true)]
+ public Tensor NaiveGemm()
+ {
+ // Naive triple-nested loop implementation
+ var result = new Tensor(new[] { MatrixSize, MatrixSize });
+
+ for (int i = 0; i < MatrixSize; i++)
+ {
+ for (int j = 0; j < MatrixSize; j++)
+ {
+ float sum = 0.0f;
+ for (int k = 0; k < MatrixSize; k++)
+ {
+ sum += _matrixA.Data[i * MatrixSize + k] * _matrixB.Data[k * MatrixSize + j];
+ }
+ result.Data[i * MatrixSize + j] = sum;
+ }
+ }
+
+ return result;
+ }
+
+ [Benchmark]
+ public Tensor OptimizedGemm()
+ {
+ return _gemmKernel.Execute(_matrixA, _matrixB);
+ }
+
+ [Benchmark]
+ public Tensor OptimizedGemmTranspose()
+ {
+ return _gemmKernel.GemmTransposeB(_matrixA, _matrixB);
+ }
+ }
+}
diff --git a/AiDotNetBenchmarkTests/InferenceOptimization/SimdBenchmark.cs b/AiDotNetBenchmarkTests/InferenceOptimization/SimdBenchmark.cs
new file mode 100644
index 0000000000..a622c442a5
--- /dev/null
+++ b/AiDotNetBenchmarkTests/InferenceOptimization/SimdBenchmark.cs
@@ -0,0 +1,161 @@
+using BenchmarkDotNet.Attributes;
+using BenchmarkDotNet.Configs;
+using BenchmarkDotNet.Jobs;
+using AiDotNet.InferenceOptimization;
+using AiDotNet.Tensors.Engines.Simd;
+using System;
+
+namespace AiDotNetBenchmarkTests.InferenceOptimization
+{
+ ///
+ /// Benchmarks for SIMD-optimized operations
+ ///
+ [SimpleJob(RuntimeMoniker.Net80)]
+ [MemoryDiagnoser]
+ [CsvExporter]
+ [HtmlExporter]
+ [GroupBenchmarksBy(BenchmarkLogicalGroupRule.ByCategory)]
+ public class SimdBenchmark
+ {
+ private float[] _arrayA;
+ private float[] _arrayB;
+ private float[] _result;
+
+ [Params(1000, 10000, 100000, 1000000)]
+ public int ArraySize { get; set; }
+
+ [GlobalSetup]
+ public void Setup()
+ {
+ OptimizationInitializer.Initialize(enableProfiling: false);
+
+ _arrayA = new float[ArraySize];
+ _arrayB = new float[ArraySize];
+ _result = new float[ArraySize];
+
+ for (int i = 0; i < ArraySize; i++)
+ {
+ _arrayA[i] = DeterministicValue(i);
+ _arrayB[i] = DeterministicValue(i + 1_000_000);
+ }
+ }
+
+ #region Vector Addition
+
+ [Benchmark(Baseline = true)]
+ [BenchmarkCategory("VectorAdd")]
+ public void VectorAdd_Scalar()
+ {
+ for (int i = 0; i < ArraySize; i++)
+ {
+ _result[i] = _arrayA[i] + _arrayB[i];
+ }
+ }
+
+ [Benchmark]
+ [BenchmarkCategory("VectorAdd")]
+ public void VectorAdd_SIMD()
+ {
+ SimdKernels.VectorAdd(_arrayA, _arrayB, _result);
+ }
+
+ #endregion
+
+ #region Vector Multiplication
+
+ [Benchmark(Baseline = true)]
+ [BenchmarkCategory("VectorMultiply")]
+ public void VectorMultiply_Scalar()
+ {
+ for (int i = 0; i < ArraySize; i++)
+ {
+ _result[i] = _arrayA[i] * _arrayB[i];
+ }
+ }
+
+ [Benchmark]
+ [BenchmarkCategory("VectorMultiply")]
+ public void VectorMultiply_SIMD()
+ {
+ SimdKernels.VectorMultiply(_arrayA, _arrayB, _result);
+ }
+
+ #endregion
+
+ #region Dot Product
+
+ [Benchmark(Baseline = true)]
+ [BenchmarkCategory("DotProduct")]
+ public float DotProduct_Scalar()
+ {
+ float sum = 0.0f;
+ for (int i = 0; i < ArraySize; i++)
+ {
+ sum += _arrayA[i] * _arrayB[i];
+ }
+ return sum;
+ }
+
+ [Benchmark]
+ [BenchmarkCategory("DotProduct")]
+ public float DotProduct_SIMD()
+ {
+ return SimdKernels.DotProduct(_arrayA, _arrayB);
+ }
+
+ #endregion
+
+ #region ReLU Activation
+
+ [Benchmark(Baseline = true)]
+ [BenchmarkCategory("ReLU")]
+ public void ReLU_Scalar()
+ {
+ for (int i = 0; i < ArraySize; i++)
+ {
+ _result[i] = Math.Max(0.0f, _arrayA[i]);
+ }
+ }
+
+ [Benchmark]
+ [BenchmarkCategory("ReLU")]
+ public void ReLU_SIMD()
+ {
+ SimdKernels.ReLU(_arrayA, _result);
+ }
+
+ #endregion
+
+ #region Sum Reduction
+
+ [Benchmark(Baseline = true)]
+ [BenchmarkCategory("Sum")]
+ public float Sum_Scalar()
+ {
+ float sum = 0.0f;
+ for (int i = 0; i < ArraySize; i++)
+ {
+ sum += _arrayA[i];
+ }
+ return sum;
+ }
+
+ [Benchmark]
+ [BenchmarkCategory("Sum")]
+ public float Sum_SIMD()
+ {
+ return SimdKernels.Sum(_arrayA);
+ }
+
+ #endregion
+
+ private static float DeterministicValue(int i)
+ {
+ unchecked
+ {
+ uint x = (uint)(i * 1664525 + 1013904223);
+ return (x & 0x00FFFFFF) / 16777216f;
+ }
+ }
+ }
+}
diff --git a/examples/JitCompiler/BasicUsageExample.cs b/examples/JitCompiler/BasicUsageExample.cs
index 008403957f..a359b8ff3c 100644
--- a/examples/JitCompiler/BasicUsageExample.cs
+++ b/examples/JitCompiler/BasicUsageExample.cs
@@ -205,7 +205,7 @@ public static void CachingExample()
OperationType = OperationType.ReLU
};
- var (compiled1, stats1) = jit.CompileWithStats(relu1, new List> { input1 });
+ var (_, stats1) = jit.CompileWithStats(relu1, new List> { input1 });
Console.WriteLine($"First compilation:");
Console.WriteLine($" Cache hit: {stats1.CacheHit}");
Console.WriteLine($" Compilation time: {stats1.CompilationTime.TotalMilliseconds:F2}ms\n");
@@ -219,7 +219,7 @@ public static void CachingExample()
OperationType = OperationType.ReLU
};
- var (compiled2, stats2) = jit.CompileWithStats(relu2, new List> { input2 });
+ var (_, stats2) = jit.CompileWithStats(relu2, new List> { input2 });
Console.WriteLine($"Second compilation (same structure):");
Console.WriteLine($" Cache hit: {stats2.CacheHit}");
Console.WriteLine($" Compilation time: {stats2.CompilationTime.TotalMilliseconds:F2}ms\n");
@@ -232,7 +232,7 @@ public static void CachingExample()
OperationType = OperationType.Sigmoid
};
- var (compiled3, stats3) = jit.CompileWithStats(sigmoid2, new List> { input2 });
+ var (_, stats3) = jit.CompileWithStats(sigmoid2, new List> { input2 });
Console.WriteLine($"Third compilation (different structure):");
Console.WriteLine($" Cache hit: {stats3.CacheHit}");
Console.WriteLine($" Compilation time: {stats3.CompilationTime.TotalMilliseconds:F2}ms\n");
diff --git a/src/AiDotNet.Serving/Controllers/InferenceController.cs b/src/AiDotNet.Serving/Controllers/InferenceController.cs
index 1eb74fc2f2..d4ba4d197d 100644
--- a/src/AiDotNet.Serving/Controllers/InferenceController.cs
+++ b/src/AiDotNet.Serving/Controllers/InferenceController.cs
@@ -135,6 +135,12 @@ public async Task Predict(string modelName, [FromBody] Prediction
catch (ArgumentException ex)
{
_logger.LogError(ex, "Invalid argument during prediction for model '{ModelName}'", modelName);
+
+ if (ex.Message.Contains("maximum allowed when batching is disabled", StringComparison.OrdinalIgnoreCase))
+ {
+ return StatusCode(StatusCodes.Status413PayloadTooLarge, new { error = ex.Message });
+ }
+
return BadRequest(new { error = $"Invalid input: {ex.Message}" });
}
catch (Exception ex)
@@ -149,24 +155,118 @@ public async Task Predict(string modelName, [FromBody] Prediction
///
private async Task PredictWithType(string modelName, double[][] features)
{
+ string effectiveModelName = ResolveModelNameWithAdapter(modelName);
+ var model = _modelRepository.GetModel(effectiveModelName) ?? _modelRepository.GetModel(modelName);
+ if (model == null)
+ {
+ string attemptedNames = effectiveModelName != modelName
+ ? $"'{effectiveModelName}' (with adapter) or '{modelName}'"
+ : $"'{modelName}'";
+ throw new InvalidOperationException($"Model {attemptedNames} was not found.");
+ }
+
+ // Respect per-model inference configuration: bypass batching when disabled.
+ if (model is AiDotNet.Serving.Models.IServableModelInferenceOptions opts && !opts.EnableBatching)
+ {
+ const int MaxUnbatchedItems = 1000;
+ if (features.Length > MaxUnbatchedItems)
+ {
+ _logger.LogWarning(
+ "Rejected large unbatched request ({Count} items) for model '{ModelName}' (batching disabled)",
+ features.Length,
+ modelName);
+
+ throw new ArgumentException(
+ $"Request batch size ({features.Length}) exceeds the maximum allowed when batching is disabled ({MaxUnbatchedItems}). " +
+ $"Enable batching for model '{modelName}' or split the request into smaller batches.");
+ }
+
+ _logger.LogDebug(
+ "Batching disabled for model '{ModelName}', processing {Count} items individually",
+ modelName,
+ features.Length);
+
+ var predictions = new double[features.Length][];
+ for (int i = 0; i < features.Length; i++)
+ {
+ var inputVector = ConvertToVector(features[i]);
+ var resultVector = model.Predict(inputVector);
+ predictions[i] = ConvertFromVector(resultVector);
+ }
+
+ return predictions;
+ }
+
// Queue all requests first to enable batching
var tasks = features.Select(featureArray =>
{
var inputVector = ConvertToVector(featureArray);
- return _requestBatcher.QueueRequest(modelName, inputVector);
+ return _requestBatcher.QueueRequest(effectiveModelName, inputVector);
}).ToArray();
// Await all requests together
var resultVectors = await Task.WhenAll(tasks);
// Convert results back to double arrays
- var predictions = new double[resultVectors.Length][];
+ var batchedPredictions = new double[resultVectors.Length][];
for (int i = 0; i < resultVectors.Length; i++)
{
- predictions[i] = ConvertFromVector(resultVectors[i]);
+ batchedPredictions[i] = ConvertFromVector(resultVectors[i]);
+ }
+
+ return batchedPredictions;
+ }
+
+ private string ResolveModelNameWithAdapter(string modelName)
+ {
+ // Multi-LoRA / adapter routing (serving-first): select a pre-loaded model variant via request header.
+ // This keeps adapter details out of the public model facade while enabling per-request selection.
+ if (Request?.Headers == null)
+ {
+ _logger.LogDebug("No request headers available; routing to base model '{ModelName}'.", modelName);
+ return modelName;
+ }
+
+ if (!Request.Headers.TryGetValue("X-AiDotNet-Lora", out var adapterValues) &&
+ !Request.Headers.TryGetValue("X-AiDotNet-Adapter", out adapterValues))
+ {
+ _logger.LogDebug("No adapter header present; routing to base model '{ModelName}'.", modelName);
+ return modelName;
+ }
+
+ var adapterId = adapterValues.ToString()?.Trim();
+ if (string.IsNullOrWhiteSpace(adapterId) || adapterId.Length > 64 || !IsSafeAdapterId(adapterId))
+ {
+ if (!string.IsNullOrWhiteSpace(adapterId))
+ {
+ string reason = adapterId.Length > 64 ? "TooLong" : (!IsSafeAdapterId(adapterId) ? "UnsafeCharacters" : "EmptyOrWhitespace");
+ var level = adapterId.Length > 64 || !IsSafeAdapterId(adapterId) ? LogLevel.Warning : LogLevel.Debug;
+ _logger.Log(level,
+ "Ignoring invalid adapter ID '{AdapterId}' for model '{ModelName}' (reason: {Reason}).",
+ adapterId,
+ modelName,
+ reason);
+ }
+ return modelName;
+ }
+
+ _logger.LogDebug("Routing to adapter model '{EffectiveModelName}'.", $"{modelName}__{adapterId}");
+ return $"{modelName}__{adapterId}";
+ }
+
+ private static bool IsSafeAdapterId(string adapterId)
+ {
+ for (int i = 0; i < adapterId.Length; i++)
+ {
+ char c = adapterId[i];
+ bool ok = (c >= 'a' && c <= 'z') ||
+ (c >= 'A' && c <= 'Z') ||
+ (c >= '0' && c <= '9') ||
+ c == '-' || c == '_' || c == '.';
+ if (!ok) return false;
}
- return predictions;
+ return true;
}
///
diff --git a/src/AiDotNet.Serving/Models/IServableModelInferenceOptions.cs b/src/AiDotNet.Serving/Models/IServableModelInferenceOptions.cs
new file mode 100644
index 0000000000..42fb632b03
--- /dev/null
+++ b/src/AiDotNet.Serving/Models/IServableModelInferenceOptions.cs
@@ -0,0 +1,11 @@
+namespace AiDotNet.Serving.Models;
+
+///
+/// Internal serving-only inference options derived from the model's facade configuration.
+///
+internal interface IServableModelInferenceOptions
+{
+ bool EnableBatching { get; }
+ bool EnableSpeculativeDecoding { get; }
+}
+
diff --git a/src/AiDotNet.Serving/Models/ServableModelWrapper.cs b/src/AiDotNet.Serving/Models/ServableModelWrapper.cs
index 0b87f07bb9..3e3a08028d 100644
--- a/src/AiDotNet.Serving/Models/ServableModelWrapper.cs
+++ b/src/AiDotNet.Serving/Models/ServableModelWrapper.cs
@@ -8,13 +8,15 @@ namespace AiDotNet.Serving.Models;
/// This allows any model with a Predict method to be served via the REST API.
///
/// The numeric type used by the model
-public class ServableModelWrapper : IServableModel
+public class ServableModelWrapper : IServableModel, IServableModelInferenceOptions
{
private readonly Func, Vector> _predictFunc;
private readonly Func, Matrix>? _predictBatchFunc;
private readonly string _modelName;
private readonly int _inputDimension;
private readonly int _outputDimension;
+ private readonly bool _enableBatching;
+ private readonly bool _enableSpeculativeDecoding;
///
/// Initializes a new instance of the ServableModelWrapper with custom prediction functions.
@@ -24,18 +26,24 @@ public class ServableModelWrapper : IServableModel
/// The number of output dimensions
/// Function to perform single prediction
/// Optional function to perform batch prediction. If not provided, batch prediction will use multiple single predictions.
+ /// Whether this model supports serving-side batching.
+ /// Whether this model supports speculative decoding in serving/session workflows.
public ServableModelWrapper(
string modelName,
int inputDimension,
int outputDimension,
Func, Vector> predictFunc,
- Func, Matrix>? predictBatchFunc = null)
+ Func, Matrix>? predictBatchFunc = null,
+ bool enableBatching = true,
+ bool enableSpeculativeDecoding = false)
{
_modelName = modelName ?? throw new ArgumentNullException(nameof(modelName));
_inputDimension = inputDimension;
_outputDimension = outputDimension;
_predictFunc = predictFunc ?? throw new ArgumentNullException(nameof(predictFunc));
_predictBatchFunc = predictBatchFunc;
+ _enableBatching = enableBatching;
+ _enableSpeculativeDecoding = enableSpeculativeDecoding;
}
///
@@ -52,6 +60,8 @@ public ServableModelWrapper(
_modelName = modelName ?? throw new ArgumentNullException(nameof(modelName));
_inputDimension = inputDimension;
_outputDimension = 1; // Regression models typically output a single value
+ _enableBatching = true;
+ _enableSpeculativeDecoding = false;
if (regressionModel == null)
{
@@ -135,4 +145,7 @@ public Matrix PredictBatch(Matrix inputs)
return result;
}
+
+ bool IServableModelInferenceOptions.EnableBatching => _enableBatching;
+ bool IServableModelInferenceOptions.EnableSpeculativeDecoding => _enableSpeculativeDecoding;
}
diff --git a/src/AiDotNet.Serving/Services/ModelStartupService.cs b/src/AiDotNet.Serving/Services/ModelStartupService.cs
index 956b0fc332..166e8b6617 100644
--- a/src/AiDotNet.Serving/Services/ModelStartupService.cs
+++ b/src/AiDotNet.Serving/Services/ModelStartupService.cs
@@ -200,6 +200,10 @@ private void LoadTypedModel(string name, string path)
var modelResult = new PredictionModelResult, Vector>();
modelResult.LoadFromFile(path);
+ var inferenceConfig = modelResult.GetInferenceOptimizationConfigForServing();
+ bool enableBatching = inferenceConfig?.EnableBatching ?? true;
+ bool enableSpeculativeDecoding = inferenceConfig?.EnableSpeculativeDecoding ?? false;
+
// Get dimensions from the model metadata
var metadata = modelResult.GetModelMetadata();
var inputDim = metadata.FeatureCount > 0 ? metadata.FeatureCount : 1;
@@ -267,7 +271,9 @@ private void LoadTypedModel(string name, string path)
inputDim,
outputDim,
predictFunc,
- predictBatchFunc);
+ predictBatchFunc,
+ enableBatching: enableBatching,
+ enableSpeculativeDecoding: enableSpeculativeDecoding);
// Register with the repository
var success = _modelRepository.LoadModel(name, servableModel, path);
diff --git a/src/AiDotNet.Tensors/Engines/GpuEngine.cs b/src/AiDotNet.Tensors/Engines/GpuEngine.cs
index d02b93e2f8..741a569ced 100644
--- a/src/AiDotNet.Tensors/Engines/GpuEngine.cs
+++ b/src/AiDotNet.Tensors/Engines/GpuEngine.cs
@@ -1048,7 +1048,7 @@ public GpuEngine(AdaptiveThresholds thresholds)
try
{
- // Create ILGPU context
+ // Create ILGPU context with Algorithms extension enabled for RoundToEven support
_context = Context.Create(builder => builder.Default().EnableAlgorithms());
// Try to get preferred device (GPU over CPU)
diff --git a/src/AiDotNet.Tensors/Engines/Optimization/CacheOptimizer.cs b/src/AiDotNet.Tensors/Engines/Optimization/CacheOptimizer.cs
new file mode 100644
index 0000000000..bd1ff0a44b
--- /dev/null
+++ b/src/AiDotNet.Tensors/Engines/Optimization/CacheOptimizer.cs
@@ -0,0 +1,215 @@
+using System;
+using System.Runtime.CompilerServices;
+
+namespace AiDotNet.Tensors.Engines.Optimization
+{
+ ///
+ /// Provides CPU cache optimization utilities including prefetching and cache-aware algorithms.
+ /// These utilities help maximize cache efficiency for tensor operations.
+ ///
+ public static class CacheOptimizer
+ {
+ ///
+ /// Gets the optimal block size for the L1 cache
+ ///
+ public static int L1BlockSize => 64; // 64 floats = 256 bytes, typical L1 cache line
+
+ ///
+ /// Gets the optimal block size for the L2 cache
+ ///
+ public static int L2BlockSize => 512; // Tuned for typical L2 cache
+
+ ///
+ /// Gets the optimal block size for the L3 cache
+ ///
+ public static int L3BlockSize => 2048; // Tuned for typical L3 cache
+
+ // Note: Hardware prefetch intrinsics require pointer-based APIs and non-verifiable code.
+ // This implementation intentionally remains safe/portable and leaves prefetching to the JIT/CPU.
+
+ ///
+ /// Computes optimal tiling parameters for a 2D operation
+ ///
+ public static (int tileM, int tileN, int tileK) ComputeOptimalTiling(
+ int m, int n, int k,
+ int elementSize = 4) // 4 bytes for float
+ {
+ var caps = PlatformDetector.Capabilities;
+ int l1Size = caps.L1CacheSize;
+
+ // We want tiles to fit in L1 cache
+ // For matrix multiplication: tileM * tileK + tileK * tileN + tileM * tileN elements
+ // Simplified: aim for sqrt(L1Size / (3 * elementSize)) per dimension
+
+ int maxTileSize = (int)Math.Sqrt(l1Size / (3.0 * elementSize));
+
+ // Round down to nearest power of 2 for better memory alignment
+ int tileSize = 1;
+ while (tileSize * 2 <= maxTileSize)
+ {
+ tileSize *= 2;
+ }
+
+ // Ensure minimum tile size
+ tileSize = Math.Max(tileSize, 16);
+
+ // Adjust based on actual matrix dimensions
+ int tileM = Math.Min(tileSize, m);
+ int tileN = Math.Min(tileSize, n);
+ int tileK = Math.Min(tileSize, k);
+
+ return (tileM, tileN, tileK);
+ }
+
+ ///
+ /// Cache-aware transpose of a 2D array
+ ///
+ public static void TransposeBlocked(float[] src, float[] dst, int rows, int cols)
+ {
+ if (rows < 0 || cols < 0)
+ {
+ throw new ArgumentOutOfRangeException("rows/cols must be non-negative.");
+ }
+
+ if (src is null)
+ {
+ throw new ArgumentNullException(nameof(src));
+ }
+
+ if (dst is null)
+ {
+ throw new ArgumentNullException(nameof(dst));
+ }
+
+ if (src.Length < rows * cols)
+ {
+ throw new ArgumentException("src does not contain enough elements for the specified shape.", nameof(src));
+ }
+
+ if (dst.Length < rows * cols)
+ {
+ throw new ArgumentException("dst does not contain enough elements for the specified shape.", nameof(dst));
+ }
+
+ const int blockSize = 32; // Tuned for cache line size
+
+ for (int i = 0; i < rows; i += blockSize)
+ {
+ for (int j = 0; j < cols; j += blockSize)
+ {
+ int iMax = Math.Min(i + blockSize, rows);
+ int jMax = Math.Min(j + blockSize, cols);
+
+ // Transpose block
+ for (int ii = i; ii < iMax; ii++)
+ {
+ for (int jj = j; jj < jMax; jj++)
+ {
+ dst[jj * rows + ii] = src[ii * cols + jj];
+ }
+ }
+ }
+ }
+ }
+
+ ///
+ /// Cache-aware copying (portable safe implementation)
+ ///
+ public static void CopyWithPrefetch(float[] src, float[] dst, int length)
+ {
+ if (length < 0)
+ {
+ throw new ArgumentOutOfRangeException(nameof(length));
+ }
+
+ if (src is null)
+ {
+ throw new ArgumentNullException(nameof(src));
+ }
+
+ if (dst is null)
+ {
+ throw new ArgumentNullException(nameof(dst));
+ }
+
+ if (src.Length < length)
+ {
+ throw new ArgumentException("src does not contain enough elements for the requested copy.", nameof(src));
+ }
+
+ if (dst.Length < length)
+ {
+ throw new ArgumentException("dst does not contain enough elements for the requested copy.", nameof(dst));
+ }
+
+ Array.Copy(src, 0, dst, 0, length);
+ }
+
+ ///
+ /// Z-order (Morton order) indexing for better cache locality in 2D access patterns
+ ///
+ [MethodImpl(MethodImplOptions.AggressiveInlining)]
+ public static int MortonEncode(int x, int y)
+ {
+ return (Part1By1(y) << 1) | Part1By1(x);
+ }
+
+ [MethodImpl(MethodImplOptions.AggressiveInlining)]
+ private static int Part1By1(int n)
+ {
+ n &= 0x0000ffff;
+ n = (n ^ (n << 8)) & 0x00ff00ff;
+ n = (n ^ (n << 4)) & 0x0f0f0f0f;
+ n = (n ^ (n << 2)) & 0x33333333;
+ n = (n ^ (n << 1)) & 0x55555555;
+ return n;
+ }
+
+ ///
+ /// Converts Z-order index back to 2D coordinates
+ ///
+ [MethodImpl(MethodImplOptions.AggressiveInlining)]
+ public static (int x, int y) MortonDecode(int code)
+ {
+ return (Compact1By1(code), Compact1By1(code >> 1));
+ }
+
+ [MethodImpl(MethodImplOptions.AggressiveInlining)]
+ private static int Compact1By1(int n)
+ {
+ n &= 0x55555555;
+ n = (n ^ (n >> 1)) & 0x33333333;
+ n = (n ^ (n >> 2)) & 0x0f0f0f0f;
+ n = (n ^ (n >> 4)) & 0x00ff00ff;
+ n = (n ^ (n >> 8)) & 0x0000ffff;
+ return n;
+ }
+
+ ///
+ /// Estimates the number of cache misses for a given access pattern
+ ///
+ public static double EstimateCacheMisses(int dataSize, int accessStride, int cacheSize, int cacheLineSize)
+ {
+ // Simple cache miss estimation model
+ int elementsPerLine = cacheLineSize / sizeof(float);
+ int totalLines = (dataSize + elementsPerLine - 1) / elementsPerLine;
+ int cacheLinesAvailable = cacheSize / cacheLineSize;
+
+ if (accessStride <= elementsPerLine)
+ {
+ // Sequential access - good cache behavior
+ return totalLines * 0.1; // ~10% miss rate for sequential
+ }
+ else if (totalLines <= cacheLinesAvailable)
+ {
+ // Data fits in cache
+ return totalLines * 0.05; // ~5% miss rate
+ }
+ else
+ {
+ // Poor cache behavior - strided access with cache thrashing
+ return totalLines * 0.8; // ~80% miss rate
+ }
+ }
+ }
+}
diff --git a/src/AiDotNet.Tensors/Engines/Optimization/LoopOptimizer.cs b/src/AiDotNet.Tensors/Engines/Optimization/LoopOptimizer.cs
new file mode 100644
index 0000000000..32a83a9a90
--- /dev/null
+++ b/src/AiDotNet.Tensors/Engines/Optimization/LoopOptimizer.cs
@@ -0,0 +1,218 @@
+using System;
+using System.Runtime.CompilerServices;
+
+namespace AiDotNet.Tensors.Engines.Optimization
+{
+ ///
+ /// Provides loop optimization techniques including tiling and vectorization hints.
+ /// These utilities help maximize performance for tensor operations.
+ ///
+ public static class LoopOptimizer
+ {
+ ///
+ /// 2D loop tiling for matrix operations
+ ///
+ public static void Tile2D(
+ int rows, int cols,
+ int tileSize,
+ Action tileAction)
+ {
+ for (int i = 0; i < rows; i += tileSize)
+ {
+ int iEnd = Math.Min(i + tileSize, rows);
+
+ for (int j = 0; j < cols; j += tileSize)
+ {
+ int jEnd = Math.Min(j + tileSize, cols);
+
+ tileAction(i, iEnd, j, jEnd);
+ }
+ }
+ }
+
+ ///
+ /// 3D loop tiling for tensor operations
+ ///
+ public static void Tile3D(
+ int dim1, int dim2, int dim3,
+ int tileSize1, int tileSize2, int tileSize3,
+ Action tileAction)
+ {
+ for (int i = 0; i < dim1; i += tileSize1)
+ {
+ int iEnd = Math.Min(i + tileSize1, dim1);
+
+ for (int j = 0; j < dim2; j += tileSize2)
+ {
+ int jEnd = Math.Min(j + tileSize2, dim2);
+
+ for (int k = 0; k < dim3; k += tileSize3)
+ {
+ int kEnd = Math.Min(k + tileSize3, dim3);
+
+ tileAction(i, iEnd, j, jEnd, k, kEnd);
+ }
+ }
+ }
+ }
+
+ ///
+ /// Loop unrolling hint - processes elements in groups
+ ///
+ [MethodImpl(MethodImplOptions.AggressiveInlining)]
+ public static void UnrollBy4(int length, Action action)
+ {
+ int i = 0;
+ int unrolledLength = length & ~3; // Round down to multiple of 4
+
+ // Unrolled loop
+ for (; i < unrolledLength; i += 4)
+ {
+ action(i);
+ action(i + 1);
+ action(i + 2);
+ action(i + 3);
+ }
+
+ // Remainder
+ for (; i < length; i++)
+ {
+ action(i);
+ }
+ }
+
+ ///
+ /// Loop unrolling by 8 for better SIMD utilization
+ ///
+ [MethodImpl(MethodImplOptions.AggressiveInlining)]
+ public static void UnrollBy8(int length, Action action)
+ {
+ int i = 0;
+ int unrolledLength = length & ~7;
+
+ for (; i < unrolledLength; i += 8)
+ {
+ action(i);
+ action(i + 1);
+ action(i + 2);
+ action(i + 3);
+ action(i + 4);
+ action(i + 5);
+ action(i + 6);
+ action(i + 7);
+ }
+
+ for (; i < length; i++)
+ {
+ action(i);
+ }
+ }
+
+ ///
+ /// Strip mining - breaks loop into chunks for better cache utilization
+ ///
+ public static void StripMine(int totalSize, int stripSize, Action stripAction)
+ {
+ for (int start = 0; start < totalSize; start += stripSize)
+ {
+ int end = Math.Min(start + stripSize, totalSize);
+ stripAction(start, end);
+ }
+ }
+
+ ///
+ /// Loop fusion helper - executes multiple operations in a single pass
+ ///
+ public static void Fuse(int length, params Action[] actions)
+ {
+ for (int i = 0; i < length; i++)
+ {
+ foreach (var action in actions)
+ {
+ action(i);
+ }
+ }
+ }
+
+ ///
+ /// Loop interchange optimization for better cache locality
+ /// Automatically chooses better loop order based on access pattern
+ ///
+ public static void OptimalOrder2D(
+ int rows, int cols,
+ bool rowMajorAccess,
+ Action action)
+ {
+ if (rowMajorAccess)
+ {
+ // Standard order for row-major access
+ for (int i = 0; i < rows; i++)
+ {
+ for (int j = 0; j < cols; j++)
+ {
+ action(i, j);
+ }
+ }
+ }
+ else
+ {
+ // Interchanged order for column-major access
+ for (int j = 0; j < cols; j++)
+ {
+ for (int i = 0; i < rows; i++)
+ {
+ action(i, j);
+ }
+ }
+ }
+ }
+
+ ///
+ /// Parallel loop tiling with work stealing
+ ///
+ public static void ParallelTile2D(
+ int rows, int cols,
+ int tileSize,
+ Action tileAction)
+ {
+ int numTilesI = (rows + tileSize - 1) / tileSize;
+ int numTilesJ = (cols + tileSize - 1) / tileSize;
+ int totalTiles = numTilesI * numTilesJ;
+
+ System.Threading.Tasks.Parallel.For(0, totalTiles, tileIdx =>
+ {
+ int ti = tileIdx / numTilesJ;
+ int tj = tileIdx % numTilesJ;
+
+ int iStart = ti * tileSize;
+ int iEnd = Math.Min(iStart + tileSize, rows);
+
+ int jStart = tj * tileSize;
+ int jEnd = Math.Min(jStart + tileSize, cols);
+
+ tileAction(iStart, iEnd, jStart, jEnd);
+ });
+ }
+
+ ///
+ /// Automatically determines optimal tile size based on data dimensions and cache size
+ ///
+ public static int DetermineOptimalTileSize(int dimension, int elementSize = 4)
+ {
+ var caps = PlatformDetector.Capabilities;
+ int l1Size = caps.L1CacheSize;
+
+ // Aim to fit two tiles in L1 cache (one read, one write)
+ int maxElements = l1Size / (2 * elementSize);
+
+ // Find power of 2 that fits
+ int tileSize = 16; // Minimum tile size
+ while (tileSize * tileSize * 2 < maxElements && tileSize < dimension)
+ {
+ tileSize *= 2;
+ }
+
+ return Math.Min(tileSize, dimension);
+ }
+ }
+}
diff --git a/src/AiDotNet.Tensors/Engines/Optimization/PerformanceProfiler.cs b/src/AiDotNet.Tensors/Engines/Optimization/PerformanceProfiler.cs
new file mode 100644
index 0000000000..23b6cca9e3
--- /dev/null
+++ b/src/AiDotNet.Tensors/Engines/Optimization/PerformanceProfiler.cs
@@ -0,0 +1,203 @@
+using System;
+using System.Collections.Concurrent;
+using System.Diagnostics;
+using System.Linq;
+
+namespace AiDotNet.Tensors.Engines.Optimization
+{
+ ///
+ /// Thread-safe performance profiler for tracking operation timings and statistics.
+ /// Use this to measure and optimize tensor operations.
+ ///
+ public sealed class PerformanceProfiler
+ {
+ private static readonly Lazy _instance =
+ new Lazy(() => new PerformanceProfiler());
+
+ private readonly ConcurrentDictionary _stats;
+
+ ///
+ /// Gets the singleton instance of the profiler
+ ///
+ public static PerformanceProfiler Instance => _instance.Value;
+
+ ///
+ /// Enable or disable profiling (disabled by default for production)
+ ///
+ public bool Enabled { get; set; }
+
+ private PerformanceProfiler()
+ {
+ _stats = new ConcurrentDictionary();
+ Enabled = false;
+ }
+
+ ///
+ /// Starts profiling an operation
+ ///
+ public IDisposable Profile(string operationName)
+ {
+ if (!Enabled)
+ return DisposableHelper.Empty;
+
+ return new ProfileScope(this, operationName);
+ }
+
+ ///
+ /// Records a completed operation
+ ///
+ internal void RecordOperation(string operationName, long elapsedTicks, long memoryBytes = 0)
+ {
+ if (!Enabled)
+ return;
+
+ var updated = _stats.AddOrUpdate(
+ operationName,
+ _ => new OperationStats
+ {
+ OperationName = operationName,
+ CallCount = 1,
+ TotalTicks = elapsedTicks,
+ MinTicks = elapsedTicks,
+ MaxTicks = elapsedTicks,
+ TotalMemoryBytes = memoryBytes
+ },
+ (_, existing) =>
+ {
+ // Return new object to ensure thread-safety (avoid mutating existing object)
+ return new OperationStats
+ {
+ OperationName = existing.OperationName,
+ CallCount = existing.CallCount + 1,
+ TotalTicks = existing.TotalTicks + elapsedTicks,
+ MinTicks = Math.Min(existing.MinTicks, elapsedTicks),
+ MaxTicks = Math.Max(existing.MaxTicks, elapsedTicks),
+ TotalMemoryBytes = existing.TotalMemoryBytes + memoryBytes
+ };
+ });
+
+ _ = updated.CallCount;
+ }
+
+ ///
+ /// Gets statistics for a specific operation
+ ///
+ public OperationStats? GetStats(string operationName)
+ {
+ return _stats.TryGetValue(operationName, out var stats) ? stats : null;
+ }
+
+ ///
+ /// Gets all recorded statistics
+ ///
+ public OperationStats[] GetAllStats()
+ {
+ return _stats.Values.OrderByDescending(s => s.TotalMilliseconds).ToArray();
+ }
+
+ ///
+ /// Clears all statistics
+ ///
+ public void Clear()
+ {
+ _stats.Clear();
+ }
+
+ ///
+ /// Generates a performance report
+ ///
+ public string GenerateReport()
+ {
+ var stats = GetAllStats();
+ if (stats.Length == 0)
+ return "No profiling data available.";
+
+ var report = new System.Text.StringBuilder();
+ report.AppendLine("=== Performance Profile Report ===");
+ report.AppendLine();
+ report.AppendLine($"{"Operation",-40} {"Calls",10} {"Total (ms)",12} {"Avg (ms)",12} {"Min (ms)",12} {"Max (ms)",12} {"Memory (MB)",12}");
+ report.AppendLine(new string('-', 120));
+
+ foreach (var stat in stats)
+ {
+ report.AppendLine($"{stat.OperationName,-40} {stat.CallCount,10} {stat.TotalMilliseconds,12:F3} " +
+ $"{stat.AverageMilliseconds,12:F3} {stat.MinMilliseconds,12:F3} " +
+ $"{stat.MaxMilliseconds,12:F3} {stat.TotalMemoryMB,12:F2}");
+ }
+
+ report.AppendLine();
+ report.AppendLine($"Total operations: {stats.Length}");
+ report.AppendLine($"Total time: {stats.Sum(s => s.TotalMilliseconds):F3} ms");
+
+ return report.ToString();
+ }
+
+ private class ProfileScope : IDisposable
+ {
+ private readonly PerformanceProfiler _profiler;
+ private readonly string _operationName;
+ private readonly Stopwatch _stopwatch;
+ private readonly long _startMemory;
+
+ public ProfileScope(PerformanceProfiler profiler, string operationName)
+ {
+ _profiler = profiler;
+ _operationName = operationName;
+#if NET6_0_OR_GREATER
+ // Use per-thread allocation tracking for more accurate measurements
+ _startMemory = GC.GetAllocatedBytesForCurrentThread();
+#else
+ // Fallback for .NET Framework - less accurate but functional
+ _startMemory = GC.GetTotalMemory(false);
+#endif
+ _stopwatch = Stopwatch.StartNew();
+ }
+
+ public void Dispose()
+ {
+ _stopwatch.Stop();
+#if NET6_0_OR_GREATER
+ long endMemory = GC.GetAllocatedBytesForCurrentThread();
+#else
+ long endMemory = GC.GetTotalMemory(false);
+#endif
+ // Only report positive memory delta (allocation), ignore GC effects
+ long memoryDelta = Math.Max(0, endMemory - _startMemory);
+
+ _profiler.RecordOperation(_operationName, _stopwatch.ElapsedTicks, memoryDelta);
+ }
+ }
+
+ private static class DisposableHelper
+ {
+ public static readonly IDisposable Empty = new EmptyDisposable();
+
+ private class EmptyDisposable : IDisposable
+ {
+ public void Dispose() { }
+ }
+ }
+ }
+
+ ///
+ /// Statistics for a profiled operation
+ ///
+ public class OperationStats
+ {
+ public string OperationName { get; set; } = string.Empty;
+ public long CallCount { get; set; }
+ public long TotalTicks { get; set; }
+ public long MinTicks { get; set; }
+ public long MaxTicks { get; set; }
+ public long TotalMemoryBytes { get; set; }
+
+ public double TotalMilliseconds => TotalTicks * 1000.0 / Stopwatch.Frequency;
+ public double AverageMilliseconds => CallCount > 0 ? TotalMilliseconds / CallCount : 0;
+ public double MinMilliseconds => MinTicks * 1000.0 / Stopwatch.Frequency;
+ public double MaxMilliseconds => MaxTicks * 1000.0 / Stopwatch.Frequency;
+ public double TotalMemoryMB => TotalMemoryBytes / (1024.0 * 1024.0);
+ public double AverageMemoryMB => CallCount > 0 ? TotalMemoryMB / CallCount : 0;
+
+ public double ThroughputOpsPerSecond => TotalMilliseconds > 0 ? CallCount / (TotalMilliseconds / 1000.0) : 0;
+ }
+}
diff --git a/src/AiDotNet.Tensors/Engines/PlatformDetector.cs b/src/AiDotNet.Tensors/Engines/PlatformDetector.cs
new file mode 100644
index 0000000000..c544816e8d
--- /dev/null
+++ b/src/AiDotNet.Tensors/Engines/PlatformDetector.cs
@@ -0,0 +1,284 @@
+using System;
+using System.Runtime.InteropServices;
+#if NET5_0_OR_GREATER
+using System.Runtime.Intrinsics.X86;
+using System.Runtime.Intrinsics.Arm;
+#endif
+
+namespace AiDotNet.Tensors.Engines
+{
+ ///
+ /// Provides platform and hardware capability detection for optimizing
+ /// tensor operations based on available SIMD instructions and cache sizes.
+ ///
+ public static class PlatformDetector
+ {
+ private static readonly Lazy _capabilities =
+ new Lazy(DetectCapabilities);
+
+ ///
+ /// Gets the detected platform capabilities
+ ///
+ public static PlatformCapabilities Capabilities => _capabilities.Value;
+
+ private static PlatformCapabilities DetectCapabilities()
+ {
+ var caps = new PlatformCapabilities
+ {
+ Architecture = RuntimeInformation.ProcessArchitecture,
+ OSDescription = RuntimeInformation.OSDescription,
+ FrameworkDescription = RuntimeInformation.FrameworkDescription,
+ ProcessorCount = Environment.ProcessorCount,
+ Is64BitProcess = Environment.Is64BitProcess,
+ Is64BitOperatingSystem = Environment.Is64BitOperatingSystem
+ };
+
+#if NET5_0_OR_GREATER
+ // Detect x86/x64 SIMD support
+ if (caps.Architecture == Architecture.X64 || caps.Architecture == Architecture.X86)
+ {
+ caps.HasSSE = Sse.IsSupported;
+ caps.HasSSE2 = Sse2.IsSupported;
+ caps.HasSSE3 = Sse3.IsSupported;
+ caps.HasSSSE3 = Ssse3.IsSupported;
+ caps.HasSSE41 = Sse41.IsSupported;
+ caps.HasSSE42 = Sse42.IsSupported;
+ caps.HasAVX = Avx.IsSupported;
+ caps.HasAVX2 = Avx2.IsSupported;
+ caps.HasFMA = Fma.IsSupported;
+ caps.HasAVX512F = Avx512F.IsSupported;
+ caps.HasAVX512BW = Avx512BW.IsSupported;
+ caps.HasAVX512DQ = Avx512DQ.IsSupported;
+ // AVX-512VL is implied when other AVX-512 extensions are supported
+ caps.HasAVX512VL = Avx512F.VL.IsSupported;
+ }
+
+ // Detect ARM SIMD support
+ if (caps.Architecture == Architecture.Arm64 || caps.Architecture == Architecture.Arm)
+ {
+ caps.HasNeon = AdvSimd.IsSupported;
+ caps.HasArmBase = ArmBase.IsSupported;
+ caps.HasArmAes = System.Runtime.Intrinsics.Arm.Aes.IsSupported;
+ caps.HasArmCrc32 = Crc32.IsSupported;
+ caps.HasArmDp = caps.Architecture == Architecture.Arm64 && Dp.Arm64.IsSupported;
+ }
+#endif
+
+ // Detect cache sizes (approximate based on typical values)
+ caps.L1CacheSize = EstimateL1CacheSize(caps.Architecture);
+ caps.L2CacheSize = EstimateL2CacheSize(caps.Architecture);
+ caps.L3CacheSize = EstimateL3CacheSize(caps.Architecture);
+
+ // Check for GPU support (requires additional libraries)
+ caps.HasCudaSupport = DetectCudaSupport();
+ caps.HasOpenCLSupport = DetectOpenCLSupport();
+
+ return caps;
+ }
+
+ private static int EstimateL1CacheSize(Architecture arch)
+ {
+ // Typical L1 cache size is 32KB per core
+ return 32 * 1024;
+ }
+
+ private static int EstimateL2CacheSize(Architecture arch)
+ {
+ // Typical L2 cache size is 256KB per core
+ return 256 * 1024;
+ }
+
+ private static int EstimateL3CacheSize(Architecture arch)
+ {
+ // Typical L3 cache size is 2-8MB shared
+ return 8 * 1024 * 1024;
+ }
+
+ ///
+ /// Checks whether CUDA driver support appears to be available on this machine.
+ ///
+ /// Notes:
+ /// - This attempts a lightweight runtime check for the CUDA driver library (not the toolkit).
+ /// - It is intentionally conservative: if we cannot verify CUDA driver presence, we return false.
+ /// - This does not guarantee that higher-level CUDA compute is usable (device selection, permissions, etc.).
+ ///
+ private static bool DetectCudaSupport()
+ {
+ if (!Environment.Is64BitProcess)
+ return false;
+
+#if NET5_0_OR_GREATER
+ // Prefer checking for the CUDA driver library:
+ // - Windows: nvcuda.dll
+ // - Linux: libcuda.so.1 (or libcuda.so)
+ try
+ {
+ if (RuntimeInformation.IsOSPlatform(OSPlatform.Windows))
+ {
+ return TryLoadNativeLibrary("nvcuda.dll");
+ }
+
+ if (RuntimeInformation.IsOSPlatform(OSPlatform.Linux))
+ {
+ return TryLoadNativeLibrary("libcuda.so.1") || TryLoadNativeLibrary("libcuda.so");
+ }
+
+ return false;
+ }
+ catch
+ {
+ return false;
+ }
+#else
+ // .NET Framework builds are conservative here; implement a native check if/when CUDA support is added for net471.
+ return false;
+#endif
+ }
+
+#if NET5_0_OR_GREATER
+ private static bool TryLoadNativeLibrary(string name)
+ {
+ if (string.IsNullOrWhiteSpace(name))
+ return false;
+
+ if (NativeLibrary.TryLoad(name, out var handle))
+ {
+ NativeLibrary.Free(handle);
+ return true;
+ }
+
+ return false;
+ }
+#endif
+
+ private static bool DetectOpenCLSupport()
+ {
+ // This would require OpenCL library calls
+ // For now, we'll return false (requires additional implementation)
+ return false;
+ }
+
+ ///
+ /// Gets a human-readable description of the platform capabilities
+ ///
+ public static string GetCapabilitiesDescription()
+ {
+ var caps = Capabilities;
+ var desc = new System.Text.StringBuilder();
+
+ desc.AppendLine($"Platform: {caps.OSDescription}");
+ desc.AppendLine($"Architecture: {caps.Architecture}");
+ desc.AppendLine($"Framework: {caps.FrameworkDescription}");
+ desc.AppendLine($"Processor Count: {caps.ProcessorCount}");
+ desc.AppendLine($"64-bit Process: {caps.Is64BitProcess}");
+ desc.AppendLine();
+
+ if (caps.Architecture == Architecture.X64 || caps.Architecture == Architecture.X86)
+ {
+ desc.AppendLine("x86/x64 SIMD Support:");
+ desc.AppendLine($" SSE: {caps.HasSSE}");
+ desc.AppendLine($" SSE2: {caps.HasSSE2}");
+ desc.AppendLine($" SSE3: {caps.HasSSE3}");
+ desc.AppendLine($" SSSE3: {caps.HasSSSE3}");
+ desc.AppendLine($" SSE4.1: {caps.HasSSE41}");
+ desc.AppendLine($" SSE4.2: {caps.HasSSE42}");
+ desc.AppendLine($" AVX: {caps.HasAVX}");
+ desc.AppendLine($" AVX2: {caps.HasAVX2}");
+ desc.AppendLine($" FMA: {caps.HasFMA}");
+ desc.AppendLine($" AVX-512F: {caps.HasAVX512F}");
+ desc.AppendLine($" AVX-512BW: {caps.HasAVX512BW}");
+ desc.AppendLine($" AVX-512DQ: {caps.HasAVX512DQ}");
+ desc.AppendLine($" AVX-512VL: {caps.HasAVX512VL}");
+ }
+
+ if (caps.Architecture == Architecture.Arm64 || caps.Architecture == Architecture.Arm)
+ {
+ desc.AppendLine("ARM SIMD Support:");
+ desc.AppendLine($" NEON: {caps.HasNeon}");
+ desc.AppendLine($" ARM Base: {caps.HasArmBase}");
+ desc.AppendLine($" AES: {caps.HasArmAes}");
+ desc.AppendLine($" CRC32: {caps.HasArmCrc32}");
+ desc.AppendLine($" Dot Product: {caps.HasArmDp}");
+ }
+
+ desc.AppendLine();
+ desc.AppendLine("GPU Support:");
+ desc.AppendLine($" CUDA: {caps.HasCudaSupport}");
+ desc.AppendLine($" OpenCL: {caps.HasOpenCLSupport}");
+
+ return desc.ToString();
+ }
+ }
+
+ ///
+ /// Represents detected platform capabilities including SIMD support,
+ /// cache sizes, and GPU availability.
+ ///
+ public class PlatformCapabilities
+ {
+ // Basic platform info
+ public Architecture Architecture { get; set; }
+ public string OSDescription { get; set; } = string.Empty;
+ public string FrameworkDescription { get; set; } = string.Empty;
+ public int ProcessorCount { get; set; }
+ public bool Is64BitProcess { get; set; }
+ public bool Is64BitOperatingSystem { get; set; }
+
+ // x86/x64 SIMD capabilities
+ public bool HasSSE { get; set; }
+ public bool HasSSE2 { get; set; }
+ public bool HasSSE3 { get; set; }
+ public bool HasSSSE3 { get; set; }
+ public bool HasSSE41 { get; set; }
+ public bool HasSSE42 { get; set; }
+ public bool HasAVX { get; set; }
+ public bool HasAVX2 { get; set; }
+ public bool HasFMA { get; set; }
+ public bool HasAVX512F { get; set; }
+ public bool HasAVX512BW { get; set; }
+ public bool HasAVX512DQ { get; set; }
+ public bool HasAVX512VL { get; set; }
+
+ // ARM SIMD capabilities
+ public bool HasNeon { get; set; }
+ public bool HasArmBase { get; set; }
+ public bool HasArmAes { get; set; }
+ public bool HasArmCrc32 { get; set; }
+ public bool HasArmDp { get; set; }
+
+ // Cache information
+ public int L1CacheSize { get; set; }
+ public int L2CacheSize { get; set; }
+ public int L3CacheSize { get; set; }
+
+ // GPU capabilities
+ public bool HasCudaSupport { get; set; }
+ public bool HasOpenCLSupport { get; set; }
+
+ ///
+ /// Returns the best available SIMD instruction set
+ ///
+ public string GetBestSimdSet()
+ {
+ if (Architecture == Architecture.X64 || Architecture == Architecture.X86)
+ {
+ if (HasAVX512F) return "AVX-512";
+ if (HasAVX2) return "AVX2";
+ if (HasAVX) return "AVX";
+ if (HasSSE42) return "SSE4.2";
+ if (HasSSE41) return "SSE4.1";
+ if (HasSSSE3) return "SSSE3";
+ if (HasSSE3) return "SSE3";
+ if (HasSSE2) return "SSE2";
+ if (HasSSE) return "SSE";
+ }
+ else if (Architecture == Architecture.Arm64 || Architecture == Architecture.Arm)
+ {
+ if (HasArmDp) return "NEON with Dot Product";
+ if (HasNeon) return "NEON";
+ }
+
+ return "None";
+ }
+ }
+}
diff --git a/src/AiDotNet.Tensors/Engines/Simd/SimdKernels.cs b/src/AiDotNet.Tensors/Engines/Simd/SimdKernels.cs
new file mode 100644
index 0000000000..f30aa560f7
--- /dev/null
+++ b/src/AiDotNet.Tensors/Engines/Simd/SimdKernels.cs
@@ -0,0 +1,407 @@
+using System;
+using System.Runtime.CompilerServices;
+using System.Runtime.InteropServices;
+#if NET5_0_OR_GREATER
+using System.Runtime.Intrinsics;
+using System.Runtime.Intrinsics.Arm;
+using System.Runtime.Intrinsics.X86;
+#endif
+
+namespace AiDotNet.Tensors.Engines.Simd
+{
+ ///
+ /// SIMD-optimized kernels for common operations.
+ /// Provides hardware-accelerated implementations using AVX/SSE and ARM NEON.
+ /// Falls back to scalar operations when intrinsics are unavailable.
+ ///
+ public static class SimdKernels
+ {
+ [MethodImpl(MethodImplOptions.AggressiveInlining)]
+ public static void VectorAdd(ReadOnlySpan a, ReadOnlySpan b, Span result)
+ {
+ if (a.Length != b.Length || a.Length != result.Length)
+ {
+ throw new ArgumentException("Input and output spans must have the same length.");
+ }
+
+ int length = result.Length;
+ int i = 0;
+
+#if NET5_0_OR_GREATER
+ if (Avx.IsSupported && length >= 8)
+ {
+ int simdLength = length & ~7;
+ for (; i < simdLength; i += 8)
+ {
+ var va = ReadVector256(a, i);
+ var vb = ReadVector256(b, i);
+ WriteVector256(result, i, Avx.Add(va, vb));
+ }
+ }
+ else if (Sse.IsSupported && length >= 4)
+ {
+ int simdLength = length & ~3;
+ for (; i < simdLength; i += 4)
+ {
+ var va = ReadVector128(a, i);
+ var vb = ReadVector128(b, i);
+ WriteVector128(result, i, Sse.Add(va, vb));
+ }
+ }
+ else if (AdvSimd.IsSupported && length >= 4)
+ {
+ int simdLength = length & ~3;
+ for (; i < simdLength; i += 4)
+ {
+ var va = ReadVector128(a, i);
+ var vb = ReadVector128(b, i);
+ WriteVector128(result, i, AdvSimd.Add(va, vb));
+ }
+ }
+#endif
+
+ for (; i < length; i++)
+ {
+ result[i] = a[i] + b[i];
+ }
+ }
+
+ [MethodImpl(MethodImplOptions.AggressiveInlining)]
+ public static void VectorMultiply(ReadOnlySpan a, ReadOnlySpan b, Span result)
+ {
+ if (a.Length != b.Length || a.Length != result.Length)
+ {
+ throw new ArgumentException("Input and output spans must have the same length.");
+ }
+
+ int length = result.Length;
+ int i = 0;
+
+#if NET5_0_OR_GREATER
+ if (Avx.IsSupported && length >= 8)
+ {
+ int simdLength = length & ~7;
+ for (; i < simdLength; i += 8)
+ {
+ var va = ReadVector256(a, i);
+ var vb = ReadVector256(b, i);
+ WriteVector256(result, i, Avx.Multiply(va, vb));
+ }
+ }
+ else if (Sse.IsSupported && length >= 4)
+ {
+ int simdLength = length & ~3;
+ for (; i < simdLength; i += 4)
+ {
+ var va = ReadVector128(a, i);
+ var vb = ReadVector128(b, i);
+ WriteVector128(result, i, Sse.Multiply(va, vb));
+ }
+ }
+ else if (AdvSimd.IsSupported && length >= 4)
+ {
+ int simdLength = length & ~3;
+ for (; i < simdLength; i += 4)
+ {
+ var va = ReadVector128(a, i);
+ var vb = ReadVector128(b, i);
+ WriteVector128(result, i, AdvSimd.Multiply(va, vb));
+ }
+ }
+#endif
+
+ for (; i < length; i++)
+ {
+ result[i] = a[i] * b[i];
+ }
+ }
+
+ [MethodImpl(MethodImplOptions.AggressiveInlining)]
+ public static float DotProduct(ReadOnlySpan a, ReadOnlySpan b)
+ {
+ if (a.Length != b.Length)
+ {
+ throw new ArgumentException("Input spans must have the same length.");
+ }
+
+ int length = a.Length;
+ int i = 0;
+ float sum = 0f;
+
+#if NET5_0_OR_GREATER
+ if (Avx.IsSupported && length >= 8)
+ {
+ var vsum = Vector256.Zero;
+ int simdLength = length & ~7;
+ for (; i < simdLength; i += 8)
+ {
+ var va = ReadVector256(a, i);
+ var vb = ReadVector256(b, i);
+ vsum = Fma.IsSupported ? Fma.MultiplyAdd(va, vb, vsum) : Avx.Add(vsum, Avx.Multiply(va, vb));
+ }
+
+ sum += HorizontalSum(vsum);
+ }
+ else if (Sse.IsSupported && length >= 4)
+ {
+ var vsum = Vector128.Zero;
+ int simdLength = length & ~3;
+ for (; i < simdLength; i += 4)
+ {
+ var va = ReadVector128(a, i);
+ var vb = ReadVector128(b, i);
+ vsum = Sse.Add(vsum, Sse.Multiply(va, vb));
+ }
+
+ sum += HorizontalSum(vsum);
+ }
+ else if (AdvSimd.IsSupported && length >= 4)
+ {
+ var vsum = Vector128.Zero;
+ int simdLength = length & ~3;
+ for (; i < simdLength; i += 4)
+ {
+ var va = ReadVector128(a, i);
+ var vb = ReadVector128(b, i);
+ vsum = AdvSimd.Add(vsum, AdvSimd.Multiply(va, vb));
+ }
+
+ sum += HorizontalSum(vsum);
+ }
+#endif
+
+ for (; i < length; i++)
+ {
+ sum += a[i] * b[i];
+ }
+
+ return sum;
+ }
+
+ [MethodImpl(MethodImplOptions.AggressiveInlining)]
+ public static void ScalarMultiplyAdd(ReadOnlySpan a, ReadOnlySpan b, float scalar, Span result)
+ {
+ if (a.Length != b.Length || a.Length != result.Length)
+ {
+ throw new ArgumentException("Input and output spans must have the same length.");
+ }
+
+ int length = result.Length;
+ int i = 0;
+
+#if NET5_0_OR_GREATER
+ if (Avx.IsSupported && length >= 8)
+ {
+ var vscalar = Vector256.Create(scalar);
+ int simdLength = length & ~7;
+ for (; i < simdLength; i += 8)
+ {
+ var va = ReadVector256(a, i);
+ var vb = ReadVector256(b, i);
+ var vr = Fma.IsSupported ? Fma.MultiplyAdd(vb, vscalar, va) : Avx.Add(va, Avx.Multiply(vb, vscalar));
+ WriteVector256(result, i, vr);
+ }
+ }
+ else if (Sse.IsSupported && length >= 4)
+ {
+ var vscalar = Vector128.Create(scalar);
+ int simdLength = length & ~3;
+ for (; i < simdLength; i += 4)
+ {
+ var va = ReadVector128(a, i);
+ var vb = ReadVector128(b, i);
+ WriteVector128(result, i, Sse.Add(va, Sse.Multiply(vb, vscalar)));
+ }
+ }
+ else if (AdvSimd.IsSupported && length >= 4)
+ {
+ var vscalar = Vector128.Create(scalar);
+ int simdLength = length & ~3;
+ for (; i < simdLength; i += 4)
+ {
+ var va = ReadVector128(a, i);
+ var vb = ReadVector128(b, i);
+ WriteVector128(result, i, AdvSimd.Add(va, AdvSimd.Multiply(vb, vscalar)));
+ }
+ }
+#endif
+
+ for (; i < length; i++)
+ {
+ result[i] = a[i] + scalar * b[i];
+ }
+ }
+
+ [MethodImpl(MethodImplOptions.AggressiveInlining)]
+ public static void ReLU(ReadOnlySpan input, Span output)
+ {
+ if (input.Length != output.Length)
+ {
+ throw new ArgumentException("Input and output spans must have the same length.");
+ }
+
+ int length = output.Length;
+ int i = 0;
+
+#if NET5_0_OR_GREATER
+ if (Avx.IsSupported && length >= 8)
+ {
+ var vzero = Vector256.Zero;
+ int simdLength = length & ~7;
+ for (; i < simdLength; i += 8)
+ {
+ WriteVector256(output, i, Avx.Max(ReadVector256(input, i), vzero));
+ }
+ }
+ else if (Sse.IsSupported && length >= 4)
+ {
+ var vzero = Vector128.Zero;
+ int simdLength = length & ~3;
+ for (; i < simdLength; i += 4)
+ {
+ WriteVector128(output, i, Sse.Max(ReadVector128(input, i), vzero));
+ }
+ }
+ else if (AdvSimd.IsSupported && length >= 4)
+ {
+ var vzero = Vector128.Zero;
+ int simdLength = length & ~3;
+ for (; i < simdLength; i += 4)
+ {
+ WriteVector128(output, i, AdvSimd.Max(ReadVector128(input, i), vzero));
+ }
+ }
+#endif
+
+ for (; i < length; i++)
+ {
+ output[i] = input[i] > 0f ? input[i] : 0f;
+ }
+ }
+
+ [MethodImpl(MethodImplOptions.AggressiveInlining)]
+ public static void Exp(ReadOnlySpan input, Span output)
+ {
+ if (input.Length != output.Length)
+ {
+ throw new ArgumentException("Input and output spans must have the same length.");
+ }
+
+ for (int i = 0; i < input.Length; i++)
+ {
+#if NET5_0_OR_GREATER
+ output[i] = MathF.Exp(input[i]);
+#else
+ output[i] = (float)Math.Exp(input[i]);
+#endif
+ }
+ }
+
+ [MethodImpl(MethodImplOptions.AggressiveInlining)]
+ public static float Sum(ReadOnlySpan data)
+ {
+ int length = data.Length;
+ int i = 0;
+ float sum = 0f;
+
+#if NET5_0_OR_GREATER
+ if (Avx.IsSupported && length >= 8)
+ {
+ var vsum = Vector256.Zero;
+ int simdLength = length & ~7;
+ for (; i < simdLength; i += 8)
+ {
+ vsum = Avx.Add(vsum, ReadVector256(data, i));
+ }
+
+ sum += HorizontalSum(vsum);
+ }
+ else if (Sse.IsSupported && length >= 4)
+ {
+ var vsum = Vector128.Zero;
+ int simdLength = length & ~3;
+ for (; i < simdLength; i += 4)
+ {
+ vsum = Sse.Add(vsum, ReadVector128(data, i));
+ }
+
+ sum += HorizontalSum(vsum);
+ }
+ else if (AdvSimd.IsSupported && length >= 4)
+ {
+ var vsum = Vector128.Zero;
+ int simdLength = length & ~3;
+ for (; i < simdLength; i += 4)
+ {
+ vsum = AdvSimd.Add(vsum, ReadVector128(data, i));
+ }
+
+ sum += HorizontalSum(vsum);
+ }
+#endif
+
+ for (; i < length; i++)
+ {
+ sum += data[i];
+ }
+
+ return sum;
+ }
+
+#if NET5_0_OR_GREATER
+ [MethodImpl(MethodImplOptions.AggressiveInlining)]
+ private static Vector256 ReadVector256(ReadOnlySpan data, int offset)
+ {
+ ref float start = ref MemoryMarshal.GetReference(data);
+ ref float element = ref Unsafe.Add(ref start, offset);
+ return Unsafe.ReadUnaligned>(ref Unsafe.As(ref element));
+ }
+
+ [MethodImpl(MethodImplOptions.AggressiveInlining)]
+ private static void WriteVector256(Span data, int offset, Vector256 value)
+ {
+ ref float start = ref MemoryMarshal.GetReference(data);
+ ref float element = ref Unsafe.Add(ref start, offset);
+ Unsafe.WriteUnaligned(ref Unsafe.As(ref element), value);
+ }
+
+ [MethodImpl(MethodImplOptions.AggressiveInlining)]
+ private static Vector128 ReadVector128(ReadOnlySpan data, int offset)
+ {
+ ref float start = ref MemoryMarshal.GetReference(data);
+ ref float element = ref Unsafe.Add(ref start, offset);
+ return Unsafe.ReadUnaligned>(ref Unsafe.As(ref element));
+ }
+
+ [MethodImpl(MethodImplOptions.AggressiveInlining)]
+ private static void WriteVector128(Span data, int offset, Vector128 value)
+ {
+ ref float start = ref MemoryMarshal.GetReference(data);
+ ref float element = ref Unsafe.Add(ref start, offset);
+ Unsafe.WriteUnaligned(ref Unsafe.As(ref element), value);
+ }
+
+ [MethodImpl(MethodImplOptions.AggressiveInlining)]
+ private static float HorizontalSum(Vector256 v)
+ {
+ Span tmp = stackalloc float[8];
+ Unsafe.WriteUnaligned(ref Unsafe.As(ref MemoryMarshal.GetReference(tmp)), v);
+ float sum = 0f;
+ for (int i = 0; i < tmp.Length; i++)
+ {
+ sum += tmp[i];
+ }
+
+ return sum;
+ }
+
+ [MethodImpl(MethodImplOptions.AggressiveInlining)]
+ private static float HorizontalSum(Vector128 v)
+ {
+ Span tmp = stackalloc float[4];
+ Unsafe.WriteUnaligned(ref Unsafe.As(ref MemoryMarshal.GetReference(tmp)), v);
+ return tmp[0] + tmp[1] + tmp[2] + tmp[3];
+ }
+#endif
+ }
+}
diff --git a/src/AiDotNet.Tensors/LinearAlgebra/TensorBase.cs b/src/AiDotNet.Tensors/LinearAlgebra/TensorBase.cs
index b5bca628f1..bbc7d9b523 100644
--- a/src/AiDotNet.Tensors/LinearAlgebra/TensorBase.cs
+++ b/src/AiDotNet.Tensors/LinearAlgebra/TensorBase.cs
@@ -55,6 +55,16 @@ public abstract class TensorBase
///
public int Rank => Shape.Length;
+ ///
+ /// Gets direct access to the underlying data array for high-performance operations.
+ ///
+ ///
+ /// Warning: This property provides direct access to internal storage.
+ /// Modifications to this array will affect the tensor. Use with caution in
+ /// performance-critical code paths like SIMD operations.
+ ///
+ public T[] Data => _data.Data;
+
///
/// Initializes a new instance of the TensorBase class with the specified shape.
///
diff --git a/src/AiDotNet.Tensors/LinearAlgebra/VectorBase.cs b/src/AiDotNet.Tensors/LinearAlgebra/VectorBase.cs
index c9c93d65ef..ef7cbfd784 100644
--- a/src/AiDotNet.Tensors/LinearAlgebra/VectorBase.cs
+++ b/src/AiDotNet.Tensors/LinearAlgebra/VectorBase.cs
@@ -74,6 +74,16 @@ protected VectorBase(IEnumerable values)
///
public int Length => _data.Length;
+ ///
+ /// Gets direct access to the underlying data array for high-performance operations.
+ ///
+ ///
+ /// Warning: This property provides direct access to internal storage.
+ /// Modifications to this array will affect the vector. Use with caution in
+ /// performance-critical code paths like SIMD operations.
+ ///
+ public T[] Data => _data;
+
///
/// Gets a value indicating whether the vector contains no elements.
///
diff --git a/src/AiDotNet.csproj b/src/AiDotNet.csproj
index 030925c8d1..31cb5a75c4 100644
--- a/src/AiDotNet.csproj
+++ b/src/AiDotNet.csproj
@@ -3,6 +3,7 @@
net8.0;net471
enable
enable
+ true
True
0.0.5-preview
Ai for .Net
diff --git a/src/Configuration/InferenceOptimizationConfig.cs b/src/Configuration/InferenceOptimizationConfig.cs
index 25dc47f1a5..4e4221e2d5 100644
--- a/src/Configuration/InferenceOptimizationConfig.cs
+++ b/src/Configuration/InferenceOptimizationConfig.cs
@@ -116,6 +116,97 @@ public class InferenceOptimizationConfig
/// Cache eviction policy (default: LRU).
public CacheEvictionPolicy KVCacheEvictionPolicy { get; set; } = CacheEvictionPolicy.LRU;
+ ///
+ /// Gets or sets whether to use a sliding window KV-cache for long contexts.
+ ///
+ ///
+ /// When enabled, only the most recent tokens are kept.
+ /// This is a common industry approach for long-context serving to cap memory usage.
+ ///
+ public bool UseSlidingWindowKVCache { get; set; } = false;
+
+ ///
+ /// Gets or sets the sliding window size in tokens when is enabled.
+ ///
+ /// Window size in tokens (default: 1024).
+ public int KVCacheWindowSize { get; set; } = 1024;
+
+ ///
+ /// Gets or sets the precision used for KV-cache storage.
+ ///
+ ///
+ ///
+ /// Industry-standard serving stores KV-cache in FP16 to halve memory usage and increase cache capacity.
+ /// The default selects FP16 when KV-cache is enabled and the numeric
+ /// type supports it.
+ ///
+ ///
+ /// For Beginners: This setting controls how much memory your model uses during autoregressive inference.
+ ///
+ /// - FP16: Uses about half the memory (recommended default)
+ /// - FP32: Uses more memory but can be slightly more numerically accurate
+ ///
+ /// Most production systems prefer FP16 KV-cache for capacity and throughput.
+ ///
+ ///
+ public KVCachePrecisionMode KVCachePrecision { get; set; } = KVCachePrecisionMode.Auto;
+
+ ///
+ /// Gets or sets the quantization mode used for KV-cache storage.
+ ///
+ ///
+ ///
+ /// KV-cache quantization can further reduce memory beyond FP16 by storing keys/values in int8 with scaling.
+ /// This is an opt-in advanced feature because it can introduce small numerical error.
+ ///
+ /// For Beginners:
+ /// - None (default): Store KV-cache in FP16/FP32 depending on .
+ /// - Int8: Store KV-cache in 8-bit integers to save memory (advanced).
+ ///
+ ///
+ public KVCacheQuantizationMode KVCacheQuantization { get; set; } = KVCacheQuantizationMode.None;
+
+ ///
+ /// Gets or sets whether to use a paged KV-cache backend (vLLM-style) for long-context / multi-sequence serving.
+ ///
+ ///
+ /// When enabled, the system may choose a paged cache implementation that allocates KV memory in fixed-size blocks.
+ /// This is the industry-standard approach for high-throughput serving where many sequences are active concurrently.
+ /// Users can disable this to force the traditional contiguous KV-cache.
+ ///
+ public bool EnablePagedKVCache { get; set; } = true;
+
+ ///
+ /// Gets or sets the block size (in tokens) for the paged KV-cache when enabled.
+ ///
+ ///
+ /// Common values are 16 or 32. Smaller blocks reduce internal fragmentation; larger blocks reduce table overhead.
+ ///
+ public int PagedKVCacheBlockSize { get; set; } = 16;
+
+ #endregion
+
+ #region Attention Settings
+
+ ///
+ /// Gets or sets whether Flash Attention is enabled (when applicable).
+ ///
+ ///
+ /// Flash Attention computes exact attention without materializing the full N×N attention matrix,
+ /// reducing memory bandwidth pressure and improving throughput for long sequences.
+ ///
+ public bool EnableFlashAttention { get; set; } = true;
+
+ ///
+ /// Gets or sets how attention masking should be applied for optimized attention implementations.
+ ///
+ ///
+ /// - Auto: Applies causal masking for known autoregressive models (e.g., text generation), otherwise no mask.
+ /// - Disabled: Never applies causal masking.
+ /// - Causal: Always applies causal masking (GPT-style).
+ ///
+ public AttentionMaskingMode AttentionMasking { get; set; } = AttentionMaskingMode.Auto;
+
#endregion
#region Batching Settings
@@ -251,6 +342,18 @@ public void Validate()
throw new InvalidOperationException(
$"SpeculationDepth must be non-negative. Got: {SpeculationDepth}");
}
+
+ if (UseSlidingWindowKVCache && KVCacheWindowSize <= 0)
+ {
+ throw new InvalidOperationException(
+ $"KVCacheWindowSize must be positive when UseSlidingWindowKVCache is enabled. Got: {KVCacheWindowSize}");
+ }
+
+ if (EnablePagedKVCache && PagedKVCacheBlockSize <= 0)
+ {
+ throw new InvalidOperationException(
+ $"PagedKVCacheBlockSize must be positive when EnablePagedKVCache is enabled. Got: {PagedKVCacheBlockSize}");
+ }
}
#endregion
@@ -294,9 +397,14 @@ public void Validate()
///
/// Options:
/// - NGram: Simple statistical model (fast, no GPU needed)
- /// - SmallNeural: Smaller version of the main model (more accurate drafts)
+ /// - SmallNeural: Smaller companion model (more accurate drafts)
///
/// NGram is usually sufficient and has near-zero overhead.
+ ///
+ ///
+ /// Note: Small neural draft models require an external companion model. In the MVP, the library
+ /// falls back to when a companion draft model is not available.
+ ///
///
///
public DraftModelType DraftModelType { get; set; } = DraftModelType.NGram;
@@ -331,9 +439,110 @@ public void Validate()
///
public bool UseTreeSpeculation { get; set; } = false;
+ ///
+ /// Gets or sets the policy for when speculative decoding should run.
+ ///
+ ///
+ /// Auto is recommended: it can back off speculative decoding under high load (e.g., large batches)
+ /// to avoid throughput regressions, while still enabling it for latency-sensitive scenarios.
+ ///
+ public SpeculationPolicy SpeculationPolicy { get; set; } = SpeculationPolicy.Auto;
+
+ ///
+ /// Gets or sets the speculative decoding method.
+ ///
+ ///
+ ///
+ /// The default currently selects .
+ ///
+ ///
+ /// For Beginners: This chooses the "style" of speculative decoding.
+ ///
+ ///
+ public SpeculativeMethod SpeculativeMethod { get; set; } = SpeculativeMethod.Auto;
+
+ #endregion
+
+ #region Inference Quantization (Advanced)
+
+ ///
+ /// Gets or sets whether weight-only INT8 quantization is enabled for inference.
+ ///
+ ///
+ ///
+ /// Weight-only quantization reduces memory bandwidth and improves cache locality by storing weights in int8
+ /// with per-output scaling. Activations remain in FP32/FP16, and accumulation is performed in float.
+ ///
+ ///
+ /// For Beginners: This makes your model weights smaller so the CPU/GPU can read them faster.
+ ///
+ ///
+ /// This is disabled by default until validated across more layer types and kernels. When enabled, the optimizer
+ /// will apply it opportunistically and fall back safely when unsupported.
+ ///
+ ///
+ public bool EnableWeightOnlyQuantization { get; set; } = false;
+
#endregion
}
+///
+/// Policies for enabling/disabling speculative decoding at runtime.
+///
+public enum SpeculationPolicy
+{
+ ///
+ /// Automatically decide based on runtime conditions (recommended).
+ ///
+ Auto,
+
+ ///
+ /// Always enable speculative decoding when configured.
+ ///
+ ForceOn,
+
+ ///
+ /// Always disable speculative decoding even if enabled in config.
+ ///
+ ForceOff,
+
+ ///
+ /// Prefer speculative decoding to reduce latency, even under moderate load.
+ ///
+ LatencyFirst,
+
+ ///
+ /// Prefer throughput and stability: use speculative decoding only when conditions are ideal.
+ ///
+ ThroughputFirst
+}
+
+///
+/// Selects the speculative decoding method.
+///
+public enum SpeculativeMethod
+{
+ ///
+ /// Automatically select the best available method (defaults to ClassicDraftModel today).
+ ///
+ Auto,
+
+ ///
+ /// Classic draft-model speculative decoding (standard).
+ ///
+ ClassicDraftModel,
+
+ ///
+ /// Medusa-style multi-head proposals (hook for future internal implementation).
+ ///
+ Medusa,
+
+ ///
+ /// EAGLE-style enhanced draft proposals (hook for future internal implementation).
+ ///
+ Eagle
+}
+
///
/// Cache eviction policies for KV cache management.
///
@@ -356,6 +565,65 @@ public enum DraftModelType
NGram,
/// Small neural network model (more accurate, uses GPU).
SmallNeural,
- /// Custom user-provided draft model.
+ /// Custom draft model (internal/serving integration).
Custom
}
+
+///
+/// Controls how attention masking is applied for optimized attention implementations.
+///
+public enum AttentionMaskingMode
+{
+ ///
+ /// Automatically select masking based on model/task heuristics.
+ ///
+ Auto,
+
+ ///
+ /// Do not apply causal masking.
+ ///
+ Disabled,
+
+ ///
+ /// Apply causal masking (autoregressive decoding).
+ ///
+ Causal
+}
+
+///
+/// Controls the numeric precision of KV-cache storage.
+///
+public enum KVCachePrecisionMode
+{
+ ///
+ /// Select an industry-standard default.
+ ///
+ ///
+ ///
+ /// Uses FP16 when KV-cache is enabled and the numeric type supports conversion; otherwise falls back to FP32.
+ ///
+ ///
+ Auto,
+
+ ///
+ /// Store KV-cache in FP16 (half precision) to reduce memory use.
+ ///
+ Float16,
+
+ ///
+ /// Store KV-cache in FP32 (single precision) for maximal numerical fidelity.
+ ///
+ Float32
+}
+
+///
+/// Controls optional KV-cache quantization for inference.
+///
+public enum KVCacheQuantizationMode
+{
+ /// No quantization (default).
+ None,
+
+ /// Signed int8 quantization with scaling (advanced, opt-in).
+ Int8
+}
diff --git a/src/Helpers/DeserializationHelper.cs b/src/Helpers/DeserializationHelper.cs
index 1c2142b55d..158e157535 100644
--- a/src/Helpers/DeserializationHelper.cs
+++ b/src/Helpers/DeserializationHelper.cs
@@ -48,6 +48,13 @@ static DeserializationHelper()
///
public static ILayer CreateLayerFromType(string layerType, int[] inputShape, int[] outputShape, Dictionary? additionalParams = null)
{
+ // Allow layerType to contain serialized constructor metadata, e.g. "MultiHeadAttentionLayer;HeadCount=8".
+ if (TryParseLayerTypeIdentifier(layerType, out var parsedTypeName, out var parsedParams))
+ {
+ layerType = parsedTypeName;
+ additionalParams = MergeParams(additionalParams, parsedParams);
+ }
+
if (!LayerTypes.TryGetValue(layerType, out Type? openGenericType))
{
throw new NotSupportedException($"Layer type {layerType} is not supported for deserialization.");
@@ -88,15 +95,159 @@ public static ILayer CreateLayerFromType(string layerType, int[] inputShap
if (genericDef == typeof(DenseLayer<>))
{
- // DenseLayer(int inputSize, int outputSize, IActivationFunction? activationFunction = null)
- // Use specific constructor to avoid ambiguity with vector activation constructor
+ instance = CreateDenseLayer(type, inputShape, outputShape, additionalParams);
+ }
+ else if (genericDef == typeof(InputLayer<>))
+ {
+ // InputLayer(int inputSize)
+ var ctor = type.GetConstructor([typeof(int)]);
+ if (ctor is null)
+ {
+ throw new InvalidOperationException("Cannot find InputLayer constructor with (int).");
+ }
+
+ instance = ctor.Invoke([inputShape[0]]);
+ }
+ else if (genericDef == typeof(ReshapeLayer<>))
+ {
+ // ReshapeLayer(int[] inputShape, int[] outputShape)
+ var ctor = type.GetConstructor([typeof(int[]), typeof(int[])]);
+ if (ctor is null)
+ {
+ throw new InvalidOperationException("Cannot find ReshapeLayer constructor with (int[], int[]).");
+ }
+
+ instance = ctor.Invoke([inputShape, outputShape]);
+ }
+ else if (genericDef == typeof(EmbeddingLayer<>))
+ {
+ // EmbeddingLayer(int vocabularySize, int embeddingDimension)
+ int embeddingDim = outputShape[0];
+ int vocabSize = TryGetInt(additionalParams, "VocabularySize")
+ ?? TryGetInt(additionalParams, "VocabSize")
+ ?? throw new InvalidOperationException("EmbeddingLayer requires VocabularySize metadata for deserialization.");
+
+ var ctor = type.GetConstructor([typeof(int), typeof(int)]);
+ if (ctor is null)
+ {
+ throw new InvalidOperationException("Cannot find EmbeddingLayer constructor with (int, int).");
+ }
+ instance = ctor.Invoke([vocabSize, embeddingDim]);
+ }
+ else if (genericDef == typeof(PositionalEncodingLayer<>))
+ {
+ // PositionalEncodingLayer(int maxSequenceLength, int embeddingSize)
+ if (inputShape.Length < 2)
+ {
+ throw new InvalidOperationException("PositionalEncodingLayer requires input shape [maxSequenceLength, embeddingSize].");
+ }
+
+ int maxSeqLen = inputShape[0];
+ int embDim = inputShape[1];
+
+ var ctor = type.GetConstructor([typeof(int), typeof(int)]);
+ if (ctor is null)
+ {
+ throw new InvalidOperationException("Cannot find PositionalEncodingLayer constructor with (int, int).");
+ }
+ instance = ctor.Invoke([maxSeqLen, embDim]);
+ }
+ else if (genericDef == typeof(DropoutLayer<>))
+ {
+ // DropoutLayer(double dropoutRate = 0.5)
+ double rate = TryGetDouble(additionalParams, "DropoutRate") ?? 0.5;
+ var ctor = type.GetConstructor([typeof(double)]);
+ if (ctor is null)
+ {
+ throw new InvalidOperationException("Cannot find DropoutLayer constructor with (double).");
+ }
+ instance = ctor.Invoke([rate]);
+ }
+ else if (genericDef == typeof(LayerNormalizationLayer<>))
+ {
+ // LayerNormalizationLayer(int featureSize, double epsilon = ...)
+ int featureSize = inputShape[0];
+ double epsilon = TryGetDouble(additionalParams, "Epsilon") ?? NumericalStabilityHelper.LargeEpsilon;
+ var ctor = type.GetConstructor([typeof(int), typeof(double)]);
+ if (ctor is null)
+ {
+ throw new InvalidOperationException("Cannot find LayerNormalizationLayer constructor with (int, double).");
+ }
+ instance = ctor.Invoke([featureSize, epsilon]);
+ }
+ else if (genericDef == typeof(MultiHeadAttentionLayer<>))
+ {
+ instance = CreateMultiHeadAttentionLayer(type, inputShape, additionalParams);
+ }
+ else if (genericDef == typeof(SelfAttentionLayer<>))
+ {
+ // SelfAttentionLayer(int sequenceLength, int embeddingDimension, int headCount = 8, IActivationFunction? = null)
+ if (inputShape.Length < 2)
+ {
+ throw new InvalidOperationException("SelfAttentionLayer requires input shape [sequenceLength, embeddingDimension].");
+ }
+
+ int seqLen = inputShape[0];
+ int embDim = inputShape[1];
+ int headCount = TryGetInt(additionalParams, "HeadCount") ?? ResolveDefaultHeadCount(embDim);
+
+ var activationFuncType = typeof(IActivationFunction<>).MakeGenericType(typeof(T));
+ var ctor = type.GetConstructor([typeof(int), typeof(int), typeof(int), activationFuncType]);
+ if (ctor is null)
+ {
+ throw new InvalidOperationException("Cannot find SelfAttentionLayer constructor with (int, int, int, IActivationFunction).");
+ }
+ object? activation = TryCreateActivationInstance(additionalParams, "ScalarActivationType", activationFuncType);
+ instance = ctor.Invoke([seqLen, embDim, headCount, activation]);
+ }
+ else if (genericDef == typeof(AttentionLayer<>))
+ {
+ // AttentionLayer(int inputSize, int attentionSize, IActivationFunction? = null)
+ int inputSize = inputShape[0];
+ int attentionSize = outputShape[0];
+
var activationFuncType = typeof(IActivationFunction<>).MakeGenericType(typeof(T));
var ctor = type.GetConstructor([typeof(int), typeof(int), activationFuncType]);
if (ctor is null)
{
- throw new InvalidOperationException($"Cannot find DenseLayer constructor with (int, int, IActivationFunction).");
+ throw new InvalidOperationException("Cannot find AttentionLayer constructor with (int, int, IActivationFunction).");
}
- instance = ctor.Invoke([inputShape[0], outputShape[0], null]);
+ object? activation = TryCreateActivationInstance(additionalParams, "ScalarActivationType", activationFuncType);
+ instance = ctor.Invoke([inputSize, attentionSize, activation]);
+ }
+ else if (genericDef == typeof(GraphAttentionLayer<>))
+ {
+ // GraphAttentionLayer(int inputFeatures, int outputFeatures, int numHeads = 1, double alpha = 0.2, double dropoutRate = 0.0, IActivationFunction? = null)
+ int inputFeatures = inputShape[0];
+ int outputFeatures = outputShape[0];
+ int numHeads = TryGetInt(additionalParams, "NumHeads") ?? 1;
+ double alpha = TryGetDouble(additionalParams, "Alpha") ?? 0.2;
+ double dropout = TryGetDouble(additionalParams, "DropoutRate") ?? 0.0;
+
+ var activationFuncType = typeof(IActivationFunction<>).MakeGenericType(typeof(T));
+ var ctor = type.GetConstructor([typeof(int), typeof(int), typeof(int), typeof(double), typeof(double), activationFuncType]);
+ if (ctor is null)
+ {
+ throw new InvalidOperationException("Cannot find GraphAttentionLayer constructor with expected signature.");
+ }
+ object? activation = TryCreateActivationInstance(additionalParams, "ScalarActivationType", activationFuncType);
+ instance = ctor.Invoke([inputFeatures, outputFeatures, numHeads, alpha, dropout, activation]);
+ }
+ else if (genericDef == typeof(AiDotNet.NeuralNetworks.Attention.FlashAttentionLayer<>))
+ {
+ instance = CreateFlashAttentionLayer(type, inputShape, additionalParams);
+ }
+ else if (genericDef == typeof(AiDotNet.Inference.CachedMultiHeadAttention<>))
+ {
+ instance = CreateCachedMultiHeadAttention(type, inputShape, additionalParams);
+ }
+ else if (genericDef == typeof(AiDotNet.Inference.PagedCachedMultiHeadAttention<>))
+ {
+ instance = CreatePagedCachedMultiHeadAttention(type, inputShape, additionalParams);
+ }
+ else if (genericDef == typeof(AiDotNet.LoRA.Adapters.MultiLoRAAdapter<>))
+ {
+ instance = CreateMultiLoRAAdapter(type, inputShape, outputShape, additionalParams);
}
else if (genericDef == typeof(ConvolutionalLayer<>))
{
@@ -139,42 +290,442 @@ public static ILayer CreateLayerFromType(string layerType, int[] inputShap
}
else if (genericDef == typeof(ActivationLayer<>))
{
- // ActivationLayer(int[] inputShape, IActivationFunction activationFunction)
+ instance = CreateActivationLayer(type, inputShape, additionalParams);
+ }
+ else
+ {
+ // Default: pass inputShape as first parameter
+ var ctor = type.GetConstructor([typeof(int[])]);
+ if (ctor is null)
+ {
+ throw new NotSupportedException(
+ $"Layer type {layerType} is not supported for deserialization (no known constructor found).");
+ }
+
+ instance = ctor.Invoke([inputShape]);
+ }
+ if (instance == null)
+ {
+ throw new InvalidOperationException($"Failed to create instance of layer type {layerType}.");
+ }
+
+ return (ILayer)instance;
+ }
+
+ private static object CreateDenseLayer(Type type, int[] inputShape, int[] outputShape, Dictionary? additionalParams)
+ {
+ // DenseLayer(int inputSize, int outputSize, IActivationFunction? activationFunction = null)
+ // Use specific constructor to avoid ambiguity with vector activation constructor.
+ var activationFuncType = typeof(IActivationFunction<>).MakeGenericType(typeof(T));
+ var ctor = type.GetConstructor([typeof(int), typeof(int), activationFuncType]);
+ if (ctor is null)
+ {
+ throw new InvalidOperationException("Cannot find DenseLayer constructor with (int, int, IActivationFunction).");
+ }
+
+ object? activation = TryCreateActivationInstance(additionalParams, "ScalarActivationType", activationFuncType);
+ return ctor.Invoke([inputShape[0], outputShape[0], activation]);
+ }
+
+ private static object CreateMultiHeadAttentionLayer(Type type, int[] inputShape, Dictionary? additionalParams)
+ {
+ // MultiHeadAttentionLayer(int sequenceLength, int embeddingDimension, int headCount, IActivationFunction? activationFunction = null)
+ if (inputShape.Length < 2)
+ {
+ throw new InvalidOperationException("MultiHeadAttentionLayer requires input shape [sequenceLength, embeddingDimension].");
+ }
+
+ int seqLen = inputShape[0];
+ int embDim = inputShape[1];
+ int headCount = TryGetInt(additionalParams, "HeadCount") ?? ResolveDefaultHeadCount(embDim);
+
+ var activationFuncType = typeof(IActivationFunction<>).MakeGenericType(typeof(T));
+ var ctor = type.GetConstructor([typeof(int), typeof(int), typeof(int), activationFuncType]);
+ if (ctor is null)
+ {
+ throw new InvalidOperationException("Cannot find MultiHeadAttentionLayer constructor with (int, int, int, IActivationFunction).");
+ }
+
+ object? activation = TryCreateActivationInstance(additionalParams, "ScalarActivationType", activationFuncType);
+ return ctor.Invoke([seqLen, embDim, headCount, activation]);
+ }
+
+ private static object CreateFlashAttentionLayer(Type type, int[] inputShape, Dictionary? additionalParams)
+ {
+ // FlashAttentionLayer(int sequenceLength, int embeddingDimension, int headCount, FlashAttentionConfig config, IActivationFunction? activationFunction = null)
+ if (inputShape.Length < 2)
+ {
+ throw new InvalidOperationException("FlashAttentionLayer requires input shape [sequenceLength, embeddingDimension].");
+ }
+
+ int seqLen = inputShape[0];
+ int embDim = inputShape[1];
+ int headCount = TryGetInt(additionalParams, "HeadCount") ?? ResolveDefaultHeadCount(embDim);
+ bool useCausal = TryGetBool(additionalParams, "UseCausalMask") ?? false;
+
+ var flashConfig = AiDotNet.NeuralNetworks.Attention.FlashAttentionConfig.Default;
+ flashConfig.UseCausalMask = useCausal;
+
+ var activationFuncType = typeof(IActivationFunction<>).MakeGenericType(typeof(T));
+ var ctor = type.GetConstructor([typeof(int), typeof(int), typeof(int), typeof(AiDotNet.NeuralNetworks.Attention.FlashAttentionConfig), activationFuncType]);
+ if (ctor is null)
+ {
+ throw new InvalidOperationException("Cannot find FlashAttentionLayer constructor with expected signature.");
+ }
+
+ object? activation = TryCreateActivationInstance(additionalParams, "ScalarActivationType", activationFuncType);
+ return ctor.Invoke([seqLen, embDim, headCount, flashConfig, activation]);
+ }
+
+ private static object CreateCachedMultiHeadAttention(Type type, int[] inputShape, Dictionary? additionalParams)
+ {
+ // CachedMultiHeadAttention(int sequenceLength, int embeddingDimension, int headCount, bool useFlashAttention, int layerIndex, bool useCausalMask, IActivationFunction? activationFunction = null)
+ if (inputShape.Length < 2)
+ {
+ throw new InvalidOperationException("CachedMultiHeadAttention requires input shape [sequenceLength, embeddingDimension].");
+ }
+
+ int seqLen = inputShape[0];
+ int embDim = inputShape[1];
+ int headCount = TryGetInt(additionalParams, "HeadCount") ?? ResolveDefaultHeadCount(embDim);
+ bool useFlash = TryGetBool(additionalParams, "UseFlashAttention") ?? true;
+ bool useCausal = TryGetBool(additionalParams, "UseCausalMask") ?? true;
+
+ var activationFuncType = typeof(IActivationFunction<>).MakeGenericType(typeof(T));
+ var ctor = type.GetConstructor([typeof(int), typeof(int), typeof(int), typeof(bool), typeof(int), typeof(bool), activationFuncType]);
+ if (ctor is null)
+ {
+ throw new InvalidOperationException("Cannot find CachedMultiHeadAttention constructor with expected signature.");
+ }
+
+ object? activation = TryCreateActivationInstance(additionalParams, "ScalarActivationType", activationFuncType);
+ return ctor.Invoke([seqLen, embDim, headCount, useFlash, 0, useCausal, activation]);
+ }
+
+ private static object CreatePagedCachedMultiHeadAttention(Type type, int[] inputShape, Dictionary? additionalParams)
+ {
+ // PagedCachedMultiHeadAttention(int sequenceLength, int embeddingDimension, int headCount, bool useCausalMask, IActivationFunction? activationFunction = null)
+ if (inputShape.Length < 2)
+ {
+ throw new InvalidOperationException("PagedCachedMultiHeadAttention requires input shape [sequenceLength, embeddingDimension].");
+ }
+
+ int seqLen = inputShape[0];
+ int embDim = inputShape[1];
+ int headCount = TryGetInt(additionalParams, "HeadCount") ?? ResolveDefaultHeadCount(embDim);
+ bool useCausal = TryGetBool(additionalParams, "UseCausalMask") ?? true;
+
+ var activationFuncType = typeof(IActivationFunction<>).MakeGenericType(typeof(T));
+ var ctor = type.GetConstructor([typeof(int), typeof(int), typeof(int), typeof(bool), activationFuncType]);
+ if (ctor is null)
+ {
+ throw new InvalidOperationException("Cannot find PagedCachedMultiHeadAttention constructor with expected signature.");
+ }
+
+ object? activation = TryCreateActivationInstance(additionalParams, "ScalarActivationType", activationFuncType);
+ return ctor.Invoke([seqLen, embDim, headCount, useCausal, activation]);
+ }
+
+ private static object CreateMultiLoRAAdapter(Type type, int[] inputShape, int[] outputShape, Dictionary? additionalParams)
+ {
+ // MultiLoRAAdapter(ILayer baseLayer, string defaultTaskName, int defaultRank, double alpha, bool freezeBaseLayer)
+ bool freezeBaseLayer = TryGetBool(additionalParams, "FreezeBaseLayer") ?? true;
+
+ string? encodedBaseLayerId = additionalParams?.TryGetValue("BaseLayerTypeId", out var baseType) == true ? baseType as string : null;
+ string baseLayerIdentifier = !string.IsNullOrWhiteSpace(encodedBaseLayerId)
+ ? Uri.UnescapeDataString(encodedBaseLayerId)
+ : "DenseLayer`1";
+
+ var baseLayer = CreateLayerFromType(baseLayerIdentifier, inputShape, outputShape, null);
+
+ static string[] ParseList(string? raw)
+ {
+ if (string.IsNullOrWhiteSpace(raw)) return Array.Empty();
+ return raw!.Split(new[] { '|' }, StringSplitOptions.RemoveEmptyEntries);
+ }
+
+ static int[] ParseIntList(string? raw)
+ {
+ var parts = ParseList(raw);
+ var result = new int[parts.Length];
+ for (int i = 0; i < parts.Length; i++)
+ {
+ result[i] = int.TryParse(parts[i], System.Globalization.NumberStyles.Integer, System.Globalization.CultureInfo.InvariantCulture, out var v) ? v : 1;
+ }
+ return result;
+ }
+
+ static double[] ParseDoubleList(string? raw)
+ {
+ var parts = ParseList(raw);
+ var result = new double[parts.Length];
+ for (int i = 0; i < parts.Length; i++)
+ {
+ result[i] = double.TryParse(parts[i], System.Globalization.NumberStyles.Float, System.Globalization.CultureInfo.InvariantCulture, out var v) ? v : -1;
+ }
+ return result;
+ }
+
+ string? tasksRaw = additionalParams?.TryGetValue("Tasks", out var tasksObj) == true ? tasksObj as string : null;
+ var encodedTasks = ParseList(tasksRaw);
+ if (encodedTasks.Length == 0)
+ {
+ encodedTasks = ["default"];
+ }
+
+ var tasks = encodedTasks.Select(Uri.UnescapeDataString).ToArray();
+ var ranks = ParseIntList(additionalParams?.TryGetValue("TaskRanks", out var ranksObj) == true ? ranksObj as string : null);
+ var alphas = ParseDoubleList(additionalParams?.TryGetValue("TaskAlphas", out var alphasObj) == true ? alphasObj as string : null);
+
+ int defaultRank = ranks.Length > 0 ? ranks[0] : 1;
+ double defaultAlpha = alphas.Length > 0 ? alphas[0] : -1;
+
+ var iLayerType = typeof(ILayer<>).MakeGenericType(typeof(T));
+ var ctor = type.GetConstructor([iLayerType, typeof(string), typeof(int), typeof(double), typeof(bool)]);
+ if (ctor is null)
+ {
+ throw new InvalidOperationException("Cannot find MultiLoRAAdapter constructor with expected signature.");
+ }
+
+ var instance = ctor.Invoke([baseLayer, tasks[0], defaultRank, defaultAlpha, freezeBaseLayer]);
+ var multi = (AiDotNet.LoRA.Adapters.MultiLoRAAdapter)instance;
+
+ for (int taskIndex = 1; taskIndex < tasks.Length; taskIndex++)
+ {
+ int rank = taskIndex < ranks.Length ? ranks[taskIndex] : defaultRank;
+ double alpha = taskIndex < alphas.Length ? alphas[taskIndex] : -1;
+ multi.AddTask(tasks[taskIndex], rank, alpha);
+ }
+
+ if (additionalParams?.TryGetValue("CurrentTask", out var currentTaskObj) == true &&
+ currentTaskObj is string currentTaskEncoded)
+ {
+ string currentTask = Uri.UnescapeDataString(currentTaskEncoded);
+ if (!string.IsNullOrWhiteSpace(currentTask))
+ {
+ multi.SetCurrentTask(currentTask);
+ }
+ }
+
+ return instance;
+ }
+
+ private static object CreateActivationLayer(Type type, int[] inputShape, Dictionary? additionalParams)
+ {
+ // ActivationLayer(int[] inputShape, IActivationFunction activationFunction)
+ var scalarActivationType = typeof(IActivationFunction<>).MakeGenericType(typeof(T));
+ var vectorActivationType = typeof(IVectorActivationFunction<>).MakeGenericType(typeof(T));
+
+ object? vectorActivation = TryCreateActivationInstance(additionalParams, "VectorActivationType", vectorActivationType);
+ object? scalarActivation = TryCreateActivationInstance(additionalParams, "ScalarActivationType", scalarActivationType);
+
+ object? activationFunction = vectorActivation ?? scalarActivation;
+
+ if (activationFunction == null)
+ {
+ // Back-compat fallback: use enum if available, otherwise default ReLU.
ActivationFunction activationFunctionEnum = additionalParams?.TryGetValue("ActivationFunction", out var af) == true
? (ActivationFunction)af : ActivationFunction.ReLU;
- // Use ActivationFunctionFactory to create the IActivationFunction from enum
+
var factoryType = typeof(ActivationFunctionFactory<>).MakeGenericType(typeof(T));
var createMethod = factoryType.GetMethod("CreateActivationFunction", BindingFlags.Public | BindingFlags.Static);
if (createMethod is null)
{
throw new InvalidOperationException("Cannot find ActivationFunctionFactory.CreateActivationFunction method.");
}
- object? activationFunction = createMethod.Invoke(null, [activationFunctionEnum]);
- if (activationFunction is null)
+
+ activationFunction = createMethod.Invoke(null, [activationFunctionEnum]);
+ }
+
+ if (activationFunction == null)
+ {
+ throw new InvalidOperationException("Failed to create activation function for ActivationLayer.");
+ }
+
+ if (vectorActivationType.IsInstanceOfType(activationFunction))
+ {
+ var ctor = type.GetConstructor([typeof(int[]), vectorActivationType]);
+ if (ctor is null)
{
- throw new InvalidOperationException($"Failed to create activation function for {activationFunctionEnum}.");
+ throw new InvalidOperationException("Cannot find ActivationLayer constructor with (int[], IVectorActivationFunction).");
}
+ return ctor.Invoke([inputShape, activationFunction]);
+ }
- // Use specific constructor to avoid ambiguity with vector activation constructor
- var activationFuncType = typeof(IActivationFunction<>).MakeGenericType(typeof(T));
- var ctor = type.GetConstructor([typeof(int[]), activationFuncType]);
- if (ctor is null)
+ var scalarCtor = type.GetConstructor([typeof(int[]), scalarActivationType]);
+ if (scalarCtor is null)
+ {
+ throw new InvalidOperationException("Cannot find ActivationLayer constructor with (int[], IActivationFunction).");
+ }
+ return scalarCtor.Invoke([inputShape, activationFunction]);
+ }
+
+ private static bool TryParseLayerTypeIdentifier(
+ string identifier,
+ out string typeName,
+ out Dictionary parameters)
+ {
+ typeName = identifier;
+ parameters = new Dictionary(StringComparer.Ordinal);
+
+ int sep = identifier.IndexOf(';');
+ if (sep < 0)
+ {
+ return false;
+ }
+
+ typeName = identifier.Substring(0, sep);
+ var parts = identifier.Substring(sep + 1).Split(new[] { ';' }, StringSplitOptions.RemoveEmptyEntries);
+ foreach (var part in parts)
+ {
+ int eq = part.IndexOf('=');
+ if (eq <= 0 || eq == part.Length - 1)
+ {
+ continue;
+ }
+
+ string key = part.Substring(0, eq);
+ string value = part.Substring(eq + 1);
+
+ if (int.TryParse(value, System.Globalization.NumberStyles.Integer, System.Globalization.CultureInfo.InvariantCulture, out int i))
+ {
+ parameters[key] = i;
+ }
+ else if (long.TryParse(value, System.Globalization.NumberStyles.Integer, System.Globalization.CultureInfo.InvariantCulture, out long l))
{
- throw new InvalidOperationException($"Cannot find ActivationLayer constructor with (int[], IActivationFunction).");
+ parameters[key] = l;
+ }
+ else if (double.TryParse(value, System.Globalization.NumberStyles.Float, System.Globalization.CultureInfo.InvariantCulture, out double d))
+ {
+ parameters[key] = d;
+ }
+ else if (bool.TryParse(value, out bool b))
+ {
+ parameters[key] = b;
+ }
+ else
+ {
+ parameters[key] = value;
}
- instance = ctor.Invoke([inputShape, activationFunction]);
}
- else
+
+ return true;
+ }
+
+ private static Dictionary MergeParams(
+ Dictionary? original,
+ Dictionary parsed)
+ {
+ if (original == null || original.Count == 0)
{
- // Default: pass inputShape as first parameter
- instance = Activator.CreateInstance(type, [inputShape]);
+ return parsed;
}
- if (instance == null)
+
+ foreach (var kvp in parsed)
{
- throw new InvalidOperationException($"Failed to create instance of layer type {layerType}.");
+ original[kvp.Key] = kvp.Value;
}
- return (ILayer)instance;
+ return original;
+ }
+
+ private static int? TryGetInt(Dictionary? parameters, string key)
+ {
+ if (parameters != null && parameters.TryGetValue(key, out var value) && value != null)
+ {
+ if (value is int i)
+ return i;
+ if (value is long l && l >= int.MinValue && l <= int.MaxValue)
+ return (int)l;
+ if (int.TryParse(value.ToString() ?? string.Empty, out int parsed))
+ return parsed;
+ }
+ return null;
+ }
+
+ private static double? TryGetDouble(Dictionary? parameters, string key)
+ {
+ if (parameters != null && parameters.TryGetValue(key, out var value) && value != null)
+ {
+ if (value is double d)
+ return d;
+ if (double.TryParse(value.ToString() ?? string.Empty, System.Globalization.NumberStyles.Float, System.Globalization.CultureInfo.InvariantCulture, out double parsed))
+ return parsed;
+ }
+ return null;
+ }
+
+ private static bool? TryGetBool(Dictionary? parameters, string key)
+ {
+ if (parameters != null && parameters.TryGetValue(key, out var value) && value != null)
+ {
+ if (value is bool b)
+ return b;
+ if (bool.TryParse(value.ToString() ?? string.Empty, out bool parsed))
+ return parsed;
+ }
+ return null;
+ }
+
+ private static object? TryCreateActivationInstance(
+ Dictionary? parameters,
+ string key,
+ Type expectedInterface)
+ {
+ if (parameters == null || !parameters.TryGetValue(key, out var value) || value == null)
+ {
+ return null;
+ }
+
+ string? typeName = value as string ?? value.ToString() ?? string.Empty;
+ if (string.IsNullOrWhiteSpace(typeName))
+ {
+ return null;
+ }
+
+ var type = Type.GetType(typeName, throwOnError: false);
+ if (type == null)
+ {
+ return null;
+ }
+
+ try
+ {
+ var instance = Activator.CreateInstance(type);
+ if (instance == null)
+ {
+ return null;
+ }
+
+ return expectedInterface.IsInstanceOfType(instance) ? instance : null;
+ }
+ catch (MissingMethodException)
+ {
+ return null;
+ }
+ catch (TargetInvocationException ex) when (ex.InnerException is MissingMethodException)
+ {
+ return null;
+ }
+ catch (Exception ex)
+ {
+ // Best-effort: deserialization should not throw if an optional activation cannot be created.
+ System.Diagnostics.Debug.WriteLine($"Unexpected error deserializing activation {typeName}: {ex.Message}");
+ return null;
+ }
+ }
+
+ private static int ResolveDefaultHeadCount(int embeddingDimension)
+ {
+ // Conservative but practical default: prefer common head counts if divisible, otherwise fall back to 1.
+ foreach (var candidate in new[] { 8, 4, 16, 12, 6, 2, 1 })
+ {
+ if (candidate > 0 && embeddingDimension % candidate == 0)
+ {
+ return candidate;
+ }
+ }
+ return 1;
}
///
@@ -203,7 +754,25 @@ public static ILayer CreateLayerFromType(string layerType, int[] inputShap
throw new InvalidOperationException($"Type {typeName} does not implement interface {typeof(TInterface).Name}");
}
- return (TInterface?)Activator.CreateInstance(type)
- ?? throw new InvalidOperationException($"Failed to create instance of type {typeName}");
+ try
+ {
+ return (TInterface?)Activator.CreateInstance(type)
+ ?? throw new InvalidOperationException($"Failed to create instance of type {typeName}");
+ }
+ catch (MissingMethodException)
+ {
+ // Some implementations require constructor arguments.
+ // Treat them as optional on deserialization and let callers provide sensible defaults.
+ return null;
+ }
+ catch (TargetInvocationException ex) when (ex.InnerException is MissingMethodException)
+ {
+ // Same as above: no parameterless ctor available.
+ return null;
+ }
+ catch (Exception ex)
+ {
+ throw new InvalidOperationException($"Failed to instantiate type {typeName}", ex);
+ }
}
}
diff --git a/src/Helpers/InferenceDiagnostics.cs b/src/Helpers/InferenceDiagnostics.cs
new file mode 100644
index 0000000000..11b501a2ab
--- /dev/null
+++ b/src/Helpers/InferenceDiagnostics.cs
@@ -0,0 +1,89 @@
+using System;
+using System.Collections.Concurrent;
+
+namespace AiDotNet.Helpers;
+
+///
+/// Internal diagnostics for inference decisions (non-user-facing).
+/// Enable by setting env var AIDOTNET_DIAGNOSTICS=1.
+///
+internal static class InferenceDiagnostics
+{
+ private const int MaxEntries = 1024;
+
+ private static readonly ConcurrentQueue Entries = new();
+
+ private static bool IsEnabled()
+ {
+ var value = Environment.GetEnvironmentVariable("AIDOTNET_DIAGNOSTICS");
+ return string.Equals(value, "1", StringComparison.OrdinalIgnoreCase) ||
+ string.Equals(value, "true", StringComparison.OrdinalIgnoreCase);
+ }
+
+ internal static void RecordDecision(string area, string feature, bool enabled, string reason)
+ {
+ if (!IsEnabled())
+ return;
+
+ Entries.Enqueue(new InferenceDiagnosticEntry(
+ TimestampUtc: DateTime.UtcNow,
+ Area: area ?? string.Empty,
+ Feature: feature ?? string.Empty,
+ Enabled: enabled,
+ Reason: reason ?? string.Empty,
+ ExceptionType: null,
+ ExceptionMessage: null));
+
+ TrimIfNeeded();
+ }
+
+ internal static void RecordException(string area, string feature, Exception ex, string reason)
+ {
+ if (!IsEnabled())
+ return;
+
+ Entries.Enqueue(new InferenceDiagnosticEntry(
+ TimestampUtc: DateTime.UtcNow,
+ Area: area ?? string.Empty,
+ Feature: feature ?? string.Empty,
+ Enabled: false,
+ Reason: reason ?? string.Empty,
+ ExceptionType: ex.GetType().FullName ?? ex.GetType().Name,
+ ExceptionMessage: ex.Message));
+
+ TrimIfNeeded();
+ }
+
+ // Intentionally internal-only: serving can use InternalsVisibleTo to read these if needed later.
+ internal static InferenceDiagnosticEntry[] Snapshot()
+ {
+ if (!IsEnabled())
+ return Array.Empty();
+
+ return Entries.ToArray();
+ }
+
+ internal static void Clear()
+ {
+ while (Entries.TryDequeue(out _))
+ {
+ }
+ }
+
+ private static void TrimIfNeeded()
+ {
+ // Best-effort: bound memory use when diagnostics are enabled.
+ while (Entries.Count > MaxEntries && Entries.TryDequeue(out _))
+ {
+ }
+ }
+
+ internal readonly record struct InferenceDiagnosticEntry(
+ DateTime TimestampUtc,
+ string Area,
+ string Feature,
+ bool Enabled,
+ string Reason,
+ string? ExceptionType,
+ string? ExceptionMessage);
+}
diff --git a/src/Inference/CachedMultiHeadAttention.cs b/src/Inference/CachedMultiHeadAttention.cs
index 59042aab2c..ef820a18db 100644
--- a/src/Inference/CachedMultiHeadAttention.cs
+++ b/src/Inference/CachedMultiHeadAttention.cs
@@ -34,12 +34,13 @@ namespace AiDotNet.Inference;
///
///
/// The numeric type for computations.
-public class CachedMultiHeadAttention : LayerBase
+internal class CachedMultiHeadAttention : LayerBase
{
private readonly int _headCount;
private readonly int _headDimension;
private readonly int _embeddingDimension;
private readonly bool _useFlashAttention;
+ private readonly bool _useCausalMask;
// Projection weights
private Matrix _queryWeights;
@@ -92,6 +93,15 @@ public class CachedMultiHeadAttention : LayerBase
///
public bool UsesFlashAttention => _useFlashAttention;
+ ///
+ /// Gets whether causal masking is enabled for attention.
+ ///
+ ///
+ /// Causal masking is required for autoregressive decoding (GPT-style), where each token may only attend
+ /// to itself and previous tokens. Disable for bidirectional attention (BERT-style) and most encoders.
+ ///
+ public bool UsesCausalMask => _useCausalMask;
+
///
/// Gets or sets the KV-Cache. Must be set before inference.
///
@@ -118,15 +128,20 @@ public int LayerIndex
/// Number of attention heads.
/// Whether to use Flash Attention algorithm.
/// Index of this layer in the transformer (for cache access).
+ /// Whether to apply causal masking (required for autoregressive decoding).
+ /// Optional activation function (defaults to identity).
public CachedMultiHeadAttention(
int sequenceLength,
int embeddingDimension,
int headCount,
bool useFlashAttention = true,
- int layerIndex = 0)
+ int layerIndex = 0,
+ bool useCausalMask = true,
+ IActivationFunction? activationFunction = null)
: base(
[sequenceLength, embeddingDimension],
- [sequenceLength, embeddingDimension])
+ [sequenceLength, embeddingDimension],
+ activationFunction ?? new IdentityActivation())
{
if (embeddingDimension % headCount != 0)
{
@@ -139,6 +154,7 @@ public CachedMultiHeadAttention(
_embeddingDimension = embeddingDimension;
_useFlashAttention = useFlashAttention;
_layerIndex = layerIndex;
+ _useCausalMask = useCausalMask;
// Initialize projection weights
_queryWeights = new Matrix(embeddingDimension, embeddingDimension);
@@ -230,13 +246,18 @@ private Tensor ForwardWithCache(Tensor input)
Tensor attentionOutput;
if (_useFlashAttention)
{
- var config = new FlashAttentionConfig { UseCausalMask = true };
- var (flashOutput, _) = FlashAttention.Forward(queries, keys, values, config);
+ var config = FlashAttentionConfig.Default;
+ config.UseCausalMask = _useCausalMask;
+
+ int seqLenKV = keys.Shape[2];
+ int seqLenQ = queries.Shape[2];
+ int queryOffset = Math.Max(0, seqLenKV - seqLenQ);
+ var (flashOutput, _) = FlashAttention.Forward(queries, keys, values, config, queryOffset: queryOffset);
attentionOutput = flashOutput;
}
else
{
- attentionOutput = StandardAttention(queries, keys, values, useCausalMask: true);
+ attentionOutput = StandardAttention(queries, keys, values, useCausalMask: _useCausalMask);
}
// Reshape back to [batch, seq, embDim]
@@ -244,9 +265,9 @@ private Tensor ForwardWithCache(Tensor input)
// Output projection
var output = attentionOutput.Multiply(_outputWeights).Add(_outputBias);
- _lastOutput = output;
+ _lastOutput = ApplyActivation(output);
- return output;
+ return _lastOutput;
}
///
@@ -272,12 +293,13 @@ private Tensor ForwardStandard(Tensor input)
if (_useFlashAttention)
{
var config = FlashAttentionConfig.Default;
+ config.UseCausalMask = _useCausalMask;
var (flashOutput, _) = FlashAttention.Forward(queries, keys, values, config);
attentionOutput = flashOutput;
}
else
{
- attentionOutput = StandardAttention(queries, keys, values, useCausalMask: false);
+ attentionOutput = StandardAttention(queries, keys, values, useCausalMask: _useCausalMask);
}
// Reshape back
@@ -285,9 +307,9 @@ private Tensor ForwardStandard(Tensor input)
// Output projection
var output = attentionOutput.Multiply(_outputWeights).Add(_outputBias);
- _lastOutput = output;
+ _lastOutput = ApplyActivation(output);
- return output;
+ return _lastOutput;
}
///
@@ -387,6 +409,8 @@ public override Tensor Backward(Tensor outputGradient)
throw new InvalidOperationException("Forward pass must be called before backward pass.");
}
+ var activationGradient = ApplyActivationDerivative(_lastOutput, outputGradient);
+
// Standard backward pass (no cache during training)
// Implementation similar to MultiHeadAttentionLayer
var inputGradient = new Tensor(_lastInput.Shape);
@@ -397,7 +421,7 @@ public override Tensor Backward(Tensor outputGradient)
_keyWeightsGradient = new Matrix(_keyWeights.Rows, _keyWeights.Columns);
_valueWeightsGradient = new Matrix(_valueWeights.Rows, _valueWeights.Columns);
_outputWeightsGradient = new Matrix(_outputWeights.Rows, _outputWeights.Columns);
- _outputBiasGradient = outputGradient.Sum([0, 1]).ToVector();
+ _outputBiasGradient = activationGradient.Sum([0, 1]).ToVector();
return inputGradient;
}
@@ -512,6 +536,7 @@ public override Dictionary GetDiagnostics()
diagnostics["HeadDimension"] = _headDimension.ToString();
diagnostics["InferenceMode"] = InferenceMode.ToString();
diagnostics["UsesFlashAttention"] = _useFlashAttention.ToString();
+ diagnostics["UsesCausalMask"] = _useCausalMask.ToString();
diagnostics["LayerIndex"] = _layerIndex.ToString();
diagnostics["CacheAttached"] = (_cache != null).ToString();
@@ -580,4 +605,14 @@ private Tensor MatrixToTensor(Matrix matrix)
}
return tensor;
}
+
+ internal override Dictionary GetMetadata()
+ {
+ return new Dictionary
+ {
+ ["HeadCount"] = _headCount.ToString(),
+ ["UseFlashAttention"] = _useFlashAttention.ToString(),
+ ["UseCausalMask"] = _useCausalMask.ToString()
+ };
+ }
}
diff --git a/src/Inference/InferenceOptimizer.cs b/src/Inference/InferenceOptimizer.cs
index 5085467b98..d185a36465 100644
--- a/src/Inference/InferenceOptimizer.cs
+++ b/src/Inference/InferenceOptimizer.cs
@@ -1,8 +1,14 @@
using AiDotNet.Configuration;
+using AiDotNet.Helpers;
+using AiDotNet.Inference.PagedAttention;
+using AiDotNet.Inference.Quantization;
using AiDotNet.Inference.SpeculativeDecoding;
using AiDotNet.NeuralNetworks;
+using AiDotNet.NeuralNetworks.Attention;
using AiDotNet.NeuralNetworks.Layers;
+using AiDotNet.Tensors.Helpers;
using AiDotNet.Tensors.LinearAlgebra;
+using System.Threading;
namespace AiDotNet.Inference;
@@ -31,10 +37,15 @@ namespace AiDotNet.Inference;
///
///
/// The numeric type for computations.
-public class InferenceOptimizer
+internal class InferenceOptimizer
{
private readonly InferenceOptimizationConfig _config;
private KVCache? _kvCache;
+ private PagedKVCache? _pagedKVCache;
+ private PagedAttentionKernel? _pagedKernel;
+ private long? _pagedSequenceId;
+ private List>? _pagedAttentionLayers;
+ private static long s_nextPagedSequenceId = DateTime.UtcNow.Ticks;
private IDraftModel? _draftModel;
private SpeculativeDecoder? _speculativeDecoder;
private bool _isInitialized;
@@ -71,6 +82,57 @@ public InferenceOptimizer()
{
}
+ ///
+ /// Creates an inference-optimized model instance based on the current configuration.
+ ///
+ /// The neural network to optimize.
+ /// Whether to clone the model before applying layer-level rewrites.
+ /// The optimized model and whether any optimizations were applied.
+ ///
+ /// This method can apply stateless layer rewrites (e.g., MultiHeadAttention -> FlashAttentionLayer)
+ /// and then initialize stateful inference features (e.g., KV-cache) on the resulting model.
+ ///
+ public (NeuralNetworkBase OptimizedModel, bool AnyOptimizationsApplied) OptimizeForInference(
+ NeuralNetworkBase model,
+ bool cloneModel = true)
+ {
+ if (model == null)
+ throw new ArgumentNullException(nameof(model));
+
+ _config.Validate();
+
+ // Clone only when we might rewrite layers; otherwise keep original reference.
+ bool mayRewriteAttention = _config.EnableFlashAttention || _config.EnableKVCache || _config.EnableWeightOnlyQuantization;
+ var workingModel = model;
+ if (cloneModel && mayRewriteAttention && HasOptimizableAttentionLayers(model))
+ {
+ try
+ {
+ // NeuralNetworkBase.Clone performs a deep copy via serialization.
+ workingModel = (NeuralNetworkBase)model.Clone();
+ }
+ catch (Exception ex)
+ {
+ // Some layer types may not yet support serialization-based cloning.
+ // Do not mutate the user's original model; just skip optimizations.
+ Console.WriteLine($"Warning: model cloning failed for inference optimizations: {ex.Message}. Skipping inference optimizations for this model instance.");
+ InferenceDiagnostics.RecordException(
+ area: "InferenceOptimizer",
+ feature: "CloneForRewrite",
+ ex: ex,
+ reason: "Clone failed; skipping all inference optimizations to avoid mutating user model.");
+ return (model, false);
+ }
+ }
+
+ bool anyApplied = ApplyAttentionOptimizations(workingModel);
+ InferenceDiagnostics.RecordDecision("InferenceOptimizer", "AttentionRewrites", enabled: anyApplied, reason: anyApplied ? "Applied" : "NoApplicableLayersOrDisabled");
+ anyApplied |= ApplyWeightOnlyQuantization(workingModel);
+ anyApplied |= Initialize(workingModel);
+
+ return (workingModel, anyApplied);
+ }
+
///
/// Initializes inference optimizations for a neural network model.
///
@@ -91,12 +153,20 @@ public bool Initialize(NeuralNetworkBase model)
if (model == null)
throw new ArgumentNullException(nameof(model));
+ _config.Validate();
+
bool anyOptimizationsApplied = false;
// Find and configure attention layers for KV caching
if (_config.EnableKVCache)
{
- anyOptimizationsApplied |= InitializeKVCache(model);
+ anyOptimizationsApplied |= _config.EnablePagedKVCache
+ ? InitializePagedKVCache(model)
+ : InitializeKVCache(model);
+ }
+ else
+ {
+ InferenceDiagnostics.RecordDecision("InferenceOptimizer", "KVCache", enabled: false, reason: "DisabledByConfig");
}
// Initialize speculative decoding if enabled
@@ -104,6 +174,10 @@ public bool Initialize(NeuralNetworkBase model)
{
anyOptimizationsApplied |= InitializeSpeculativeDecoding(model);
}
+ else
+ {
+ InferenceDiagnostics.RecordDecision("InferenceOptimizer", "SpeculativeDecoding", enabled: false, reason: "DisabledByConfig");
+ }
_isInitialized = true;
return anyOptimizationsApplied;
@@ -138,7 +212,8 @@ private bool InitializeKVCache(NeuralNetworkBase model)
var firstLayer = attentionLayers[0];
int numHeads = firstLayer.HeadCount;
int headDim = firstLayer.HeadDimension;
- int maxSeqLen = EstimateMaxSequenceLength();
+ int numLayers = attentionLayers.Count;
+ int maxSeqLen = EstimateMaxSequenceLength(numLayers, numHeads, headDim);
// Create KV cache configuration
var cacheConfig = new KVCacheConfig
@@ -148,7 +223,12 @@ private bool InitializeKVCache(NeuralNetworkBase model)
HeadDimension = headDim,
MaxSequenceLength = maxSeqLen,
MaxBatchSize = _config.MaxBatchSize,
- PreAllocate = true
+ PreAllocate = true,
+ UseSlidingWindow = _config.UseSlidingWindowKVCache,
+ WindowSize = _config.UseSlidingWindowKVCache
+ ? Math.Min(_config.KVCacheWindowSize, maxSeqLen)
+ : 1024,
+ DataType = ResolveKVCacheDataType()
};
// Create and attach KV cache
@@ -164,21 +244,482 @@ private bool InitializeKVCache(NeuralNetworkBase model)
return true;
}
+ private CacheDataType ResolveKVCacheDataType()
+ {
+ bool fp16Capable = typeof(T) == typeof(float) || typeof(T) == typeof(double) || typeof(T) == typeof(Half);
+ bool int8Capable = fp16Capable;
+
+ CacheDataType resolved;
+ if (_config.KVCacheQuantization == KVCacheQuantizationMode.Int8 && int8Capable)
+ {
+ resolved = CacheDataType.Int8;
+ }
+ else
+ {
+ resolved = _config.KVCachePrecision switch
+ {
+ KVCachePrecisionMode.Float32 => CacheDataType.Float32,
+ KVCachePrecisionMode.Float16 => fp16Capable ? CacheDataType.Float16 : CacheDataType.Float32,
+ _ => fp16Capable ? CacheDataType.Float16 : CacheDataType.Float32
+ };
+ }
+
+ InferenceDiagnostics.RecordDecision(
+ area: "InferenceOptimizer",
+ feature: "KVCachePrecision",
+ enabled: resolved == CacheDataType.Float16 || resolved == CacheDataType.Int8,
+ reason: $"Precision={_config.KVCachePrecision};Quant={_config.KVCacheQuantization};Resolved={resolved};Type={typeof(T).Name}");
+
+ return resolved;
+ }
+
+ private bool InitializePagedKVCache(NeuralNetworkBase model)
+ {
+ var attentionLayers = new List>();
+ int layerIndex = 0;
+
+ foreach (var layer in model.Layers)
+ {
+ if (layer is PagedCachedMultiHeadAttention pagedAttention)
+ {
+ pagedAttention.LayerIndex = layerIndex;
+ attentionLayers.Add(pagedAttention);
+ layerIndex++;
+ }
+ }
+
+ if (attentionLayers.Count == 0)
+ {
+ // No paged attention layers present; fall back to contiguous cache if applicable.
+ return InitializeKVCache(model);
+ }
+
+ var firstLayer = attentionLayers[0];
+ int numHeads = firstLayer.HeadCount;
+ int headDim = firstLayer.HeadDimension;
+ int numLayers = attentionLayers.Count;
+
+ long availableBytes = (long)_config.KVCacheMaxSizeMB * 1024 * 1024;
+ int blockSize = _config.PagedKVCacheBlockSize;
+
+ _pagedKVCache = PagedKVCache.FromMemorySize(availableBytes, numLayers, numHeads, headDim, blockSize);
+ _pagedKernel = new PagedAttentionKernel(_pagedKVCache, new PagedAttentionConfig
+ {
+ NumHeads = numHeads,
+ HeadDimension = headDim,
+ BlockSize = blockSize,
+ MaxBatchSize = _config.MaxBatchSize
+ });
+
+ // Allocate a fresh sequence ID for this optimized model instance (one model == one sequence).
+ if (!TryAllocatePagedSequenceId(_pagedKVCache, initialTokens: 0, out long sequenceId))
+ {
+ InferenceDiagnostics.RecordDecision(
+ area: "InferenceOptimizer",
+ feature: "PagedKVCache",
+ enabled: false,
+ reason: "AllocateSequenceFailed(OutOfMemoryOrExhausted)");
+ _pagedKVCache = null;
+ _pagedKernel = null;
+ _pagedAttentionLayers = null;
+ _pagedSequenceId = null;
+ return false;
+ }
+
+ _pagedSequenceId = sequenceId;
+ _pagedAttentionLayers = attentionLayers;
+
+ foreach (var layer in attentionLayers)
+ {
+ layer.Kernel = _pagedKernel;
+ layer.SequenceId = sequenceId;
+ layer.InferenceMode = true;
+ }
+
+ return true;
+ }
+
+ private static bool TryAllocatePagedSequenceId(PagedKVCache cache, int initialTokens, out long sequenceId)
+ {
+ const int maxAttempts = 1024;
+ var spin = new SpinWait();
+
+ for (int attempt = 0; attempt < maxAttempts; attempt++)
+ {
+ sequenceId = Interlocked.Increment(ref s_nextPagedSequenceId);
+ if (cache.AllocateSequence(sequenceId, initialTokens))
+ {
+ return true;
+ }
+
+ spin.SpinOnce();
+ }
+
+ sequenceId = 0;
+ return false;
+ }
+
+ private static bool TryAllocatePagedSequenceId(PagedKVCache cache, long preferredId, int initialTokens, out long sequenceId)
+ {
+ if (cache.AllocateSequence(preferredId, initialTokens))
+ {
+ sequenceId = preferredId;
+ return true;
+ }
+
+ return TryAllocatePagedSequenceId(cache, initialTokens, out sequenceId);
+ }
+
+ private bool HasOptimizableAttentionLayers(NeuralNetworkBase model)
+ {
+ foreach (var layer in model.Layers)
+ {
+ if (layer is MultiHeadAttentionLayer || layer is FlashAttentionLayer || layer is SelfAttentionLayer)
+ return true;
+
+ if (_config.EnableWeightOnlyQuantization &&
+ typeof(T) == typeof(float) &&
+ layer is DenseLayer)
+ {
+ return true;
+ }
+ }
+
+ return false;
+ }
+
+ private bool ApplyAttentionOptimizations(NeuralNetworkBase model)
+ {
+ bool useCausalMask = ResolveCausalMask(model);
+ InferenceDiagnostics.RecordDecision(
+ area: "InferenceOptimizer",
+ feature: "CausalMask",
+ enabled: useCausalMask,
+ reason: _config.AttentionMasking == AttentionMaskingMode.Auto ? "Auto" : _config.AttentionMasking.ToString());
+
+ // KV-cache is only beneficial for incremental decoding patterns; default to enabling it only when causal masking applies.
+ bool enableKVCache = _config.EnableKVCache && useCausalMask;
+ bool enablePagedKVCache = enableKVCache && _config.EnablePagedKVCache;
+ bool enableFlashAttention = _config.EnableFlashAttention;
+
+ bool anyRewritten = false;
+
+ for (int i = 0; i < model.Layers.Count; i++)
+ {
+ var layer = model.Layers[i];
+
+ if (layer is SelfAttentionLayer selfAttention && (enableKVCache || enableFlashAttention))
+ {
+ var converted = TryConvertSelfAttentionToMultiHead(selfAttention);
+ if (converted != null)
+ {
+ model.Layers[i] = converted;
+ anyRewritten = true;
+
+ // Re-process this index under MultiHeadAttention rules.
+ i--;
+ continue;
+ }
+
+ InferenceDiagnostics.RecordDecision(
+ area: "InferenceOptimizer",
+ feature: "SelfAttentionRewrite",
+ enabled: false,
+ reason: "UnsupportedSelfAttentionLayer(HeadCountOrShape)");
+ continue;
+ }
+
+ if (layer is MultiHeadAttentionLayer mha)
+ {
+ var inputShape = mha.GetInputShape();
+ if (inputShape.Length < 2)
+ {
+ continue;
+ }
+
+ int seqLen = inputShape[0];
+ int embDim = inputShape[1];
+ int headCount = mha.HeadCount;
+ var activation = mha.ScalarActivation;
+
+ if (enableKVCache)
+ {
+ if (enablePagedKVCache)
+ {
+ var paged = new PagedCachedMultiHeadAttention(
+ sequenceLength: seqLen,
+ embeddingDimension: embDim,
+ headCount: headCount,
+ useCausalMask: useCausalMask,
+ activationFunction: activation);
+ paged.EnableWeightOnlyQuantization = _config.EnableWeightOnlyQuantization;
+ paged.SetParameters(mha.GetParameters());
+ model.Layers[i] = paged;
+ }
+ else
+ {
+ var cached = new CachedMultiHeadAttention(
+ sequenceLength: seqLen,
+ embeddingDimension: embDim,
+ headCount: headCount,
+ useFlashAttention: enableFlashAttention,
+ layerIndex: 0,
+ useCausalMask: useCausalMask,
+ activationFunction: activation);
+ cached.SetParameters(mha.GetParameters());
+ model.Layers[i] = cached;
+ }
+ anyRewritten = true;
+ continue;
+ }
+
+ if (enableFlashAttention)
+ {
+ var flashConfig = FlashAttentionConfig.Default;
+ flashConfig.UseCausalMask = useCausalMask;
+
+ var flashLayer = new FlashAttentionLayer(
+ sequenceLength: seqLen,
+ embeddingDimension: embDim,
+ headCount: headCount,
+ config: flashConfig,
+ activationFunction: activation);
+ flashLayer.SetParameters(mha.GetParameters());
+ model.Layers[i] = flashLayer;
+ anyRewritten = true;
+ }
+
+ continue;
+ }
+
+ if (layer is FlashAttentionLayer flash && enableKVCache)
+ {
+ var inputShape = flash.GetInputShape();
+ if (inputShape.Length < 2)
+ {
+ continue;
+ }
+
+ int seqLen = inputShape[0];
+ int embDim = inputShape[1];
+ int headCount = flash.HeadCount;
+ var activation = flash.ScalarActivation;
+
+ if (enablePagedKVCache)
+ {
+ var paged = new PagedCachedMultiHeadAttention(
+ sequenceLength: seqLen,
+ embeddingDimension: embDim,
+ headCount: headCount,
+ useCausalMask: useCausalMask,
+ activationFunction: activation);
+ paged.EnableWeightOnlyQuantization = _config.EnableWeightOnlyQuantization;
+ paged.SetParameters(flash.GetParameters());
+ model.Layers[i] = paged;
+ }
+ else
+ {
+ var cached = new CachedMultiHeadAttention(
+ sequenceLength: seqLen,
+ embeddingDimension: embDim,
+ headCount: headCount,
+ useFlashAttention: enableFlashAttention,
+ layerIndex: 0,
+ useCausalMask: useCausalMask,
+ activationFunction: activation);
+ cached.SetParameters(flash.GetParameters());
+ model.Layers[i] = cached;
+ }
+ anyRewritten = true;
+ }
+ }
+
+ return anyRewritten;
+ }
+
+ private bool ApplyWeightOnlyQuantization(NeuralNetworkBase model)
+ {
+ if (!_config.EnableWeightOnlyQuantization)
+ {
+ InferenceDiagnostics.RecordDecision("InferenceOptimizer", "WeightOnlyQuantization", enabled: false, reason: "DisabledByConfig");
+ return false;
+ }
+
+ if (typeof(T) != typeof(float))
+ {
+ InferenceDiagnostics.RecordDecision("InferenceOptimizer", "WeightOnlyQuantization", enabled: false, reason: $"UnsupportedType({typeof(T).Name})");
+ return false;
+ }
+
+ bool any = false;
+ for (int i = 0; i < model.Layers.Count; i++)
+ {
+ if (model.Layers[i] is DenseLayer dense)
+ {
+ try
+ {
+ var replacement = dense.VectorActivation != null
+ ? new QuantizedDenseLayer(dense, dense.VectorActivation)
+ : new QuantizedDenseLayer(dense);
+
+ model.Layers[i] = (AiDotNet.Interfaces.ILayer)(object)replacement;
+ any = true;
+ }
+ catch (Exception ex)
+ {
+ InferenceDiagnostics.RecordException("InferenceOptimizer", "WeightOnlyQuantization", ex, "DenseLayerQuantizationFailed;FallbackToFP");
+ }
+ }
+ }
+
+ InferenceDiagnostics.RecordDecision("InferenceOptimizer", "WeightOnlyQuantization", enabled: any, reason: any ? "Applied(DenseLayer)" : "NoApplicableLayers");
+ return any;
+ }
+
+ private MultiHeadAttentionLayer? TryConvertSelfAttentionToMultiHead(SelfAttentionLayer layer)
+ {
+ var inputShape = layer.GetInputShape();
+ if (inputShape.Length < 2)
+ return null;
+
+ int seqLen = inputShape[0];
+ int embDim = inputShape[1];
+ if (seqLen <= 0 || embDim <= 0)
+ return null;
+
+ int headCount = TryGetHeadCountFromMetadata(layer) ?? 0;
+ if (headCount <= 0)
+ return null;
+
+ if (embDim % headCount != 0)
+ return null;
+
+ // SelfAttentionLayer has Q/K/V projections plus bias, but no output projection.
+ // We convert it into a MultiHeadAttentionLayer with an identity output projection so that
+ // downstream inference rewrites (FlashAttention / KV-cache) can be applied consistently.
+ var activation = layer.ScalarActivation;
+ var mha = new MultiHeadAttentionLayer(seqLen, embDim, headCount, activationFunction: activation);
+
+ var selfParams = layer.GetParameters();
+ int projSize = embDim * embDim;
+ int expectedSelf = (3 * projSize) + embDim;
+ if (selfParams.Length != expectedSelf)
+ return null;
+
+ var numOps = MathHelper.GetNumericOperations();
+
+ // MultiHead params: Q, K, V, O, bias
+ var combined = new Vector((4 * projSize) + embDim);
+ int idx = 0;
+
+ // Copy Q/K/V (3 * projSize)
+ for (int i = 0; i < 3 * projSize; i++)
+ combined[idx++] = selfParams[i];
+
+ // Output weights: identity matrix (embDim x embDim) flattened row-major
+ for (int r = 0; r < embDim; r++)
+ {
+ for (int c = 0; c < embDim; c++)
+ {
+ combined[idx++] = r == c ? numOps.One : numOps.Zero;
+ }
+ }
+
+ // Output bias (embDim)
+ for (int i = 0; i < embDim; i++)
+ combined[idx++] = selfParams[(3 * projSize) + i];
+
+ mha.SetParameters(combined);
+ return mha;
+ }
+
+ private static int? TryGetHeadCountFromMetadata(ILayer layer)
+ {
+ if (layer is not LayerBase layerBase)
+ return null;
+
+ if (!layerBase.GetMetadata().TryGetValue("HeadCount", out var raw) || string.IsNullOrWhiteSpace(raw))
+ return null;
+
+ return int.TryParse(raw, out var parsed) ? parsed : null;
+ }
+
+ private bool ResolveCausalMask(NeuralNetworkBase model)
+ {
+ return _config.AttentionMasking switch
+ {
+ AttentionMaskingMode.Causal => true,
+ AttentionMaskingMode.Disabled => false,
+ _ => InferCausalFromModel(model)
+ };
+ }
+
+ private bool InferCausalFromModel(NeuralNetworkBase model)
+ {
+ // Default to causal when the user enables generation-oriented inference features.
+ // This matches industry-standard expectations for autoregressive decoding and avoids
+ // relying on users to set TaskType explicitly.
+ if (_config.EnableKVCache || _config.EnableSpeculativeDecoding)
+ return true;
+
+ // Otherwise, keep heuristics conservative to avoid changing semantics for non-generative models.
+ return model.Architecture.TaskType == NeuralNetworkTaskType.TextGeneration;
+ }
+
///
/// Estimates the maximum sequence length based on config and memory constraints.
///
- private int EstimateMaxSequenceLength()
+ /// Number of attention layers in the model.
+ /// Number of attention heads per layer.
+ /// Dimension of each attention head.
+ /// Maximum sequence length that fits within the configured memory budget.
+ private int EstimateMaxSequenceLength(int numLayers, int numHeads, int headDim)
{
- // Calculate based on available memory
- // Formula: maxSeqLen = (maxMemoryMB * 1024 * 1024) / (numLayers * numHeads * headDim * 2 * bytesPerElement)
- // Using a simplified estimate
+ // KV cache memory per token = numLayers * numHeads * headDim * 2 (K and V) * bytesPerElement
+ // For batch size, multiply by maxBatchSize
+ // Total: maxSeqLen * numLayers * numHeads * headDim * 2 * bytesPerElement * batchSize <= maxMemoryBytes
+
long maxMemoryBytes = (long)_config.KVCacheMaxSizeMB * 1024 * 1024;
- // Default reasonable sequence length
- const int defaultMaxSeqLen = 2048;
+ // Estimate bytes per element based on type T
+ int bytesPerElement = EstimateBytesPerElement();
+
+ // Memory per token per batch item = numLayers * numHeads * headDim * 2 * bytesPerElement
+ long memoryPerToken = (long)numLayers * numHeads * headDim * 2 * bytesPerElement;
+
+ // Account for batch size
+ long memoryPerTokenWithBatch = memoryPerToken * _config.MaxBatchSize;
+
+ // Prevent division by zero
+ if (memoryPerTokenWithBatch <= 0)
+ {
+ return 2048; // Reasonable default
+ }
+
+ // Calculate maximum sequence length
+ long calculatedMaxSeqLen = maxMemoryBytes / memoryPerTokenWithBatch;
- // Cap at reasonable maximum
- return Math.Min(defaultMaxSeqLen, 8192);
+ // Apply reasonable bounds using MathHelper.Clamp for net471 compatibility
+ const long minSeqLen = 128;
+ const long maxSeqLen = 32768; // Reasonable upper bound
+
+ return (int)MathHelper.Clamp(calculatedMaxSeqLen, minSeqLen, maxSeqLen);
+ }
+
+ ///
+ /// Estimates bytes per element based on the generic type T.
+ ///
+ private static int EstimateBytesPerElement()
+ {
+ // Common numeric types used in neural networks
+ var type = typeof(T);
+ if (type == typeof(float)) return 4;
+ if (type == typeof(double)) return 8;
+ if (type == typeof(Half)) return 2;
+ if (type == typeof(decimal)) return 16;
+
+ // Default to float size if unknown
+ return 4;
}
///
@@ -186,23 +727,56 @@ private int EstimateMaxSequenceLength()
///
private bool InitializeSpeculativeDecoding(NeuralNetworkBase model)
{
- // Create draft model based on configuration
- IDraftModel? draftModel = _config.DraftModelType switch
+ // Facade-friendly behavior: speculative decoding configuration must never crash inference.
+ // If a requested draft model is unavailable, fall back to an N-gram draft model and record diagnostics.
+ try
{
- DraftModelType.NGram => CreateNGramDraftModel(),
- DraftModelType.SmallNeural => CreateNeuralDraftModel(model),
- _ => null
- };
+ // For Custom draft models, an internal caller can provide one via SetCustomDraftModel().
+ if (_config.DraftModelType == DraftModelType.Custom)
+ {
+ if (_draftModel != null)
+ {
+ InferenceDiagnostics.RecordDecision("InferenceOptimizer", "SpeculativeDraftModel", enabled: true, reason: "CustomProvided");
+ return true;
+ }
+
+ _draftModel = CreateNGramDraftModel();
+ InferenceDiagnostics.RecordDecision("InferenceOptimizer", "SpeculativeDraftModel", enabled: _draftModel != null, reason: "CustomNotProvided_FallbackToNGram");
+ return _draftModel != null;
+ }
- if (draftModel == null)
+ IDraftModel? draftModel = _config.DraftModelType switch
+ {
+ DraftModelType.NGram => CreateNGramDraftModel(),
+ DraftModelType.SmallNeural => CreateNeuralDraftModel(model),
+ _ => CreateNGramDraftModel()
+ };
+
+ _draftModel = draftModel ?? CreateNGramDraftModel();
+ InferenceDiagnostics.RecordDecision(
+ "InferenceOptimizer",
+ "SpeculativeDraftModel",
+ enabled: _draftModel != null,
+ reason: draftModel != null ? _config.DraftModelType.ToString() : $"Unavailable({_config.DraftModelType})_FallbackToNGram");
+
+ return _draftModel != null;
+ }
+ catch (Exception ex)
{
- return false;
+ InferenceDiagnostics.RecordException("InferenceOptimizer", "SpeculativeDecoding", ex, "Draft model init failed; falling back to NGram.");
+ try
+ {
+ _draftModel = CreateNGramDraftModel();
+ InferenceDiagnostics.RecordDecision("InferenceOptimizer", "SpeculativeDraftModel", enabled: _draftModel != null, reason: "ExceptionFallbackToNGram");
+ return _draftModel != null;
+ }
+ catch
+ {
+ InferenceDiagnostics.RecordDecision("InferenceOptimizer", "SpeculativeDraftModel", enabled: false, reason: "FallbackFailed");
+ _draftModel = null;
+ return false;
+ }
}
-
- // Note: SpeculativeDecoder requires a target forward function
- // This will be set when actually doing inference via CreateSpeculativeDecoder
- _draftModel = draftModel;
- return true;
}
///
@@ -217,10 +791,19 @@ private bool InitializeSpeculativeDecoding(NeuralNetworkBase model)
///
/// Creates a small neural network draft model.
///
+ ///
+ /// SmallNeural draft models require a pre-trained companion model that is smaller
+ /// and faster than the target model but trained on similar data. This cannot be
+ /// automatically generated from the target model.
+ ///
+ ///
+ /// Always thrown because SmallNeural draft models require external pre-trained models.
+ ///
private IDraftModel? CreateNeuralDraftModel(NeuralNetworkBase model)
{
- // For neural draft models, we would need a pre-trained smaller model
- // This is a placeholder - in production, this would load a companion model
+ // SmallNeural draft models require a separate pre-trained smaller model. We do not expose
+ // draft model wiring via the public facade in the MVP, so treat this as unavailable.
+ InferenceDiagnostics.RecordDecision("InferenceOptimizer", "SpeculativeDraftModel", enabled: false, reason: "SmallNeuralUnavailable_FallbackToNGram");
return null;
}
@@ -258,6 +841,10 @@ public void DisableInferenceMode(NeuralNetworkBase model)
{
cachedAttention.InferenceMode = false;
}
+ else if (layer is PagedCachedMultiHeadAttention pagedAttention)
+ {
+ pagedAttention.InferenceMode = false;
+ }
}
}
@@ -267,6 +854,54 @@ public void DisableInferenceMode(NeuralNetworkBase model)
public void ClearCache()
{
_kvCache?.Clear();
+ if (_pagedKVCache != null && _pagedSequenceId.HasValue)
+ {
+ try
+ {
+ _pagedKVCache.FreeSequence(_pagedSequenceId.Value);
+ }
+ catch
+ {
+ // Best-effort cleanup.
+ }
+
+ // Re-allocate with the same ID if possible; otherwise allocate a new one.
+ if (!TryAllocatePagedSequenceId(_pagedKVCache, _pagedSequenceId.Value, initialTokens: 0, out long allocated))
+ {
+ InferenceDiagnostics.RecordDecision(
+ area: "InferenceOptimizer",
+ feature: "PagedKVCache",
+ enabled: false,
+ reason: "ClearCacheAllocateSequenceFailed(OutOfMemoryOrExhausted)");
+
+ // Safe fallback: disable paged inference mode on layers and keep session alive.
+ if (_pagedAttentionLayers != null)
+ {
+ foreach (var layer in _pagedAttentionLayers)
+ {
+ layer.InferenceMode = false;
+ layer.Kernel = null;
+ layer.ResetState();
+ }
+ }
+
+ _pagedSequenceId = null;
+ return;
+ }
+
+ _pagedSequenceId = allocated;
+
+ if (_pagedAttentionLayers != null && _pagedSequenceId.HasValue)
+ {
+ foreach (var layer in _pagedAttentionLayers)
+ {
+ layer.SequenceId = _pagedSequenceId.Value;
+ layer.ResetState();
+ layer.InferenceMode = true;
+ layer.Kernel ??= _pagedKernel;
+ }
+ }
+ }
}
///
@@ -280,7 +915,10 @@ public Dictionary GetStatistics()
["IsInitialized"] = _isInitialized,
["KVCacheEnabled"] = _config.EnableKVCache,
["SpeculativeDecodingEnabled"] = _config.EnableSpeculativeDecoding,
- ["BatchingEnabled"] = _config.EnableBatching
+ ["BatchingEnabled"] = _config.EnableBatching,
+ ["PagedKVCacheInitialized"] = _pagedKVCache != null,
+ ["PagedAttentionLayerCount"] = _pagedAttentionLayers?.Count ?? 0,
+ ["PagedAttentionWeightOnlyQuantizationEnabled"] = _pagedAttentionLayers?.Any(l => l.EnableWeightOnlyQuantization) ?? false
};
if (_kvCache != null)
@@ -310,6 +948,34 @@ public Dictionary GetStatistics()
///
public IDraftModel? DraftModel => _draftModel;
+ ///
+ /// Sets a custom draft model for speculative decoding.
+ ///
+ /// The custom draft model implementation.
+ ///
+ /// For Beginners: Use this method when you have your own draft model implementation.
+ ///
+ /// This is required when using DraftModelType.Custom or when you want to replace the
+ /// default NGram draft model with a more sophisticated model.
+ ///
+ /// Your custom draft model must implement IDraftModel<T> and provide:
+ /// - Draft token generation
+ /// - Probability estimation for speculative decoding verification
+ ///
+ /// Example:
+ ///
+ /// var optimizer = new InferenceOptimizer<float>(config);
+ /// optimizer.SetCustomDraftModel(myCustomDraftModel);
+ /// optimizer.Initialize(mainModel);
+ ///
+ ///
+ ///
+ /// Thrown when draftModel is null.
+ public void SetCustomDraftModel(IDraftModel draftModel)
+ {
+ _draftModel = draftModel ?? throw new ArgumentNullException(nameof(draftModel));
+ }
+
///
/// Creates a speculative decoder with the given target forward function.
///
@@ -339,7 +1005,13 @@ public Dictionary GetStatistics()
var speculativeConfig = new SpeculativeDecodingConfig
{
NumDraftTokens = _config.SpeculationDepth,
- UseTreeSpeculation = _config.UseTreeSpeculation
+ UseTreeSpeculation = _config.UseTreeSpeculation ||
+ _config.SpeculativeMethod == SpeculativeMethod.Medusa ||
+ _config.SpeculativeMethod == SpeculativeMethod.Eagle,
+ AdaptiveDraftLength = _config.SpeculationPolicy == SpeculationPolicy.Auto,
+ TreeBranchFactor = _config.SpeculativeMethod == SpeculativeMethod.Medusa ? 4 : 2,
+ MaxTreeDepth = Math.Max(1, _config.SpeculationDepth),
+ MinAcceptanceRate = MathHelper.GetNumericOperations().FromDouble(0.5)
};
_speculativeDecoder = new SpeculativeDecoder(_draftModel, targetForward, speculativeConfig);
diff --git a/src/Inference/KVCache.cs b/src/Inference/KVCache.cs
index b391faa38e..bbf017f1c5 100644
--- a/src/Inference/KVCache.cs
+++ b/src/Inference/KVCache.cs
@@ -28,7 +28,7 @@ namespace AiDotNet.Inference;
///
///
/// The numeric type for cache storage (typically float or double).
-public class KVCache
+internal class KVCache
{
private static readonly INumericOperations NumOps = MathHelper.GetNumericOperations();
@@ -38,8 +38,24 @@ public class KVCache
private readonly Tensor[] _keyCache;
private readonly Tensor[] _valueCache;
- // Current sequence length for each batch item
- private readonly int[] _sequenceLengths;
+ // Optional FP16 cache storage (used when Config.DataType == Float16 and T is float/double)
+ private readonly Tensor[]? _keyCacheFp16;
+ private readonly Tensor[]? _valueCacheFp16;
+ private readonly bool _useFp16Storage;
+ private readonly Func? _toHalf;
+ private readonly Func? _fromHalf;
+
+ // Optional int8 quantized cache storage (used when Config.DataType == Int8).
+ private readonly Tensor[]? _keyCacheInt8;
+ private readonly Tensor[]? _valueCacheInt8;
+ private readonly bool _useInt8Storage;
+ private readonly float[]? _keyScaleInt8;
+ private readonly float[]? _valueScaleInt8;
+ private readonly Func? _toFloat;
+ private readonly Func? _fromFloat;
+
+ // Current sequence length for each layer and batch item: [layer][batch]
+ private readonly int[][] _sequenceLengths;
// Statistics
private long _cacheHits;
@@ -54,7 +70,7 @@ public class KVCache
///
/// Gets the current number of cached tokens for batch item 0.
///
- public int CurrentLength => _sequenceLengths[0];
+ public int CurrentLength => _sequenceLengths.Length > 0 ? _sequenceLengths[0][0] : 0;
///
/// Gets the maximum sequence length this cache can hold.
@@ -86,7 +102,66 @@ public KVCache(KVCacheConfig config)
_keyCache = new Tensor[config.NumLayers];
_valueCache = new Tensor[config.NumLayers];
- _sequenceLengths = new int[config.MaxBatchSize];
+
+ if (config.DataType == CacheDataType.Int8)
+ {
+ // Only enable int8 storage when we can safely convert between T and float.
+ if (typeof(T) == typeof(float))
+ {
+ _useInt8Storage = true;
+ _toFloat = value => (float)(object)value!;
+ _fromFloat = value => (T)(object)value;
+ }
+ else if (typeof(T) == typeof(double))
+ {
+ _useInt8Storage = true;
+ _toFloat = value => (float)(double)(object)value!;
+ _fromFloat = value => (T)(object)(double)value;
+ }
+ else if (typeof(T) == typeof(Half))
+ {
+ _useInt8Storage = true;
+ _toFloat = value => (float)(Half)(object)value!;
+ _fromFloat = value => (T)(object)(Half)value;
+ }
+ }
+
+ if (config.DataType == CacheDataType.Float16 && typeof(T) != typeof(Half))
+ {
+ // Only enable FP16 storage when we can safely convert between T and Half.
+ if (typeof(T) == typeof(float))
+ {
+ _useFp16Storage = true;
+ _toHalf = value => (Half)(float)(object)value!;
+ _fromHalf = value => (T)(object)(float)value;
+ }
+ else if (typeof(T) == typeof(double))
+ {
+ _useFp16Storage = true;
+ _toHalf = value => (Half)(double)(object)value!;
+ _fromHalf = value => (T)(object)(double)(float)value;
+ }
+ }
+
+ if (_useFp16Storage)
+ {
+ _keyCacheFp16 = new Tensor[config.NumLayers];
+ _valueCacheFp16 = new Tensor[config.NumLayers];
+ }
+
+ if (_useInt8Storage)
+ {
+ _keyCacheInt8 = new Tensor[config.NumLayers];
+ _valueCacheInt8 = new Tensor[config.NumLayers];
+ _keyScaleInt8 = new float[config.NumLayers];
+ _valueScaleInt8 = new float[config.NumLayers];
+ }
+
+ _sequenceLengths = new int[config.NumLayers][];
+ for (int layer = 0; layer < config.NumLayers; layer++)
+ {
+ _sequenceLengths[layer] = new int[config.MaxBatchSize];
+ }
if (config.PreAllocate)
{
@@ -121,8 +196,23 @@ private void AllocateCaches()
for (int layer = 0; layer < _config.NumLayers; layer++)
{
- _keyCache[layer] = new Tensor(shape);
- _valueCache[layer] = new Tensor(shape);
+ if (_useInt8Storage)
+ {
+ _keyCacheInt8![layer] = new Tensor(shape);
+ _valueCacheInt8![layer] = new Tensor(shape);
+ _keyScaleInt8![layer] = 0f;
+ _valueScaleInt8![layer] = 0f;
+ }
+ else if (_useFp16Storage)
+ {
+ _keyCacheFp16![layer] = new Tensor(shape);
+ _valueCacheFp16![layer] = new Tensor(shape);
+ }
+ else
+ {
+ _keyCache[layer] = new Tensor(shape);
+ _valueCache[layer] = new Tensor(shape);
+ }
}
}
@@ -166,10 +256,15 @@ private void AllocateCaches()
HandleSlidingWindowEviction(layerIndex, batchSize, newSeqLen);
}
+ if (_useInt8Storage)
+ {
+ EnsureInt8Scales(layerIndex, newKeys, newValues, batchSize, newSeqLen);
+ }
+
// Append new entries
for (int b = 0; b < batchSize; b++)
{
- int currentLen = _sequenceLengths[b];
+ int currentLen = _sequenceLengths[layerIndex][b];
int newLen = currentLen + newSeqLen;
if (newLen > _config.MaxSequenceLength)
@@ -187,13 +282,28 @@ private void AllocateCaches()
int targetPos = currentLen + s;
for (int d = 0; d < _config.HeadDimension; d++)
{
- _keyCache[layerIndex][new[] { b, h, targetPos, d }] = newKeys[new[] { b, h, s, d }];
- _valueCache[layerIndex][new[] { b, h, targetPos, d }] = newValues[new[] { b, h, s, d }];
+ if (_useInt8Storage)
+ {
+ var keyScale = _keyScaleInt8![layerIndex];
+ var valueScale = _valueScaleInt8![layerIndex];
+ _keyCacheInt8![layerIndex][new[] { b, h, targetPos, d }] = QuantizeToInt8(_toFloat!(newKeys[new[] { b, h, s, d }]), keyScale);
+ _valueCacheInt8![layerIndex][new[] { b, h, targetPos, d }] = QuantizeToInt8(_toFloat!(newValues[new[] { b, h, s, d }]), valueScale);
+ }
+ else if (_useFp16Storage)
+ {
+ _keyCacheFp16![layerIndex][new[] { b, h, targetPos, d }] = _toHalf!(newKeys[new[] { b, h, s, d }]);
+ _valueCacheFp16![layerIndex][new[] { b, h, targetPos, d }] = _toHalf!(newValues[new[] { b, h, s, d }]);
+ }
+ else
+ {
+ _keyCache[layerIndex][new[] { b, h, targetPos, d }] = newKeys[new[] { b, h, s, d }];
+ _valueCache[layerIndex][new[] { b, h, targetPos, d }] = newValues[new[] { b, h, s, d }];
+ }
}
}
}
- _sequenceLengths[b] = newLen;
+ _sequenceLengths[layerIndex][b] = newLen;
_cacheMisses += newSeqLen;
}
@@ -211,7 +321,7 @@ private void AllocateCaches()
{
ValidateLayerIndex(layerIndex);
- if (_keyCache[layerIndex] == null)
+ if (!IsLayerAllocated(layerIndex))
{
throw new InvalidOperationException($"Layer {layerIndex} cache not initialized. Call Append first.");
}
@@ -220,7 +330,7 @@ private void AllocateCaches()
int maxLen = 0;
for (int b = 0; b < batchSize; b++)
{
- if (_sequenceLengths[b] > maxLen) maxLen = _sequenceLengths[b];
+ if (_sequenceLengths[layerIndex][b] > maxLen) maxLen = _sequenceLengths[layerIndex][b];
}
if (maxLen == 0)
@@ -238,15 +348,30 @@ private void AllocateCaches()
// Copy cached values
for (int b = 0; b < batchSize; b++)
{
- int seqLen = _sequenceLengths[b];
+ int seqLen = _sequenceLengths[layerIndex][b];
for (int h = 0; h < _config.NumHeads; h++)
{
for (int s = 0; s < seqLen; s++)
{
for (int d = 0; d < _config.HeadDimension; d++)
{
- keys[new[] { b, h, s, d }] = _keyCache[layerIndex][new[] { b, h, s, d }];
- values[new[] { b, h, s, d }] = _valueCache[layerIndex][new[] { b, h, s, d }];
+ if (_useInt8Storage)
+ {
+ float keyScale = _keyScaleInt8![layerIndex];
+ float valueScale = _valueScaleInt8![layerIndex];
+ keys[new[] { b, h, s, d }] = _fromFloat!(DequantizeInt8(_keyCacheInt8![layerIndex][new[] { b, h, s, d }], keyScale));
+ values[new[] { b, h, s, d }] = _fromFloat!(DequantizeInt8(_valueCacheInt8![layerIndex][new[] { b, h, s, d }], valueScale));
+ }
+ else if (_useFp16Storage)
+ {
+ keys[new[] { b, h, s, d }] = _fromHalf!(_keyCacheFp16![layerIndex][new[] { b, h, s, d }]);
+ values[new[] { b, h, s, d }] = _fromHalf!(_valueCacheFp16![layerIndex][new[] { b, h, s, d }]);
+ }
+ else
+ {
+ keys[new[] { b, h, s, d }] = _keyCache[layerIndex][new[] { b, h, s, d }];
+ values[new[] { b, h, s, d }] = _valueCache[layerIndex][new[] { b, h, s, d }];
+ }
}
}
}
@@ -270,6 +395,12 @@ public void Update(int layerIndex, int[] positions, Tensor keys, Tensor va
int batchSize = keys.Shape[0];
int numPositions = positions.Length;
+ if (_useInt8Storage)
+ {
+ EnsureCacheAllocated(layerIndex);
+ EnsureInt8Scales(layerIndex, keys, values, batchSize, numPositions);
+ }
+
for (int b = 0; b < batchSize; b++)
{
for (int p = 0; p < numPositions; p++)
@@ -285,8 +416,23 @@ public void Update(int layerIndex, int[] positions, Tensor keys, Tensor va
{
for (int d = 0; d < _config.HeadDimension; d++)
{
- _keyCache[layerIndex][new[] { b, h, pos, d }] = keys[new[] { b, h, p, d }];
- _valueCache[layerIndex][new[] { b, h, pos, d }] = values[new[] { b, h, p, d }];
+ if (_useInt8Storage)
+ {
+ var keyScale = _keyScaleInt8![layerIndex];
+ var valueScale = _valueScaleInt8![layerIndex];
+ _keyCacheInt8![layerIndex][new[] { b, h, pos, d }] = QuantizeToInt8(_toFloat!(keys[new[] { b, h, p, d }]), keyScale);
+ _valueCacheInt8![layerIndex][new[] { b, h, pos, d }] = QuantizeToInt8(_toFloat!(values[new[] { b, h, p, d }]), valueScale);
+ }
+ else if (_useFp16Storage)
+ {
+ _keyCacheFp16![layerIndex][new[] { b, h, pos, d }] = _toHalf!(keys[new[] { b, h, p, d }]);
+ _valueCacheFp16![layerIndex][new[] { b, h, pos, d }] = _toHalf!(values[new[] { b, h, p, d }]);
+ }
+ else
+ {
+ _keyCache[layerIndex][new[] { b, h, pos, d }] = keys[new[] { b, h, p, d }];
+ _valueCache[layerIndex][new[] { b, h, pos, d }] = values[new[] { b, h, p, d }];
+ }
}
}
}
@@ -307,18 +453,25 @@ public void Truncate(int newLength, int batchIndex = -1)
if (batchIndex == -1)
{
- for (int b = 0; b < _sequenceLengths.Length; b++)
+ for (int layer = 0; layer < _sequenceLengths.Length; layer++)
{
- _sequenceLengths[b] = Math.Min(_sequenceLengths[b], newLength);
+ for (int b = 0; b < _sequenceLengths[layer].Length; b++)
+ {
+ _sequenceLengths[layer][b] = Math.Min(_sequenceLengths[layer][b], newLength);
+ }
}
}
else
{
- if (batchIndex < 0 || batchIndex >= _sequenceLengths.Length)
+ if (batchIndex < 0 || (_sequenceLengths.Length > 0 && batchIndex >= _sequenceLengths[0].Length))
{
throw new ArgumentOutOfRangeException(nameof(batchIndex));
}
- _sequenceLengths[batchIndex] = Math.Min(_sequenceLengths[batchIndex], newLength);
+
+ for (int layer = 0; layer < _sequenceLengths.Length; layer++)
+ {
+ _sequenceLengths[layer][batchIndex] = Math.Min(_sequenceLengths[layer][batchIndex], newLength);
+ }
}
}
@@ -327,9 +480,18 @@ public void Truncate(int newLength, int batchIndex = -1)
///
public void Clear()
{
- for (int b = 0; b < _sequenceLengths.Length; b++)
+ for (int layer = 0; layer < _sequenceLengths.Length; layer++)
{
- _sequenceLengths[b] = 0;
+ for (int b = 0; b < _sequenceLengths[layer].Length; b++)
+ {
+ _sequenceLengths[layer][b] = 0;
+ }
+
+ if (_useInt8Storage)
+ {
+ _keyScaleInt8![layer] = 0f;
+ _valueScaleInt8![layer] = 0f;
+ }
}
// Reset statistics
@@ -343,11 +505,15 @@ public void Clear()
///
public void Clear(int batchIndex)
{
- if (batchIndex < 0 || batchIndex >= _sequenceLengths.Length)
+ if (batchIndex < 0 || (_sequenceLengths.Length > 0 && batchIndex >= _sequenceLengths[0].Length))
{
throw new ArgumentOutOfRangeException(nameof(batchIndex));
}
- _sequenceLengths[batchIndex] = 0;
+
+ for (int layer = 0; layer < _sequenceLengths.Length; layer++)
+ {
+ _sequenceLengths[layer][batchIndex] = 0;
+ }
}
///
@@ -355,11 +521,12 @@ public void Clear(int batchIndex)
///
public int GetSequenceLength(int batchIndex = 0)
{
- if (batchIndex < 0 || batchIndex >= _sequenceLengths.Length)
+ if (batchIndex < 0 || (_sequenceLengths.Length > 0 && batchIndex >= _sequenceLengths[0].Length))
{
throw new ArgumentOutOfRangeException(nameof(batchIndex));
}
- return _sequenceLengths[batchIndex];
+
+ return _sequenceLengths.Length > 0 ? _sequenceLengths[0][batchIndex] : 0;
}
///
@@ -370,14 +537,26 @@ public long GetCurrentMemoryUsage()
long totalElements = 0;
for (int layer = 0; layer < _config.NumLayers; layer++)
{
- if (_keyCache[layer] != null)
+ if (IsLayerAllocated(layer))
{
- totalElements += _keyCache[layer].Length + _valueCache[layer].Length;
+ if (_useInt8Storage)
+ {
+ totalElements += _keyCacheInt8![layer].Length + _valueCacheInt8![layer].Length;
+ }
+ else if (_useFp16Storage)
+ {
+ totalElements += _keyCacheFp16![layer].Length + _valueCacheFp16![layer].Length;
+ }
+ else
+ {
+ totalElements += _keyCache[layer].Length + _valueCache[layer].Length;
+ }
}
}
int bytesPerElement = _config.DataType switch
{
+ CacheDataType.Int8 => 1,
CacheDataType.Float16 => 2,
CacheDataType.Float32 => 4,
CacheDataType.Float64 => 8,
@@ -395,6 +574,9 @@ public Dictionary GetStatistics()
{
return new Dictionary
{
+ ["DataType"] = _config.DataType.ToString(),
+ ["UseInt8Storage"] = _useInt8Storage,
+ ["UseFp16Storage"] = _useFp16Storage,
["CacheHits"] = _cacheHits,
["CacheMisses"] = _cacheMisses,
["Evictions"] = _evictions,
@@ -403,7 +585,9 @@ public Dictionary GetStatistics()
: 0.0,
["CurrentMemoryMB"] = GetCurrentMemoryUsage() / (1024.0 * 1024.0),
["MaxMemoryMB"] = _config.EstimateMemoryBytes() / (1024.0 * 1024.0),
- ["SequenceLengths"] = _sequenceLengths.ToArray()
+ ["SequenceLengths"] = _sequenceLengths.Length > 0
+ ? _sequenceLengths[0].ToArray()
+ : Array.Empty()
};
}
@@ -417,11 +601,11 @@ public void CopyBatchState(int sourceBatch, int destBatch)
if (destBatch < 0 || destBatch >= _config.MaxBatchSize)
throw new ArgumentOutOfRangeException(nameof(destBatch));
- int seqLen = _sequenceLengths[sourceBatch];
-
for (int layer = 0; layer < _config.NumLayers; layer++)
{
- if (_keyCache[layer] == null) continue;
+ if (!IsLayerAllocated(layer)) continue;
+
+ int seqLen = _sequenceLengths[layer][sourceBatch];
for (int h = 0; h < _config.NumHeads; h++)
{
@@ -429,16 +613,130 @@ public void CopyBatchState(int sourceBatch, int destBatch)
{
for (int d = 0; d < _config.HeadDimension; d++)
{
- _keyCache[layer][new[] { destBatch, h, s, d }] =
- _keyCache[layer][new[] { sourceBatch, h, s, d }];
- _valueCache[layer][new[] { destBatch, h, s, d }] =
- _valueCache[layer][new[] { sourceBatch, h, s, d }];
+ if (_useInt8Storage)
+ {
+ _keyCacheInt8![layer][new[] { destBatch, h, s, d }] =
+ _keyCacheInt8![layer][new[] { sourceBatch, h, s, d }];
+ _valueCacheInt8![layer][new[] { destBatch, h, s, d }] =
+ _valueCacheInt8![layer][new[] { sourceBatch, h, s, d }];
+ }
+ else if (_useFp16Storage)
+ {
+ _keyCacheFp16![layer][new[] { destBatch, h, s, d }] =
+ _keyCacheFp16![layer][new[] { sourceBatch, h, s, d }];
+ _valueCacheFp16![layer][new[] { destBatch, h, s, d }] =
+ _valueCacheFp16![layer][new[] { sourceBatch, h, s, d }];
+ }
+ else
+ {
+ _keyCache[layer][new[] { destBatch, h, s, d }] =
+ _keyCache[layer][new[] { sourceBatch, h, s, d }];
+ _valueCache[layer][new[] { destBatch, h, s, d }] =
+ _valueCache[layer][new[] { sourceBatch, h, s, d }];
+ }
}
}
}
+
+ _sequenceLengths[layer][destBatch] = seqLen;
+ }
+ }
+
+ private void EnsureInt8Scales(int layerIndex, Tensor newKeys, Tensor newValues, int batchSize, int newSeqLen)
+ {
+ if (!_useInt8Storage)
+ {
+ return;
}
- _sequenceLengths[destBatch] = seqLen;
+ float maxAbsK = 0f;
+ float maxAbsV = 0f;
+
+ for (int b = 0; b < batchSize; b++)
+ {
+ for (int h = 0; h < _config.NumHeads; h++)
+ {
+ for (int s = 0; s < newSeqLen; s++)
+ {
+ for (int d = 0; d < _config.HeadDimension; d++)
+ {
+ float k = _toFloat!(newKeys[new[] { b, h, s, d }]);
+ float v = _toFloat!(newValues[new[] { b, h, s, d }]);
+ float ak = Math.Abs(k);
+ float av = Math.Abs(v);
+ if (ak > maxAbsK) maxAbsK = ak;
+ if (av > maxAbsV) maxAbsV = av;
+ }
+ }
+ }
+ }
+
+ EnsureInt8ScaleForLayer(layerIndex, isKey: true, maxAbs: maxAbsK);
+ EnsureInt8ScaleForLayer(layerIndex, isKey: false, maxAbs: maxAbsV);
+ }
+
+ private void EnsureInt8ScaleForLayer(int layerIndex, bool isKey, float maxAbs)
+ {
+ float requiredScale = maxAbs > 0f ? (maxAbs / 127f) : 1f;
+ if (requiredScale <= 0f) requiredScale = 1f;
+
+ float currentScale = isKey ? _keyScaleInt8![layerIndex] : _valueScaleInt8![layerIndex];
+
+ if (currentScale <= 0f)
+ {
+ if (isKey) _keyScaleInt8![layerIndex] = requiredScale;
+ else _valueScaleInt8![layerIndex] = requiredScale;
+ return;
+ }
+
+ if (requiredScale > currentScale)
+ {
+ RescaleInt8Layer(layerIndex, isKey, currentScale, requiredScale);
+ if (isKey) _keyScaleInt8![layerIndex] = requiredScale;
+ else _valueScaleInt8![layerIndex] = requiredScale;
+ }
+ }
+
+ private void RescaleInt8Layer(int layerIndex, bool isKey, float oldScale, float newScale)
+ {
+ if (oldScale <= 0f || newScale <= 0f || Math.Abs(newScale - oldScale) < float.Epsilon)
+ {
+ return;
+ }
+
+ var cache = isKey ? _keyCacheInt8![layerIndex] : _valueCacheInt8![layerIndex];
+
+ for (int b = 0; b < _sequenceLengths[layerIndex].Length; b++)
+ {
+ int seqLen = _sequenceLengths[layerIndex][b];
+ for (int h = 0; h < _config.NumHeads; h++)
+ {
+ for (int s = 0; s < seqLen; s++)
+ {
+ for (int d = 0; d < _config.HeadDimension; d++)
+ {
+ sbyte q = cache[new[] { b, h, s, d }];
+ float value = q * oldScale;
+ cache[new[] { b, h, s, d }] = QuantizeToInt8(value, newScale);
+ }
+ }
+ }
+ }
+ }
+
+ private static sbyte QuantizeToInt8(float value, float scale)
+ {
+ if (scale <= 0f) scale = 1f;
+ int q = (int)Math.Round(value / scale);
+ if (q > 127) q = 127;
+ if (q < -127) q = -127;
+ return (sbyte)q;
+ }
+
+ private static float DequantizeInt8(sbyte value, float scale)
+ {
+ if (scale <= 0f) scale = 1f;
+ return value * scale;
}
private void ValidateLayerIndex(int layerIndex)
@@ -480,7 +778,7 @@ private void ValidateInputShapes(Tensor keys, Tensor values)
private void EnsureCacheAllocated(int layerIndex)
{
- if (_keyCache[layerIndex] == null)
+ if (!IsLayerAllocated(layerIndex))
{
var shape = new[]
{
@@ -490,8 +788,23 @@ private void EnsureCacheAllocated(int layerIndex)
_config.HeadDimension
};
- _keyCache[layerIndex] = new Tensor(shape);
- _valueCache[layerIndex] = new Tensor(shape);
+ if (_useFp16Storage)
+ {
+ _keyCacheFp16![layerIndex] = new Tensor(shape);
+ _valueCacheFp16![layerIndex] = new Tensor(shape);
+ }
+ else if (_useInt8Storage)
+ {
+ _keyCacheInt8![layerIndex] = new Tensor(shape);
+ _valueCacheInt8![layerIndex] = new Tensor(shape);
+ _keyScaleInt8![layerIndex] = 0f;
+ _valueScaleInt8![layerIndex] = 0f;
+ }
+ else
+ {
+ _keyCache[layerIndex] = new Tensor(shape);
+ _valueCache[layerIndex] = new Tensor(shape);
+ }
}
}
@@ -499,7 +812,7 @@ private void HandleSlidingWindowEviction(int layerIndex, int batchSize, int newS
{
for (int b = 0; b < batchSize; b++)
{
- int currentLen = _sequenceLengths[b];
+ int currentLen = _sequenceLengths[layerIndex][b];
int newLen = currentLen + newSeqLen;
if (newLen > _config.WindowSize)
@@ -517,18 +830,43 @@ private void HandleSlidingWindowEviction(int layerIndex, int batchSize, int newS
int srcPos = evictCount + s;
for (int d = 0; d < _config.HeadDimension; d++)
{
- _keyCache[layerIndex][new[] { b, h, s, d }] =
- _keyCache[layerIndex][new[] { b, h, srcPos, d }];
- _valueCache[layerIndex][new[] { b, h, s, d }] =
- _valueCache[layerIndex][new[] { b, h, srcPos, d }];
+ if (_useInt8Storage)
+ {
+ _keyCacheInt8![layerIndex][new[] { b, h, s, d }] =
+ _keyCacheInt8![layerIndex][new[] { b, h, srcPos, d }];
+ _valueCacheInt8![layerIndex][new[] { b, h, s, d }] =
+ _valueCacheInt8![layerIndex][new[] { b, h, srcPos, d }];
+ }
+ else if (_useFp16Storage)
+ {
+ _keyCacheFp16![layerIndex][new[] { b, h, s, d }] =
+ _keyCacheFp16![layerIndex][new[] { b, h, srcPos, d }];
+ _valueCacheFp16![layerIndex][new[] { b, h, s, d }] =
+ _valueCacheFp16![layerIndex][new[] { b, h, srcPos, d }];
+ }
+ else
+ {
+ _keyCache[layerIndex][new[] { b, h, s, d }] =
+ _keyCache[layerIndex][new[] { b, h, srcPos, d }];
+ _valueCache[layerIndex][new[] { b, h, s, d }] =
+ _valueCache[layerIndex][new[] { b, h, srcPos, d }];
+ }
}
}
}
}
- _sequenceLengths[b] = keepCount;
+ _sequenceLengths[layerIndex][b] = keepCount;
_evictions += evictCount;
}
}
}
+
+ private bool IsLayerAllocated(int layerIndex)
+ {
+ if (_useInt8Storage)
+ return _keyCacheInt8![layerIndex] != null;
+
+ return _useFp16Storage ? _keyCacheFp16![layerIndex] != null : _keyCache[layerIndex] != null;
+ }
}
diff --git a/src/Inference/KVCacheConfig.cs b/src/Inference/KVCacheConfig.cs
index fe549f3d0b..b7a3448c4d 100644
--- a/src/Inference/KVCacheConfig.cs
+++ b/src/Inference/KVCacheConfig.cs
@@ -21,7 +21,7 @@ namespace AiDotNet.Inference;
/// which don't change once computed for a given position.
///
///
-public class KVCacheConfig
+internal class KVCacheConfig
{
///
/// Maximum sequence length the cache can hold.
@@ -113,14 +113,15 @@ public long EstimateMemoryBytes()
long elementsPerLayer = (long)MaxBatchSize * NumHeads * MaxSequenceLength * HeadDimension;
long totalElements = elementsPerLayer * NumLayers * 2; // K and V
- int bytesPerElement = DataType switch
- {
- CacheDataType.Float16 => 2,
- CacheDataType.Float32 => 4,
- CacheDataType.Float64 => 8,
- CacheDataType.BFloat16 => 2,
- _ => 4
- };
+ int bytesPerElement = DataType switch
+ {
+ CacheDataType.Int8 => 1,
+ CacheDataType.Float16 => 2,
+ CacheDataType.Float32 => 4,
+ CacheDataType.Float64 => 8,
+ CacheDataType.BFloat16 => 2,
+ _ => 4
+ };
return totalElements * bytesPerElement;
}
@@ -187,8 +188,11 @@ public static KVCacheConfig ForModel(string modelSize)
///
/// Data types supported for KV-Cache storage.
///
-public enum CacheDataType
+internal enum CacheDataType
{
+ /// Signed 8-bit integer quantization (int8) with scaling.
+ Int8,
+
/// Half precision (16-bit float).
Float16,
@@ -205,7 +209,7 @@ public enum CacheDataType
///
/// Device placement options for KV-Cache.
///
-public enum CacheDevice
+internal enum CacheDevice
{
/// Automatically select based on available hardware.
Auto,
diff --git a/src/Inference/PagedAttention/BlockManager.cs b/src/Inference/PagedAttention/BlockManager.cs
index 582c7ac6ee..89614b0803 100644
--- a/src/Inference/PagedAttention/BlockManager.cs
+++ b/src/Inference/PagedAttention/BlockManager.cs
@@ -24,7 +24,7 @@ namespace AiDotNet.Inference.PagedAttention;
///
///
/// The numeric type for tensor computations.
-public class BlockManager
+internal class BlockManager
{
private readonly BlockManagerConfig _config;
private readonly object _lock = new();
@@ -338,7 +338,7 @@ public void Reset()
///
/// Configuration for the block manager.
///
-public class BlockManagerConfig
+internal class BlockManagerConfig
{
///
/// Number of tokens per block.
@@ -434,7 +434,7 @@ public static BlockManagerConfig ForModel(string modelName, long availableMemory
///
/// Statistics about the block manager state.
///
-public class BlockManagerStats
+internal class BlockManagerStats
{
/// Total number of blocks in the pool.
public int TotalBlocks { get; set; }
diff --git a/src/Inference/PagedAttention/BlockTable.cs b/src/Inference/PagedAttention/BlockTable.cs
index 3533f0adac..8c7ec16d33 100644
--- a/src/Inference/PagedAttention/BlockTable.cs
+++ b/src/Inference/PagedAttention/BlockTable.cs
@@ -21,7 +21,7 @@ namespace AiDotNet.Inference.PagedAttention;
/// - Swapping to disk (move a chapter to storage, update the table)
///
///
-public class BlockTable
+internal class BlockTable
{
private readonly int _blockSize;
private readonly List _physicalBlockIds;
@@ -236,7 +236,7 @@ public override string ToString()
/// Manages block tables for multiple sequences.
///
/// The numeric type.
-public class BlockTableManager
+internal class BlockTableManager
{
private readonly BlockManager _blockManager;
private readonly Dictionary _blockTables;
diff --git a/src/Inference/PagedAttention/PagedAttentionKernel.cs b/src/Inference/PagedAttention/PagedAttentionKernel.cs
index ccca7cf260..356e27e730 100644
--- a/src/Inference/PagedAttention/PagedAttentionKernel.cs
+++ b/src/Inference/PagedAttention/PagedAttentionKernel.cs
@@ -1,4 +1,6 @@
+using System.Buffers;
using System.Runtime.CompilerServices;
+using AiDotNet.Inference.Quantization;
namespace AiDotNet.Inference.PagedAttention;
@@ -23,7 +25,7 @@ namespace AiDotNet.Inference.PagedAttention;
///
///
/// The numeric type for tensor computations.
-public class PagedAttentionKernel
+internal class PagedAttentionKernel
{
private readonly PagedKVCache _kvCache;
private readonly PagedAttentionConfig _config;
@@ -300,15 +302,22 @@ public void UpdateCache(
int position,
int layer)
{
- // Ensure capacity
- if (!_kvCache.HasCapacityFor(sequenceId, 1))
+ // Ensure logical length and capacity for this position.
+ int requiredLength = position + 1;
+ int currentLength = _kvCache.GetSequenceLength(sequenceId);
+ if (requiredLength > currentLength)
{
- _kvCache.ExtendSequence(sequenceId, 1);
+ int additionalTokens = requiredLength - currentLength;
+ if (!_kvCache.ExtendSequence(sequenceId, additionalTokens))
+ {
+ throw new InvalidOperationException(
+ $"Failed to extend PagedKVCache sequence {sequenceId} to length {requiredLength}.");
+ }
}
// Convert and write
- var keyT = ConvertSpan(key);
- var valueT = ConvertSpan(value);
+ var keyT = ConvertArray(key);
+ var valueT = ConvertArray(value);
_kvCache.WriteKey(sequenceId, position, layer, keyT);
_kvCache.WriteValue(sequenceId, position, layer, valueT);
@@ -343,27 +352,94 @@ public void Forward(
int projDim = numHeads * headDim;
float scale = 1.0f / MathF.Sqrt(headDim);
- // Project Q, K, V
- var query = new float[projDim];
- var key = new float[projDim];
- var value = new float[projDim];
+ var pool = ArrayPool.Shared;
+ var queryBuf = pool.Rent(projDim);
+ var keyBuf = pool.Rent(projDim);
+ var valueBuf = pool.Rent(projDim);
+ var attnBuf = pool.Rent(projDim);
- // Q = hidden @ wQ
- MatVecMul(hiddenStates, wQ, query.AsSpan(), hiddenDim, projDim);
- // K = hidden @ wK
- MatVecMul(hiddenStates, wK, key.AsSpan(), hiddenDim, projDim);
- // V = hidden @ wV
- MatVecMul(hiddenStates, wV, value.AsSpan(), hiddenDim, projDim);
+ try
+ {
+ var query = queryBuf.AsSpan(0, projDim);
+ var key = keyBuf.AsSpan(0, projDim);
+ var value = valueBuf.AsSpan(0, projDim);
+ var attnOutput = attnBuf.AsSpan(0, projDim);
+
+ // Q = hidden @ wQ
+ MatVecMul(hiddenStates, wQ, query, hiddenDim, projDim);
+ // K = hidden @ wK
+ MatVecMul(hiddenStates, wK, key, hiddenDim, projDim);
+ // V = hidden @ wV
+ MatVecMul(hiddenStates, wV, value, hiddenDim, projDim);
+
+ // Update cache with new K, V
+ UpdateCache(key, value, sequenceId, position, layer);
+
+ // Compute attention
+ ComputeTiledPagedAttention(query, sequenceId, layer, attnOutput, scale);
+
+ // Project output: out = attn @ wO
+ MatVecMul(attnOutput, wO, output, projDim, hiddenDim);
+ }
+ finally
+ {
+ pool.Return(queryBuf);
+ pool.Return(keyBuf);
+ pool.Return(valueBuf);
+ pool.Return(attnBuf);
+ }
+ }
+
+ public void ForwardQuantized(
+ ReadOnlySpan hiddenStates,
+ in Int8WeightOnlyQuantization.QuantizedWeights wQ,
+ in Int8WeightOnlyQuantization.QuantizedWeights wK,
+ in Int8WeightOnlyQuantization.QuantizedWeights wV,
+ in Int8WeightOnlyQuantization.QuantizedWeights wO,
+ long sequenceId,
+ int position,
+ int layer,
+ Span output)
+ {
+ int hiddenDim = hiddenStates.Length;
+ int numHeads = _config.NumHeads;
+ int headDim = _config.HeadDimension;
+ int projDim = numHeads * headDim;
+ float scale = 1.0f / MathF.Sqrt(headDim);
- // Update cache with new K, V
- UpdateCache(key.AsSpan(), value.AsSpan(), sequenceId, position, layer);
+ if (wQ.Cols != hiddenDim || wK.Cols != hiddenDim || wV.Cols != hiddenDim || wO.Cols != projDim)
+ {
+ throw new ArgumentException("Quantized weight dimensions do not match expected shapes.");
+ }
- // Compute attention
- var attnOutput = new float[projDim];
- ComputeTiledPagedAttention(query.AsSpan(), sequenceId, layer, attnOutput.AsSpan(), scale);
+ var pool = ArrayPool.Shared;
+ var queryBuf = pool.Rent(projDim);
+ var keyBuf = pool.Rent(projDim);
+ var valueBuf = pool.Rent(projDim);
+ var attnBuf = pool.Rent(projDim);
- // Project output: out = attn @ wO
- MatVecMul(attnOutput.AsSpan(), wO, output, projDim, hiddenDim);
+ try
+ {
+ var query = queryBuf.AsSpan(0, projDim);
+ var key = keyBuf.AsSpan(0, projDim);
+ var value = valueBuf.AsSpan(0, projDim);
+ var attnOutput = attnBuf.AsSpan(0, projDim);
+
+ MatVecMulInt8(hiddenStates, wQ, query);
+ MatVecMulInt8(hiddenStates, wK, key);
+ MatVecMulInt8(hiddenStates, wV, value);
+
+ UpdateCache(key, value, sequenceId, position, layer);
+ ComputeTiledPagedAttention(query, sequenceId, layer, attnOutput, scale);
+ MatVecMulInt8(attnOutput, wO, output);
+ }
+ finally
+ {
+ pool.Return(queryBuf);
+ pool.Return(keyBuf);
+ pool.Return(valueBuf);
+ pool.Return(attnBuf);
+ }
}
private static void MatVecMul(ReadOnlySpan vec, ReadOnlySpan mat, Span output, int inDim, int outDim)
@@ -381,6 +457,32 @@ private static void MatVecMul(ReadOnlySpan vec, ReadOnlySpan mat,
}
}
+ private static void MatVecMulInt8(ReadOnlySpan vec, in Int8WeightOnlyQuantization.QuantizedWeights mat, Span output)
+ {
+ int rows = mat.Rows;
+ int cols = mat.Cols;
+
+ if (vec.Length != cols)
+ throw new ArgumentException("Input vector length must match quantized matrix column count.", nameof(vec));
+ if (output.Length < rows)
+ throw new ArgumentException("Output span too small for quantized matvec.", nameof(output));
+
+ var weights = mat.Weights;
+ var scales = mat.Scales;
+
+ for (int r = 0; r < rows; r++)
+ {
+ int baseIdx = r * cols;
+ float sum = 0f;
+ for (int c = 0; c < cols; c++)
+ {
+ sum += weights[baseIdx + c] * vec[c];
+ }
+
+ output[r] = sum * scales[r];
+ }
+ }
+
private static float ToFloat(T value)
{
if (typeof(T) == typeof(float))
@@ -405,15 +507,13 @@ private static T FromFloat(float value)
return (T)Convert.ChangeType(value, typeof(T))!;
}
- private static ReadOnlySpan ConvertSpan(ReadOnlySpan source)
+ private static T[] ConvertArray(ReadOnlySpan source)
{
if (typeof(T) == typeof(float))
{
- // Safe: We've verified T == float at runtime
- // Reinterpret the array using object cast
- var floatArray = source.ToArray();
- var tArray = (T[])(object)floatArray;
- return new ReadOnlySpan(tArray);
+ // Safe: runtime-verified T == float.
+ // Return a rooted array so GC cannot collect it while spans are in use.
+ return (T[])(object)source.ToArray();
}
var result = new T[source.Length];
@@ -428,7 +528,7 @@ private static ReadOnlySpan ConvertSpan(ReadOnlySpan source)
///
/// Configuration for paged attention kernel.
///
-public class PagedAttentionConfig
+internal class PagedAttentionConfig
{
/// Number of attention heads.
public int NumHeads { get; set; } = 32;
@@ -453,7 +553,7 @@ public class PagedAttentionConfig
/// Integrates PagedAttention with ContinuousBatcher for high-throughput serving.
///
/// Numeric type.
-public class PagedAttentionServer : IDisposable
+internal class PagedAttentionServer : IDisposable
{
private readonly PagedKVCache _kvCache;
private readonly PagedAttentionKernel _kernel;
diff --git a/src/Inference/PagedAttention/PagedKVCache.cs b/src/Inference/PagedAttention/PagedKVCache.cs
index 52af9c28ec..bcffc9863d 100644
--- a/src/Inference/PagedAttention/PagedKVCache.cs
+++ b/src/Inference/PagedAttention/PagedKVCache.cs
@@ -23,7 +23,7 @@ namespace AiDotNet.Inference.PagedAttention;
///
///
/// The numeric type for tensor computations.
-public class PagedKVCache : IDisposable
+internal class PagedKVCache : IDisposable
{
private readonly PagedKVCacheConfig _config;
private readonly BlockManager _blockManager;
@@ -87,7 +87,21 @@ public PagedKVCache(PagedKVCacheConfig config)
// Allocate physical storage
long totalElements = _elementsPerBlock * config.NumBlocks;
- _kvStorage = new T[totalElements];
+ if (totalElements > int.MaxValue)
+ throw new ArgumentOutOfRangeException(nameof(config), $"PagedKVCache requires totalElements <= {int.MaxValue}, but got {totalElements}. Reduce NumBlocks or memory size.");
+
+ try
+ {
+ _kvStorage = new T[(int)totalElements];
+ }
+ catch (OutOfMemoryException ex)
+ {
+ throw new InvalidOperationException(
+ $"Failed to allocate PagedKVCache storage ({totalElements} elements). " +
+ "This can happen when requesting very large contiguous memory blocks (e.g., multi-GB) in environments with tighter single-object limits. " +
+ "Reduce available memory/NumBlocks or use a runtime that supports larger allocations.",
+ ex);
+ }
_sequenceMetadata = new Dictionary();
}
@@ -120,7 +134,10 @@ public bool AllocateSequence(long sequenceId, int initialTokens)
if (_sequenceMetadata.ContainsKey(sequenceId))
return false;
- int blocksNeeded = _blockManager.BlocksForTokens(initialTokens);
+ // Allocate at least one block up-front so the first token write (position 0) always has capacity.
+ // Actual "current length" bookkeeping still starts at initialTokens.
+ int blocksNeeded = _blockManager.BlocksForTokens(Math.Max(1, initialTokens));
+ blocksNeeded = Math.Max(1, blocksNeeded);
var table = _blockTableManager.CreateBlockTable(sequenceId, blocksNeeded);
if (table == null)
@@ -444,7 +461,7 @@ private class SequenceMetadata
///
/// Configuration for PagedKVCache.
///
-public class PagedKVCacheConfig
+internal class PagedKVCacheConfig
{
///
/// Number of tokens per block.
@@ -521,7 +538,7 @@ public static PagedKVCacheConfig ForModel(string modelName, long availableBytes,
///
/// Statistics about the paged KV cache.
///
-public class PagedKVCacheStats
+internal class PagedKVCacheStats
{
/// Number of active sequences.
public int ActiveSequences { get; set; }
diff --git a/src/Inference/PagedCachedMultiHeadAttention.cs b/src/Inference/PagedCachedMultiHeadAttention.cs
new file mode 100644
index 0000000000..d93a98292a
--- /dev/null
+++ b/src/Inference/PagedCachedMultiHeadAttention.cs
@@ -0,0 +1,545 @@
+using AiDotNet.Inference.PagedAttention;
+using AiDotNet.NeuralNetworks.Attention;
+using AiDotNet.NeuralNetworks.Layers;
+using AiDotNet.Tensors.LinearAlgebra;
+using System.Buffers;
+using AiDotNet.Inference.Quantization;
+
+namespace AiDotNet.Inference;
+
+///
+/// Multi-head attention layer backed by PagedKVCache for efficient multi-sequence inference.
+///
+///
+/// This layer is intended for inference-time usage. When is enabled
+/// and a is attached, it uses PagedKVCache to avoid reallocations and
+/// allow many independent sequences to grow efficiently.
+///
+/// Limitation: This layer currently supports batchSize == 1 per sequence to avoid cache mixing.
+/// For concurrent serving, create one sequence per request (distinct values).
+///
+///
+internal class PagedCachedMultiHeadAttention : LayerBase
+{
+ private readonly int _headCount;
+ private readonly int _headDimension;
+ private readonly int _embeddingDimension;
+ private readonly bool _useCausalMask;
+
+ private Matrix _queryWeights;
+ private Matrix _keyWeights;
+ private Matrix _valueWeights;
+ private Matrix _outputWeights;
+ private Vector _outputBias;
+
+ private Tensor? _lastInput;
+ private Tensor? _lastOutput;
+ private int _currentPosition;
+
+ private readonly FlashAttentionConfig _flashConfig;
+
+ private readonly object _kernelWeightsLock = new();
+ private float[]? _cachedWQ;
+ private float[]? _cachedWK;
+ private float[]? _cachedWV;
+ private float[]? _cachedWO;
+ private Int8WeightOnlyQuantization.QuantizedWeights? _cachedWQInt8;
+ private Int8WeightOnlyQuantization.QuantizedWeights? _cachedWKInt8;
+ private Int8WeightOnlyQuantization.QuantizedWeights? _cachedWVInt8;
+ private Int8WeightOnlyQuantization.QuantizedWeights? _cachedWOInt8;
+
+ internal bool EnableWeightOnlyQuantization { get; set; }
+
+ ///
+ /// Gets whether this layer supports training.
+ ///
+ public override bool SupportsTraining => false;
+
+ ///
+ /// Gets the number of attention heads.
+ ///
+ public int HeadCount => _headCount;
+
+ ///
+ /// Gets the dimension of each attention head.
+ ///
+ public int HeadDimension => _headDimension;
+
+ ///
+ /// Gets or sets the layer index for KV-cache addressing.
+ ///
+ public int LayerIndex { get; set; }
+
+ ///
+ /// Gets or sets whether the layer is in inference mode (uses paged cache).
+ ///
+ public bool InferenceMode { get; set; }
+
+ ///
+ /// Gets or sets the PagedAttention kernel (owns the paged cache).
+ ///
+ public PagedAttentionKernel? Kernel { get; set; }
+
+ ///
+ /// Gets or sets the sequence ID used for this layer's cache operations.
+ ///
+ public long SequenceId { get; set; }
+
+ public PagedCachedMultiHeadAttention(
+ int sequenceLength,
+ int embeddingDimension,
+ int headCount,
+ bool useCausalMask,
+ IActivationFunction? activationFunction = null)
+ : base(
+ [sequenceLength, embeddingDimension],
+ [sequenceLength, embeddingDimension],
+ activationFunction ?? new IdentityActivation())
+ {
+ if (embeddingDimension % headCount != 0)
+ {
+ throw new ArgumentException(
+ $"Embedding dimension ({embeddingDimension}) must be divisible by head count ({headCount}).",
+ nameof(headCount));
+ }
+
+ _embeddingDimension = embeddingDimension;
+ _headCount = headCount;
+ _headDimension = embeddingDimension / headCount;
+ _useCausalMask = useCausalMask;
+
+ _queryWeights = new Matrix(embeddingDimension, embeddingDimension);
+ _keyWeights = new Matrix(embeddingDimension, embeddingDimension);
+ _valueWeights = new Matrix(embeddingDimension, embeddingDimension);
+ _outputWeights = new Matrix(embeddingDimension, embeddingDimension);
+ _outputBias = new Vector(embeddingDimension);
+
+ _flashConfig = FlashAttentionConfig.Default;
+ _flashConfig.UseCausalMask = useCausalMask;
+ }
+
+ public override Tensor Forward(Tensor input)
+ {
+ _lastInput = input;
+
+ if (!InferenceMode || Kernel == null)
+ {
+ var statelessOutput = ForwardStateless(input);
+ _lastOutput = statelessOutput;
+ return statelessOutput;
+ }
+
+ // Inference mode: update cache and compute attention token-by-token.
+ // This supports both prefill (seqLen>1) and decode (seqLen==1) by iterating tokens.
+ if (input.Shape.Length < 3)
+ {
+ throw new ArgumentException("Expected input shape [batch, seqLen, embeddingDim].", nameof(input));
+ }
+
+ int batchSize = input.Shape[0];
+ int seqLen = input.Shape[1];
+ int embDim = input.Shape[2];
+
+ if (embDim != _embeddingDimension)
+ {
+ throw new ArgumentException($"Expected embeddingDim={_embeddingDimension}, got {embDim}.", nameof(input));
+ }
+
+ if (batchSize != 1)
+ {
+ // PagedAttentionKernel supports batched attention, but this layer's state model is per-sequence.
+ // Keep it strict for now to avoid cache mixing.
+ throw new NotSupportedException("PagedCachedMultiHeadAttention currently supports batchSize==1 per sequence.");
+ }
+
+ var output = new Tensor([batchSize, seqLen, embDim]);
+
+ // Materialize weights to float spans for the paged kernel.
+ // Note: This is intentionally conservative and prioritizes correctness.
+ // PagedAttentionKernel's MatVecMul expects matrices stored as [outDim, inDim] row-major.
+ // Our weights are stored as [inDim, outDim], so we pass a transposed layout.
+ EnsureKernelWeightCache();
+ var wQ = _cachedWQ!;
+ var wK = _cachedWK!;
+ var wV = _cachedWV!;
+ var wO = _cachedWO!;
+
+ // Process each token sequentially to ensure causal behavior during prefill.
+ var pool = ArrayPool.Shared;
+ var hiddenBuffer = pool.Rent(embDim);
+ var tokenOutBuffer = pool.Rent(embDim);
+
+ try
+ {
+ var hidden = hiddenBuffer.AsSpan(0, embDim);
+ var tokenOut = tokenOutBuffer.AsSpan(0, embDim);
+
+ var wQInt8 = _cachedWQInt8;
+ var wKInt8 = _cachedWKInt8;
+ var wVInt8 = _cachedWVInt8;
+ var wOInt8 = _cachedWOInt8;
+
+ bool useQuantized = EnableWeightOnlyQuantization &&
+ typeof(T) == typeof(float) &&
+ wQInt8.HasValue &&
+ wKInt8.HasValue &&
+ wVInt8.HasValue &&
+ wOInt8.HasValue;
+
+ for (int t = 0; t < seqLen; t++)
+ {
+ for (int d = 0; d < embDim; d++)
+ {
+ hidden[d] = Convert.ToSingle(input[0, t, d]);
+ }
+
+ if (useQuantized)
+ {
+ Kernel.ForwardQuantized(
+ hiddenStates: hidden,
+ wQ: wQInt8!.Value,
+ wK: wKInt8!.Value,
+ wV: wVInt8!.Value,
+ wO: wOInt8!.Value,
+ sequenceId: SequenceId,
+ position: _currentPosition,
+ layer: LayerIndex,
+ output: tokenOut);
+ }
+ else
+ {
+ Kernel.Forward(
+ hiddenStates: hidden,
+ wQ: wQ,
+ wK: wK,
+ wV: wV,
+ wO: wO,
+ sequenceId: SequenceId,
+ position: _currentPosition,
+ layer: LayerIndex,
+ output: tokenOut);
+ }
+
+ // Add bias and activation.
+ for (int d = 0; d < embDim; d++)
+ {
+ T value = NumOps.FromDouble(tokenOut[d]);
+ value = NumOps.Add(value, _outputBias[d]);
+ output[0, t, d] = ScalarActivation!.Activate(value);
+ }
+
+ _currentPosition++;
+ }
+ }
+ finally
+ {
+ pool.Return(hiddenBuffer);
+ pool.Return(tokenOutBuffer);
+ }
+
+ _lastOutput = output;
+ return output;
+ }
+
+ private Tensor ForwardStateless(Tensor input)
+ {
+ // Stateless fallback using FlashAttention.
+ // Compute Q,K,V projections.
+ var (q, k, v) = ComputeQkv(input);
+
+ // FlashAttention expects [B, H, S, D]
+ var qh = SplitHeads(q);
+ var kh = SplitHeads(k);
+ var vh = SplitHeads(v);
+
+ var (attn, _) = FlashAttention.Forward(qh, kh, vh, _flashConfig);
+
+ // Merge heads back to [B, S, E]
+ var merged = MergeHeads(attn);
+
+ // Output projection + bias + activation.
+ // Use the tensor/matrix multiply path to leverage optimized kernels.
+ var projected = merged.Multiply(_outputWeights);
+
+ int batch = projected.Shape[0];
+ int seqLen = projected.Shape[1];
+ var output = new Tensor([batch, seqLen, _embeddingDimension]);
+
+ for (int b = 0; b < batch; b++)
+ {
+ for (int s = 0; s < seqLen; s++)
+ {
+ for (int o = 0; o < _embeddingDimension; o++)
+ {
+ T value = NumOps.Add(projected[b, s, o], _outputBias[o]);
+ output[b, s, o] = ScalarActivation!.Activate(value);
+ }
+ }
+ }
+
+ return output;
+ }
+
+ private (Tensor Q, Tensor K, Tensor V) ComputeQkv(Tensor input)
+ {
+ // Use the tensor/matrix multiply path to leverage optimized kernels.
+ var q = input.Multiply(_queryWeights);
+ var k = input.Multiply(_keyWeights);
+ var v = input.Multiply(_valueWeights);
+ return (q, k, v);
+ }
+
+ private Tensor SplitHeads(Tensor x)
+ {
+ int batchSize = x.Shape[0];
+ int seqLen = x.Shape[1];
+ var reshaped = new Tensor([batchSize, _headCount, seqLen, _headDimension]);
+
+ for (int b = 0; b < batchSize; b++)
+ {
+ for (int s = 0; s < seqLen; s++)
+ {
+ for (int h = 0; h < _headCount; h++)
+ {
+ int baseOffset = h * _headDimension;
+ for (int d = 0; d < _headDimension; d++)
+ {
+ reshaped[b, h, s, d] = x[b, s, baseOffset + d];
+ }
+ }
+ }
+ }
+
+ return reshaped;
+ }
+
+ private Tensor MergeHeads(Tensor x)
+ {
+ int batchSize = x.Shape[0];
+ int seqLen = x.Shape[2];
+ var merged = new Tensor