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([batchSize, seqLen, _embeddingDimension]); + + 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++) + { + merged[b, s, baseOffset + d] = x[b, h, s, d]; + } + } + } + } + + return merged; + } + + private static float[] MatrixToFloatForKernel(Matrix matrix) + { + int inDim = matrix.Rows; + int outDim = matrix.Columns; + var data = new float[outDim * inDim]; + + for (int o = 0; o < outDim; o++) + { + int rowOffset = o * inDim; + for (int i = 0; i < inDim; i++) + { + data[rowOffset + i] = Convert.ToSingle(matrix[i, o]); + } + } + + return data; + } + + public override Vector GetParameters() + { + int totalParams = _queryWeights.Rows * _queryWeights.Columns * 4 + _outputBias.Length; + var parameters = new Vector(totalParams); + int index = 0; + + foreach (var matrix in new[] { _queryWeights, _keyWeights, _valueWeights, _outputWeights }) + { + for (int i = 0; i < matrix.Rows; i++) + { + for (int j = 0; j < matrix.Columns; j++) + { + parameters[index++] = matrix[i, j]; + } + } + } + + for (int i = 0; i < _outputBias.Length; i++) + { + parameters[index++] = _outputBias[i]; + } + + return parameters; + } + + public override void SetParameters(Vector parameters) + { + int expectedParams = _queryWeights.Rows * _queryWeights.Columns * 4 + _outputBias.Length; + if (parameters.Length != expectedParams) + { + throw new ArgumentException($"Expected {expectedParams} parameters, got {parameters.Length}"); + } + + int index = 0; + + foreach (var matrix in new[] { _queryWeights, _keyWeights, _valueWeights, _outputWeights }) + { + for (int i = 0; i < matrix.Rows; i++) + { + for (int j = 0; j < matrix.Columns; j++) + { + matrix[i, j] = parameters[index++]; + } + } + } + + for (int i = 0; i < _outputBias.Length; i++) + { + _outputBias[i] = parameters[index++]; + } + + InvalidateKernelWeightCache(); + } + + private void EnsureKernelWeightCache() + { + bool enableQuantization = EnableWeightOnlyQuantization && typeof(T) == typeof(float); + + bool hasDenseWeights = _cachedWQ != null && _cachedWK != null && _cachedWV != null && _cachedWO != null; + bool hasQuantizedWeights = _cachedWQInt8.HasValue && _cachedWKInt8.HasValue && _cachedWVInt8.HasValue && _cachedWOInt8.HasValue; + + if (hasDenseWeights && (!enableQuantization || hasQuantizedWeights)) + { + return; + } + + float[]? localWQ = null; + float[]? localWK = null; + float[]? localWV = null; + float[]? localWO = null; + + Int8WeightOnlyQuantization.QuantizedWeights? localWQInt8 = null; + Int8WeightOnlyQuantization.QuantizedWeights? localWKInt8 = null; + Int8WeightOnlyQuantization.QuantizedWeights? localWVInt8 = null; + Int8WeightOnlyQuantization.QuantizedWeights? localWOInt8 = null; + + // First, determine what's missing and take a quick snapshot inside the lock. + lock (_kernelWeightsLock) + { + hasDenseWeights = _cachedWQ != null && _cachedWK != null && _cachedWV != null && _cachedWO != null; + hasQuantizedWeights = _cachedWQInt8.HasValue && _cachedWKInt8.HasValue && _cachedWVInt8.HasValue && _cachedWOInt8.HasValue; + + if (hasDenseWeights && (!enableQuantization || hasQuantizedWeights)) + { + return; + } + + if (_cachedWQ == null) localWQ = MatrixToFloatForKernel(_queryWeights); + if (_cachedWK == null) localWK = MatrixToFloatForKernel(_keyWeights); + if (_cachedWV == null) localWV = MatrixToFloatForKernel(_valueWeights); + if (_cachedWO == null) localWO = MatrixToFloatForKernel(_outputWeights); + + if (!enableQuantization) + { + _cachedWQInt8 = null; + _cachedWKInt8 = null; + _cachedWVInt8 = null; + _cachedWOInt8 = null; + } + } + + // Compute expensive quantization outside the lock to minimize contention. + if (enableQuantization) + { + int projDim = _headCount * _headDimension; + int hiddenDim = _embeddingDimension; + + var wq = _cachedWQ ?? localWQ; + var wk = _cachedWK ?? localWK; + var wv = _cachedWV ?? localWV; + var wo = _cachedWO ?? localWO; + + if (wq != null && wk != null && wv != null && wo != null) + { + localWQInt8 = Int8WeightOnlyQuantization.QuantizePerRow(wq, projDim, hiddenDim); + localWKInt8 = Int8WeightOnlyQuantization.QuantizePerRow(wk, projDim, hiddenDim); + localWVInt8 = Int8WeightOnlyQuantization.QuantizePerRow(wv, projDim, hiddenDim); + localWOInt8 = Int8WeightOnlyQuantization.QuantizePerRow(wo, hiddenDim, projDim); + } + } + + // Publish results under lock (double-checked to avoid overwriting). + lock (_kernelWeightsLock) + { + _cachedWQ ??= localWQ; + _cachedWK ??= localWK; + _cachedWV ??= localWV; + _cachedWO ??= localWO; + + if (enableQuantization) + { + _cachedWQInt8 ??= localWQInt8; + _cachedWKInt8 ??= localWKInt8; + _cachedWVInt8 ??= localWVInt8; + _cachedWOInt8 ??= localWOInt8; + } + } + } + + private void InvalidateKernelWeightCache() + { + lock (_kernelWeightsLock) + { + _cachedWQ = null; + _cachedWK = null; + _cachedWV = null; + _cachedWO = null; + _cachedWQInt8 = null; + _cachedWKInt8 = null; + _cachedWVInt8 = null; + _cachedWOInt8 = null; + } + } + + public override void ResetState() + { + _lastInput = null; + _lastOutput = null; + _currentPosition = 0; + } + + public override Tensor Backward(Tensor outputGradient) + { + throw new NotSupportedException($"{nameof(PagedCachedMultiHeadAttention)} is intended for inference-time usage only."); + } + + public override void UpdateParameters(T learningRate) + { + throw new NotSupportedException($"{nameof(PagedCachedMultiHeadAttention)} is intended for inference-time usage only."); + } + + public override bool SupportsJitCompilation => false; + + public override Autodiff.ComputationNode ExportComputationGraph(List> inputNodes) + { + throw new NotSupportedException($"{nameof(PagedCachedMultiHeadAttention)} does not support JIT compilation."); + } + + internal override Dictionary GetMetadata() + { + return new Dictionary + { + ["HeadCount"] = _headCount.ToString(), + ["UseCausalMask"] = _useCausalMask.ToString(), + ["EnableWeightOnlyQuantization"] = EnableWeightOnlyQuantization.ToString() + }; + } +} diff --git a/src/Inference/Quantization/Int8WeightOnlyQuantization.cs b/src/Inference/Quantization/Int8WeightOnlyQuantization.cs new file mode 100644 index 0000000000..ed82923d19 --- /dev/null +++ b/src/Inference/Quantization/Int8WeightOnlyQuantization.cs @@ -0,0 +1,100 @@ +using AiDotNet.Tensors.LinearAlgebra; + +namespace AiDotNet.Inference.Quantization; + +internal static class Int8WeightOnlyQuantization +{ + internal readonly struct QuantizedWeights + { + public QuantizedWeights(sbyte[] weights, float[] scales, int rows, int cols) + { + Weights = weights; + Scales = scales; + Rows = rows; + Cols = cols; + } + + public sbyte[] Weights { get; } + public float[] Scales { get; } + public int Rows { get; } + public int Cols { get; } + } + + public static QuantizedWeights QuantizePerRow(Tensor weights) + { + if (weights.Rank != 2) + throw new ArgumentException("Expected 2D weight tensor.", nameof(weights)); + + int rows = weights.Shape[0]; + int cols = weights.Shape[1]; + + var q = new sbyte[rows * cols]; + var scales = new float[rows]; + + for (int r = 0; r < rows; r++) + { + float maxAbs = 0f; + int baseIdx = r * cols; + for (int c = 0; c < cols; c++) + { + float v = weights[r, c]; + float av = MathF.Abs(v); + if (av > maxAbs) + maxAbs = av; + } + + float scale = maxAbs > 0f ? (maxAbs / 127f) : 1f; + scales[r] = scale; + + float inv = 1f / scale; + for (int c = 0; c < cols; c++) + { + float v = weights[r, c] * inv; + int qi = (int)MathF.Round(v); + if (qi > 127) qi = 127; + if (qi < -127) qi = -127; + q[baseIdx + c] = (sbyte)qi; + } + } + + return new QuantizedWeights(q, scales, rows, cols); + } + + public static QuantizedWeights QuantizePerRow(ReadOnlySpan weights, int rows, int cols) + { + if (rows <= 0) throw new ArgumentOutOfRangeException(nameof(rows)); + if (cols <= 0) throw new ArgumentOutOfRangeException(nameof(cols)); + if (weights.Length < rows * cols) throw new ArgumentException("Weight span too small for given dimensions.", nameof(weights)); + + var q = new sbyte[rows * cols]; + var scales = new float[rows]; + + for (int r = 0; r < rows; r++) + { + float maxAbs = 0f; + int baseIdx = r * cols; + for (int c = 0; c < cols; c++) + { + float v = weights[baseIdx + c]; + float av = MathF.Abs(v); + if (av > maxAbs) + maxAbs = av; + } + + float scale = maxAbs > 0f ? (maxAbs / 127f) : 1f; + scales[r] = scale; + + float inv = 1f / scale; + for (int c = 0; c < cols; c++) + { + float v = weights[baseIdx + c] * inv; + int qi = (int)MathF.Round(v); + if (qi > 127) qi = 127; + if (qi < -127) qi = -127; + q[baseIdx + c] = (sbyte)qi; + } + } + + return new QuantizedWeights(q, scales, rows, cols); + } +} diff --git a/src/Inference/Quantization/QuantizedDenseLayer.cs b/src/Inference/Quantization/QuantizedDenseLayer.cs new file mode 100644 index 0000000000..14c8a21909 --- /dev/null +++ b/src/Inference/Quantization/QuantizedDenseLayer.cs @@ -0,0 +1,154 @@ +using AiDotNet.Autodiff; +using AiDotNet.NeuralNetworks.Layers; +using AiDotNet.Tensors.LinearAlgebra; + +namespace AiDotNet.Inference.Quantization; + +/// +/// Inference-only dense layer that uses weight-only INT8 quantization (per-output scaling). +/// +internal sealed class QuantizedDenseLayer : LayerBase +{ + private readonly int _inputSize; + private readonly int _outputSize; + private readonly sbyte[] _weightsInt8; // row-major [out, in] + private readonly float[] _rowScales; // per out + private readonly float[] _biases; + + public QuantizedDenseLayer(DenseLayer source) + : base( + inputShape: source.GetInputShape(), + outputShape: source.GetOutputShape(), + scalarActivation: source.ScalarActivation ?? new AiDotNet.ActivationFunctions.IdentityActivation()) + { + _inputSize = source.GetInputShape()[0]; + _outputSize = source.GetOutputShape()[0]; + + if (source.VectorActivation != null) + throw new InvalidOperationException("QuantizedDenseLayer scalar-activation ctor called for a vector-activation layer."); + + var weights = source.GetWeights(); + var biases = source.GetBiases(); + if (weights == null || biases == null) + throw new ArgumentException("Dense layer must expose weights and biases.", nameof(source)); + + var q = Int8WeightOnlyQuantization.QuantizePerRow(weights); + _weightsInt8 = q.Weights; + _rowScales = q.Scales; + + _biases = new float[biases.Length]; + for (int i = 0; i < _biases.Length; i++) + { + _biases[i] = biases[i]; + } + } + + public QuantizedDenseLayer(DenseLayer source, IVectorActivationFunction vectorActivation) + : base( + inputShape: source.GetInputShape(), + outputShape: source.GetOutputShape(), + vectorActivation: vectorActivation) + { + _inputSize = source.GetInputShape()[0]; + _outputSize = source.GetOutputShape()[0]; + + var weights = source.GetWeights(); + var biases = source.GetBiases(); + if (weights == null || biases == null) + throw new ArgumentException("Dense layer must expose weights and biases.", nameof(source)); + + var q = Int8WeightOnlyQuantization.QuantizePerRow(weights); + _weightsInt8 = q.Weights; + _rowScales = q.Scales; + + _biases = new float[biases.Length]; + for (int i = 0; i < _biases.Length; i++) + { + _biases[i] = biases[i]; + } + } + + public override bool SupportsTraining => false; + + public override bool SupportsJitCompilation => false; + + public override int ParameterCount => 0; + + public override Tensor? GetWeights() => null; + + public override Tensor? GetBiases() => null; + + public override Tensor Forward(Tensor input) + { + bool inputWas1D = false; + Tensor flat; + if (input.Rank == 1) + { + inputWas1D = true; + flat = input.Reshape(1, input.Shape[0]); + } + else if (input.Rank == 2) + { + flat = input; + } + else + { + int batch = input.Shape[0]; + int features = input.Length / batch; + flat = input.Reshape(batch, features); + } + + int batchSize = flat.Shape[0]; + int featuresIn = flat.Shape[1]; + if (featuresIn != _inputSize) + throw new ArgumentException($"QuantizedDenseLayer input size mismatch. Expected {_inputSize}, got {featuresIn}."); + + var output = new Tensor(new[] { batchSize, _outputSize }); + + for (int b = 0; b < batchSize; b++) + { + for (int o = 0; o < _outputSize; o++) + { + float sum = _biases[o]; + float scale = _rowScales[o]; + int wBase = o * _inputSize; + for (int i = 0; i < _inputSize; i++) + { + sum += flat[b, i] * (_weightsInt8[wBase + i] * scale); + } + output[b, o] = sum; + } + } + + var activated = ApplyActivation(output); + if (inputWas1D) + { + return activated.Reshape(_outputSize); + } + + return activated; + } + + public override Tensor Backward(Tensor outputGradient) + => throw new NotSupportedException("QuantizedDenseLayer is inference-only."); + + public override void UpdateParameters(float learningRate) + => throw new NotSupportedException("QuantizedDenseLayer is inference-only."); + + public override void UpdateParameters(Vector parameters) + => throw new NotSupportedException("QuantizedDenseLayer is inference-only."); + + public override Vector GetParameters() + => Vector.Empty(); + + public override void ResetState() + { + // Inference-only; no recurrent state to clear. + } + + public override ComputationNode ExportComputationGraph(List> inputNodes) + { + // WOQ is a runtime inference rewrite; we intentionally don't support JIT graph export here. + throw new NotSupportedException("QuantizedDenseLayer does not support JIT compilation."); + } +} diff --git a/src/Inference/SpeculativeDecoding/DraftResult.cs b/src/Inference/SpeculativeDecoding/DraftResult.cs index e1e0ad3753..5968bf0537 100644 --- a/src/Inference/SpeculativeDecoding/DraftResult.cs +++ b/src/Inference/SpeculativeDecoding/DraftResult.cs @@ -6,7 +6,7 @@ namespace AiDotNet.Inference.SpeculativeDecoding; /// Result of draft token generation. /// /// The numeric type. -public class DraftResult +internal class DraftResult { /// /// Gets the generated draft tokens. diff --git a/src/Inference/SpeculativeDecoding/IDraftModel.cs b/src/Inference/SpeculativeDecoding/IDraftModel.cs index ffb6f886ee..9464806c7b 100644 --- a/src/Inference/SpeculativeDecoding/IDraftModel.cs +++ b/src/Inference/SpeculativeDecoding/IDraftModel.cs @@ -13,7 +13,7 @@ namespace AiDotNet.Inference.SpeculativeDecoding; /// /// /// The numeric type for computations. -public interface IDraftModel +internal interface IDraftModel { /// /// Gets the maximum number of tokens this draft model can generate in one call. diff --git a/src/Inference/SpeculativeDecoding/NGramDraftModel.cs b/src/Inference/SpeculativeDecoding/NGramDraftModel.cs index 6438c9f427..ebad5fdeff 100644 --- a/src/Inference/SpeculativeDecoding/NGramDraftModel.cs +++ b/src/Inference/SpeculativeDecoding/NGramDraftModel.cs @@ -17,7 +17,7 @@ namespace AiDotNet.Inference.SpeculativeDecoding; /// /// /// The numeric type. -public class NGramDraftModel : IDraftModel +internal class NGramDraftModel : IDraftModel { private static readonly INumericOperations NumOps = MathHelper.GetNumericOperations(); diff --git a/src/Inference/SpeculativeDecoding/NeuralDraftModel.cs b/src/Inference/SpeculativeDecoding/NeuralDraftModel.cs index 7017598431..65128d7151 100644 --- a/src/Inference/SpeculativeDecoding/NeuralDraftModel.cs +++ b/src/Inference/SpeculativeDecoding/NeuralDraftModel.cs @@ -13,7 +13,7 @@ namespace AiDotNet.Inference.SpeculativeDecoding; /// /// /// The numeric type. -public class NeuralDraftModel : IDraftModel +internal class NeuralDraftModel : IDraftModel { private static readonly INumericOperations NumOps = MathHelper.GetNumericOperations(); diff --git a/src/Inference/SpeculativeDecoding/SpeculativeDecoder.cs b/src/Inference/SpeculativeDecoding/SpeculativeDecoder.cs index 9b2f063ce7..f609e84106 100644 --- a/src/Inference/SpeculativeDecoding/SpeculativeDecoder.cs +++ b/src/Inference/SpeculativeDecoding/SpeculativeDecoder.cs @@ -30,7 +30,7 @@ namespace AiDotNet.Inference.SpeculativeDecoding; /// /// /// The numeric type for computations. -public class SpeculativeDecoder +internal class SpeculativeDecoder { private static readonly INumericOperations NumOps = MathHelper.GetNumericOperations(); @@ -38,6 +38,10 @@ public class SpeculativeDecoder private readonly Func, Matrix> _targetForward; private readonly SpeculativeDecodingConfig _config; private readonly Random _random; + private readonly int _maxDraftTokens; + private readonly int _maxTreeDepth; + private int _currentDraftTokens; + private int _currentMaxTreeDepth; // Statistics private long _totalTokensGenerated; @@ -45,6 +49,10 @@ public class SpeculativeDecoder private long _acceptedDraftTokens; private long _totalVerificationCalls; + // Tree speculation tracking (when enabled) + private long _treeTotalNodes; + private long _treeAcceptedNodes; + /// /// Gets the configuration. /// @@ -53,9 +61,9 @@ public class SpeculativeDecoder /// /// Gets the draft acceptance rate. /// - public double AcceptanceRate => _totalDraftTokens > 0 - ? (double)_acceptedDraftTokens / _totalDraftTokens - : 0; + public double AcceptanceRate => _config.UseTreeSpeculation + ? (_treeTotalNodes > 0 ? (double)_treeAcceptedNodes / _treeTotalNodes : 0) + : (_totalDraftTokens > 0 ? (double)_acceptedDraftTokens / _totalDraftTokens : 0); /// /// Gets the average tokens generated per verification call. @@ -64,6 +72,29 @@ public class SpeculativeDecoder ? (double)_totalTokensGenerated / _totalVerificationCalls : 0; + /// + /// Gets the total amount of draft work proposed so far. + /// + /// + /// For classic speculation, this counts draft tokens. For tree speculation, this counts explored nodes. + /// + internal long TotalDraftTokens => _config.UseTreeSpeculation ? _treeTotalNodes : _totalDraftTokens; + + /// + /// Gets the current adaptive draft length. + /// + internal int CurrentDraftTokens => _currentDraftTokens; + + /// + /// Gets the current adaptive tree depth. + /// + internal int CurrentMaxTreeDepth => _currentMaxTreeDepth; + + /// + /// Gets the total number of verification calls performed so far. + /// + internal long TotalVerificationCalls => _totalVerificationCalls; + /// /// Creates a speculative decoder. /// @@ -79,6 +110,10 @@ public SpeculativeDecoder( _draftModel = draftModel ?? throw new ArgumentNullException(nameof(draftModel)); _targetForward = targetForward ?? throw new ArgumentNullException(nameof(targetForward)); _config = config ?? new SpeculativeDecodingConfig(); + _maxDraftTokens = Math.Max(1, _config.NumDraftTokens); + _maxTreeDepth = Math.Max(1, _config.MaxTreeDepth); + _currentDraftTokens = _maxDraftTokens; + _currentMaxTreeDepth = _maxTreeDepth; _random = _config.Seed.HasValue ? new Random(_config.Seed.Value) : new Random(); } @@ -98,6 +133,11 @@ public async Task GenerateAsync( int? eosToken = null, CancellationToken cancellationToken = default) { + if (_config.UseTreeSpeculation) + { + return await GenerateTreeAsync(inputTokens, maxNewTokens, temperature, eosToken, cancellationToken).ConfigureAwait(false); + } + var tokens = new List(inputTokens.Length + maxNewTokens); for (int i = 0; i < inputTokens.Length; i++) { @@ -112,7 +152,7 @@ public async Task GenerateAsync( cancellationToken.ThrowIfCancellationRequested(); // Determine how many draft tokens to generate - int numDraft = Math.Min(_config.NumDraftTokens, maxNewTokens - generated); + int numDraft = Math.Min(_currentDraftTokens, maxNewTokens - generated); // Generate draft tokens var currentTokens = new Vector(tokens.ToArray()); @@ -238,6 +278,11 @@ public async Task GenerateAsync( BonusToken = true }); } + + if (_config.AdaptiveDraftLength) + { + AdjustDraftLength(); + } } done: @@ -261,6 +306,94 @@ public async Task GenerateAsync( }; } + private async Task GenerateTreeAsync( + Vector inputTokens, + int maxNewTokens, + T temperature, + int? eosToken, + CancellationToken cancellationToken) + { + List> BatchTargetForward(List> sequences) + { + var results = new List>(sequences.Count); + for (int i = 0; i < sequences.Count; i++) + { + results.Add(_targetForward(sequences[i])); + } + return results; + } + + var treeConfig = new TreeSpeculativeConfig + { + BranchFactor = Math.Max(1, _config.TreeBranchFactor), + MaxDepth = _currentMaxTreeDepth, + Seed = _config.Seed + }; + + var decoder = new TreeSpeculativeDecoder(_draftModel, BatchTargetForward, treeConfig); + var treeResult = await decoder.GenerateAsync(inputTokens, maxNewTokens, temperature, eosToken, cancellationToken).ConfigureAwait(false); + + long nodes = 0; + long accepted = 0; + var stepStats = new List(treeResult.StepStatistics.Count); + for (int i = 0; i < treeResult.StepStatistics.Count; i++) + { + var s = treeResult.StepStatistics[i]; + nodes += s.TreeNodes; + accepted += s.BestPathLength; + stepStats.Add(new StepStatistics + { + DraftTokens = s.TreeNodes, + AcceptedTokens = s.BestPathLength, + ResampledToken = false, + BonusToken = false + }); + } + + _treeTotalNodes += nodes; + _treeAcceptedNodes += accepted; + + if (_config.AdaptiveDraftLength) + { + AdjustDraftLength(); + } + + return new SpeculativeResult + { + Tokens = treeResult.Tokens, + NewTokens = treeResult.NewTokens, + NumGenerated = treeResult.NumGenerated, + AcceptanceRate = AcceptanceRate, + TokensPerVerification = treeResult.StepStatistics.Count > 0 ? (double)treeResult.NumGenerated / treeResult.StepStatistics.Count : 0, + StepStatistics = stepStats + }; + } + + private void AdjustDraftLength() + { + double minAccept = NumOps.ToDouble(_config.MinAcceptanceRate); + double ar = AcceptanceRate; + long work = _config.UseTreeSpeculation ? _treeTotalNodes : _totalDraftTokens; + + if (work < 8) + { + return; + } + + if (ar < minAccept) + { + _currentDraftTokens = Math.Max(1, _currentDraftTokens - 1); + _currentMaxTreeDepth = Math.Max(1, _currentMaxTreeDepth - 1); + return; + } + + if (ar >= minAccept + 0.2) + { + _currentDraftTokens = Math.Min(_maxDraftTokens, _currentDraftTokens + 1); + _currentMaxTreeDepth = Math.Min(_maxTreeDepth, _currentMaxTreeDepth + 1); + } + } + /// /// Synchronous generation method. /// @@ -282,6 +415,10 @@ public void ResetStatistics() _totalDraftTokens = 0; _acceptedDraftTokens = 0; _totalVerificationCalls = 0; + _treeTotalNodes = 0; + _treeAcceptedNodes = 0; + _currentDraftTokens = _maxDraftTokens; + _currentMaxTreeDepth = _maxTreeDepth; _draftModel.Reset(); } diff --git a/src/Inference/SpeculativeDecoding/SpeculativeDecodingConfig.cs b/src/Inference/SpeculativeDecoding/SpeculativeDecodingConfig.cs index b95b7bbbd6..c88293d9d7 100644 --- a/src/Inference/SpeculativeDecoding/SpeculativeDecodingConfig.cs +++ b/src/Inference/SpeculativeDecoding/SpeculativeDecodingConfig.cs @@ -6,7 +6,7 @@ namespace AiDotNet.Inference.SpeculativeDecoding; /// Configuration for speculative decoding. /// /// The numeric type for threshold values. -public class SpeculativeDecodingConfig +internal class SpeculativeDecodingConfig { private static readonly INumericOperations NumOps = MathHelper.GetNumericOperations(); diff --git a/src/Inference/SpeculativeDecoding/SpeculativeDecodingStats.cs b/src/Inference/SpeculativeDecoding/SpeculativeDecodingStats.cs index cb5c69cee4..3551739176 100644 --- a/src/Inference/SpeculativeDecoding/SpeculativeDecodingStats.cs +++ b/src/Inference/SpeculativeDecoding/SpeculativeDecodingStats.cs @@ -3,7 +3,7 @@ namespace AiDotNet.Inference.SpeculativeDecoding; /// /// Overall statistics for speculative decoding. /// -public class SpeculativeDecodingStats +internal class SpeculativeDecodingStats { /// Total tokens generated. public long TotalTokensGenerated { get; set; } diff --git a/src/Inference/SpeculativeDecoding/SpeculativeResult.cs b/src/Inference/SpeculativeDecoding/SpeculativeResult.cs index 967b1d431a..f28389fc09 100644 --- a/src/Inference/SpeculativeDecoding/SpeculativeResult.cs +++ b/src/Inference/SpeculativeDecoding/SpeculativeResult.cs @@ -5,7 +5,7 @@ namespace AiDotNet.Inference.SpeculativeDecoding; /// /// Result of speculative decoding generation. /// -public class SpeculativeResult +internal class SpeculativeResult { /// /// All tokens (input + generated). diff --git a/src/Inference/SpeculativeDecoding/StepStatistics.cs b/src/Inference/SpeculativeDecoding/StepStatistics.cs index 3fafdfef7b..20df557120 100644 --- a/src/Inference/SpeculativeDecoding/StepStatistics.cs +++ b/src/Inference/SpeculativeDecoding/StepStatistics.cs @@ -3,7 +3,7 @@ namespace AiDotNet.Inference.SpeculativeDecoding; /// /// Statistics for a single decoding step. /// -public class StepStatistics +internal class StepStatistics { /// Number of draft tokens generated. public int DraftTokens { get; set; } diff --git a/src/Inference/SpeculativeDecoding/TreeSpeculativeConfig.cs b/src/Inference/SpeculativeDecoding/TreeSpeculativeConfig.cs index 2311b60e5b..b5e4417d01 100644 --- a/src/Inference/SpeculativeDecoding/TreeSpeculativeConfig.cs +++ b/src/Inference/SpeculativeDecoding/TreeSpeculativeConfig.cs @@ -3,7 +3,7 @@ namespace AiDotNet.Inference.SpeculativeDecoding; /// /// Configuration for tree-based speculative decoding. /// -public class TreeSpeculativeConfig +internal class TreeSpeculativeConfig { /// Number of branches per node. public int BranchFactor { get; set; } = 2; diff --git a/src/Inference/SpeculativeDecoding/TreeSpeculativeDecoder.cs b/src/Inference/SpeculativeDecoding/TreeSpeculativeDecoder.cs index 2d5b317733..e2161fb803 100644 --- a/src/Inference/SpeculativeDecoding/TreeSpeculativeDecoder.cs +++ b/src/Inference/SpeculativeDecoding/TreeSpeculativeDecoder.cs @@ -27,7 +27,7 @@ namespace AiDotNet.Inference.SpeculativeDecoding; /// /// /// The numeric type. -public class TreeSpeculativeDecoder +internal class TreeSpeculativeDecoder { private static readonly INumericOperations NumOps = MathHelper.GetNumericOperations(); diff --git a/src/Inference/SpeculativeDecoding/TreeSpeculativeResult.cs b/src/Inference/SpeculativeDecoding/TreeSpeculativeResult.cs index 463dcc4e7b..ae19a2adc0 100644 --- a/src/Inference/SpeculativeDecoding/TreeSpeculativeResult.cs +++ b/src/Inference/SpeculativeDecoding/TreeSpeculativeResult.cs @@ -5,7 +5,7 @@ namespace AiDotNet.Inference.SpeculativeDecoding; /// /// Result of tree-based speculative decoding. /// -public class TreeSpeculativeResult +internal class TreeSpeculativeResult { /// /// All tokens (input + generated). diff --git a/src/Inference/SpeculativeDecoding/TreeStepStatistics.cs b/src/Inference/SpeculativeDecoding/TreeStepStatistics.cs index ca5d53f491..6ce7c203e7 100644 --- a/src/Inference/SpeculativeDecoding/TreeStepStatistics.cs +++ b/src/Inference/SpeculativeDecoding/TreeStepStatistics.cs @@ -3,7 +3,7 @@ namespace AiDotNet.Inference.SpeculativeDecoding; /// /// Statistics for a tree speculation step. /// -public class TreeStepStatistics +internal class TreeStepStatistics { /// Number of nodes in tree. public int TreeNodes { get; set; } diff --git a/src/InferenceOptimization/ARCHITECTURE.md b/src/InferenceOptimization/ARCHITECTURE.md new file mode 100644 index 0000000000..517e955471 --- /dev/null +++ b/src/InferenceOptimization/ARCHITECTURE.md @@ -0,0 +1,60 @@ +# AiDotNet Inference Optimization Architecture + +This document describes the internal structure of the `AiDotNet.InferenceOptimization` module and its extension points. + +## Design Goals + +- Hardware-aware CPU acceleration (SIMD when available) +- Deterministic behavior with safe fallbacks +- Low overhead when optimizations are disabled +- Extensible operator/kernels surface for future backends +- Thread-safe initialization and registration +- Optional profiling hooks for diagnosis + +## Key Components + +### OptimizationInitializer + +Responsibilities: +- One entrypoint to initialize platform detection and (optionally) profiling +- Ensures module initialization is safe to call multiple times + +### PlatformDetector + +Responsibilities: +- Detects process architecture and SIMD availability (x86/x64 and ARM) +- Exposes `PlatformCapabilities` used for selecting implementations + +Notes: +- Capability checks are runtime-based; unsupported intrinsics must always fall back to scalar implementations. + +### CustomOperatorRegistry + +Responsibilities: +- Registers multiple implementations per operation name +- Chooses the best supported implementation at runtime +- Caches the selection to avoid repeated capability checks + +### Kernels (`Kernels/*`) + +Responsibilities: +- Optimized building blocks for critical inference workloads: + - GEMM / matmul + - attention + - convolution + +Notes: +- Kernels are implemented with safe, span-based loops and use platform intrinsics only behind runtime capability checks. + +### CPU Helpers (`AiDotNet.Tensors/Engines/Optimization/*`) + +Responsibilities: +- Cache-aware helpers (tiling/transposition heuristics) +- Loop tiling/unrolling utilities where beneficial +- Optional profiling (`PerformanceProfiler`) for hotspot tracking + +## Integration Points + +- `AiDotNet.Inference.InferenceOptimizer` selects and applies inference-time implementations (e.g., attention variants, paged KV-cache) based on `InferenceOptimizationConfig`. +- `AiDotNet.Models.Results.PredictionModelResult` exposes facade-friendly entrypoints (`Predict`, `BeginInferenceSession`) while keeping internal complexity non-user-facing by default. + diff --git a/src/InferenceOptimization/CustomOperatorRegistry.cs b/src/InferenceOptimization/CustomOperatorRegistry.cs new file mode 100644 index 0000000000..cf661f5446 --- /dev/null +++ b/src/InferenceOptimization/CustomOperatorRegistry.cs @@ -0,0 +1,205 @@ +using System; +using System.Collections.Concurrent; +using System.Collections.Generic; +using System.Linq; + +namespace AiDotNet.InferenceOptimization +{ + /// + /// Thread-safe registry for managing custom operators with automatic fallback + /// + public sealed class CustomOperatorRegistry + { + private static readonly Lazy _instance = + new Lazy(() => new CustomOperatorRegistry()); + + private readonly ConcurrentDictionary> _operators; + private readonly ConcurrentDictionary _selectedOperators; + private readonly ConcurrentDictionary _operatorVersions; + + /// + /// Gets the singleton instance of the registry + /// + public static CustomOperatorRegistry Instance => _instance.Value; + + private CustomOperatorRegistry() + { + _operators = new ConcurrentDictionary>(); + _selectedOperators = new ConcurrentDictionary(); + _operatorVersions = new ConcurrentDictionary(); + } + + /// + /// Registers a custom operator + /// + public void Register(ICustomOperator op) + { + if (op == null) + throw new ArgumentNullException(nameof(op)); + + // Bump the version after the operator set is updated. + // This avoids stale cached selections without requiring coarse locking. + void BumpVersion() => _operatorVersions.AddOrUpdate(op.Name, 1, (_, v) => v + 1); + + // Use AddOrUpdate with factory that always creates a new sorted list + // This ensures thread-safety by never mutating existing lists + _operators.AddOrUpdate( + op.Name, + _ => new List { op }, + (_, existingList) => + { + // Create a new list with all existing operators plus the new one + // This avoids race conditions from modifying the existing list + List newList; + lock (existingList) + { + newList = new List(existingList) { op }; + } + newList.Sort((a, b) => b.Priority.CompareTo(a.Priority)); + return newList; + }); + + BumpVersion(); + } + + /// + /// Gets the best available operator for the given name + /// + public ICustomOperator? GetOperator(string name) + { + if (string.IsNullOrEmpty(name)) + throw new ArgumentException("Operator name cannot be null or empty", nameof(name)); + + while (true) + { + long version = _operatorVersions.GetOrAdd(name, 0); + + if (_selectedOperators.TryGetValue(name, out var existing) && existing.Version == version) + { + return existing.Operator is NullOperator ? null : existing.Operator; + } + + var selected = SelectOperatorOrNull(name); + + // Only publish the cached selection if the operator set version did not change while we were selecting. + if (_operatorVersions.TryGetValue(name, out var current) && current == version) + { + _selectedOperators[name] = new SelectedOperatorEntry(version, selected); + return selected is NullOperator ? null : selected; + } + + // Operator set changed while selecting; retry to avoid caching a stale choice. + } + } + + private ICustomOperator SelectOperatorOrNull(string name) + { + if (!_operators.TryGetValue(name, out var candidates)) + return new NullOperator(); + + lock (candidates) + { + // Find the highest priority supported operator + var result = candidates.FirstOrDefault(op => op.IsSupported()); + return result ?? new NullOperator(); + } + } + + /// + /// Gets a typed operator + /// + public ICustomOperator? GetOperator(string name) where T : struct + { + return GetOperator(name) as ICustomOperator; + } + + /// + /// Internal marker type for null operators + /// + private sealed class NullOperator : ICustomOperator + { + public string Name => string.Empty; + public string Version => string.Empty; + public int Priority => int.MinValue; + public bool IsSupported() => false; + public double EstimatedSpeedup() => 0; + } + + /// + /// Checks if an operator is available + /// + public bool HasOperator(string name) + { + return GetOperator(name) != null; + } + + /// + /// Unregisters all operators with the given name + /// + public void Unregister(string name) + { + _operators.TryRemove(name, out _); + _selectedOperators.TryRemove(name, out _); + _operatorVersions.TryRemove(name, out _); + } + + /// + /// Gets all registered operator names + /// + public IEnumerable GetRegisteredOperatorNames() + { + return _operators.Keys.ToArray(); + } + + /// + /// Gets detailed information about all registered operators + /// + public Dictionary> GetOperatorInfo() + { + var result = new Dictionary>(); + + foreach (var kvp in _operators) + { + lock (kvp.Value) + { + result[kvp.Key] = kvp.Value.Select(op => new OperatorInfo + { + Name = op.Name, + Version = op.Version, + Priority = op.Priority, + IsSupported = op.IsSupported(), + EstimatedSpeedup = op.EstimatedSpeedup(), + Type = op.GetType().FullName ?? op.GetType().Name + }).ToList(); + } + } + + return result; + } + + /// + /// Clears all registered operators + /// + public void Clear() + { + _operators.Clear(); + _selectedOperators.Clear(); + _operatorVersions.Clear(); + } + + private readonly record struct SelectedOperatorEntry(long Version, ICustomOperator Operator); + } + + /// + /// Information about a registered operator + /// + public class OperatorInfo + { + public string Name { get; set; } = string.Empty; + public string Version { get; set; } = string.Empty; + public int Priority { get; set; } + public bool IsSupported { get; set; } + public double EstimatedSpeedup { get; set; } + public string Type { get; set; } = string.Empty; + } +} diff --git a/src/InferenceOptimization/Examples/OptimizationExample.cs b/src/InferenceOptimization/Examples/OptimizationExample.cs deleted file mode 100644 index 049b0cccf5..0000000000 --- a/src/InferenceOptimization/Examples/OptimizationExample.cs +++ /dev/null @@ -1,242 +0,0 @@ -using AiDotNet.InferenceOptimization.Core; -using AiDotNet.Interfaces; - -namespace AiDotNet.InferenceOptimization.Examples; - -/// -/// Example usage of the inference optimization system. -/// -public class OptimizationExample -{ - /// - /// Example 1: Basic optimization of a simple CNN - /// - public static void BasicCNNOptimization() - { - Console.WriteLine("=== Example 1: Basic CNN Optimization ===\n"); - - // Create a simple CNN (pseudo-code, adapt to your model structure) - var layers = new List> - { - // Convolutional layer + BatchNorm + ReLU (will be fused) - // MaxPooling - // Another Conv + BatchNorm + ReLU (will be fused) - // Flatten - // Dense + Bias + ReLU (will be fused) - // Output Dense - }; - - // Build optimization graph - var graphBuilder = new GraphBuilder(); - var graph = graphBuilder.BuildFromLayers(layers); - - Console.WriteLine($"Original Graph: {graph.GetStatistics()}\n"); - - // Optimize with Standard level - var options = OptimizationOptions.FromLevel(OptimizationLevel.Standard); - options.PrintStatistics = true; - - var optimizer = new GraphOptimizer(options); - optimizer.Optimize(graph); - - Console.WriteLine("\nOptimization complete!"); - } - - /// - /// Example 2: Aggressive optimization for production deployment - /// - public static void ProductionOptimization() - { - Console.WriteLine("=== Example 2: Production Optimization ===\n"); - - // Create your model layers - var layers = new List>(); // Your layers here - - var graphBuilder = new GraphBuilder(); - var graph = graphBuilder.BuildFromLayers(layers); - - // Use Aggressive optimization for production - var options = new OptimizationOptions - { - Level = OptimizationLevel.Aggressive, - EnableOperatorFusion = true, - EnableMemoryReuse = true, - EnableCSE = true, - EnableInPlaceOptimization = true, - TargetLayout = "NCHW", // Optimize for GPU - PrintStatistics = true, - ValidateAfterEachPass = true - }; - - var optimizer = new GraphOptimizer(options); - optimizer.Optimize(graph); - - Console.WriteLine("\nProduction-ready optimized graph created!"); - } - - /// - /// Example 3: Custom optimization pass - /// - public static void CustomPassExample() - { - Console.WriteLine("=== Example 3: Custom Optimization Pass ===\n"); - - var graphBuilder = new GraphBuilder(); - var layers = new List>(); // Your layers - var graph = graphBuilder.BuildFromLayers(layers); - - // Create optimizer - var optimizer = new GraphOptimizer(); - - // Add custom pass (implement your own IOptimizationPass) - // optimizer.AddPass(new MyCustomPass()); - - optimizer.Optimize(graph); - - Console.WriteLine("Custom optimization applied!"); - } - - /// - /// Example 4: Comparing different optimization levels - /// - public static void CompareOptimizationLevels() - { - Console.WriteLine("=== Example 4: Comparing Optimization Levels ===\n"); - - var graphBuilder = new GraphBuilder(); - var layers = new List>(); // Your layers - var originalGraph = graphBuilder.BuildFromLayers(layers); - - var levels = new[] - { - OptimizationLevel.None, - OptimizationLevel.Basic, - OptimizationLevel.Standard, - OptimizationLevel.Aggressive, - OptimizationLevel.Maximum - }; - - foreach (var level in levels) - { - Console.WriteLine($"\n--- Testing {level} Level ---"); - - var options = OptimizationOptions.FromLevel(level); - options.PrintStatistics = true; - - var optimizer = new GraphOptimizer(options); - optimizer.Optimize(originalGraph.Clone()); - - Console.WriteLine($"Level {level} complete\n"); - } - } - - /// - /// Example 5: Transformer model optimization - /// - public static void TransformerOptimization() - { - Console.WriteLine("=== Example 5: Transformer Optimization ===\n"); - - // Build transformer graph - var graphBuilder = new GraphBuilder(); - var layers = new List> - { - // Multi-head attention (will be fused) - // Layer normalization - // Feed-forward: Dense + Bias + GELU (will be fused) - // Dense + Bias (will be fused) - // Layer normalization - // etc. - }; - - var graph = graphBuilder.BuildFromLayers(layers); - - // Optimize for transformer - var options = new OptimizationOptions - { - Level = OptimizationLevel.Aggressive, - EnableOperatorFusion = true, - EnableMemoryReuse = true, - PrintStatistics = true - }; - - var optimizer = new GraphOptimizer(options); - optimizer.Optimize(graph); - - Console.WriteLine("\nTransformer optimized!"); - Console.WriteLine("Expected speedup: 2-3x"); - } - - /// - /// Example 6: Memory-constrained optimization - /// - public static void MemoryConstrainedOptimization() - { - Console.WriteLine("=== Example 6: Memory-Constrained Optimization ===\n"); - - var graphBuilder = new GraphBuilder(); - var layers = new List>(); // Your layers - var graph = graphBuilder.BuildFromLayers(layers); - - // Prioritize memory optimizations - var options = new OptimizationOptions - { - Level = OptimizationLevel.Aggressive, - EnableMemoryReuse = true, - EnableInPlaceOptimization = true, - EnableOperatorFusion = true, // Also reduces memory - PrintStatistics = true - }; - - var optimizer = new GraphOptimizer(options); - optimizer.Optimize(graph); - - Console.WriteLine("\nMemory-optimized graph created!"); - Console.WriteLine("Expected memory reduction: 30-50%"); - } - - /// - /// Example 7: Inspect optimization passes - /// - public static void InspectPasses() - { - Console.WriteLine("=== Example 7: Inspect Optimization Passes ===\n"); - - var optimizer = new GraphOptimizer( - OptimizationOptions.FromLevel(OptimizationLevel.Aggressive) - ); - - var passes = optimizer.GetPasses(); - - Console.WriteLine($"Total passes: {passes.Count}\n"); - - foreach (var pass in passes) - { - Console.WriteLine($"- {pass.Name} ({pass.PassType})"); - } - } - - public static void Main(string[] args) - { - // Run all examples - BasicCNNOptimization(); - Console.WriteLine("\n" + new string('=', 60) + "\n"); - - ProductionOptimization(); - Console.WriteLine("\n" + new string('=', 60) + "\n"); - - CustomPassExample(); - Console.WriteLine("\n" + new string('=', 60) + "\n"); - - CompareOptimizationLevels(); - Console.WriteLine("\n" + new string('=', 60) + "\n"); - - TransformerOptimization(); - Console.WriteLine("\n" + new string('=', 60) + "\n"); - - MemoryConstrainedOptimization(); - Console.WriteLine("\n" + new string('=', 60) + "\n"); - - InspectPasses(); - } -} diff --git a/src/InferenceOptimization/ICustomOperator.cs b/src/InferenceOptimization/ICustomOperator.cs new file mode 100644 index 0000000000..7342855347 --- /dev/null +++ b/src/InferenceOptimization/ICustomOperator.cs @@ -0,0 +1,48 @@ +using System; +using AiDotNet.LinearAlgebra; + +namespace AiDotNet.InferenceOptimization +{ + /// + /// Defines the contract for custom operators with hardware-specific optimizations + /// + public interface ICustomOperator + { + /// + /// Gets the unique name of the operator + /// + string Name { get; } + + /// + /// Gets the version of the operator implementation + /// + string Version { get; } + + /// + /// Gets the priority level (higher values are preferred) + /// + int Priority { get; } + + /// + /// Determines if the operator can run on the current platform + /// + bool IsSupported(); + + /// + /// Estimates the relative performance gain over reference implementation + /// + /// Expected speedup multiplier (e.g., 2.0 for 2x speedup) + double EstimatedSpeedup(); + } + + /// + /// Base interface for custom operators that work with tensors + /// + public interface ICustomOperator : ICustomOperator where T : struct + { + /// + /// Executes the operator on input tensors + /// + Tensor Execute(params Tensor[] inputs); + } +} diff --git a/src/InferenceOptimization/Kernels/AttentionKernel.cs b/src/InferenceOptimization/Kernels/AttentionKernel.cs new file mode 100644 index 0000000000..1525264fd9 --- /dev/null +++ b/src/InferenceOptimization/Kernels/AttentionKernel.cs @@ -0,0 +1,328 @@ +using System; +using System.Threading.Tasks; +using AiDotNet.LinearAlgebra; +using AiDotNet.Tensors.Engines.Simd; + +namespace AiDotNet.InferenceOptimization.Kernels +{ + /// + /// Fused attention kernel for transformer models + /// Implements optimized scaled dot-product attention: softmax(QK^T/sqrt(d_k))V + /// + public class AttentionKernel : ICustomOperator + { + public string Name => "FusedAttention"; + public string Version => "1.0.0"; + public int Priority => 100; + + public AttentionKernel() { } + + public bool IsSupported() + { + return true; + } + + public double EstimatedSpeedup() + { + // Fused attention reduces memory traffic significantly + return 2.5; + } + + public Tensor Execute(params Tensor[] inputs) + { + if (inputs == null || inputs.Length < 3) + throw new ArgumentException("Attention requires Q, K, V tensors"); + + var q = inputs[0]; // [batch_size, seq_len_q, d_k] + var k = inputs[1]; // [batch_size, seq_len_k, d_k] + var v = inputs[2]; // [batch_size, seq_len_v, d_v] + + bool useMask = inputs.Length > 3; + Tensor? mask = useMask ? inputs[3] : null; + + return ExecuteInternal(q, k, v, mask, maskBatchModulo: q.Shape.Length == 3 ? q.Shape[0] : 0); + } + + private Tensor ExecuteInternal( + Tensor q, + Tensor k, + Tensor v, + Tensor? mask, + int maskBatchModulo) + { + if (q.Shape.Length != 3 || k.Shape.Length != 3 || v.Shape.Length != 3) + throw new ArgumentException("Attention requires 3D tensors [batch, seq_len, features]"); + + int batchSize = q.Shape[0]; + int seqLenQ = q.Shape[1]; + int seqLenK = k.Shape[1]; + int dK = q.Shape[2]; + int dV = v.Shape[2]; + + if (k.Shape[0] != batchSize || v.Shape[0] != batchSize) + throw new ArgumentException("Q, K, and V must have the same batch size"); + + if (k.Shape[2] != dK) + throw new ArgumentException("Q and K must have same feature dimension"); + + if (v.Shape[1] != seqLenK) + throw new ArgumentException("K and V must have same sequence length"); + + if (mask != null) + { + if (mask.Shape.Length != 3) + throw new ArgumentException("Attention mask must be a 3D tensor [batch, seq_len_q, seq_len_k]"); + + if (mask.Shape[1] != seqLenQ || mask.Shape[2] != seqLenK) + throw new ArgumentException("Attention mask must match [batch, seq_len_q, seq_len_k]"); + + if (maskBatchModulo <= 0) + { + if (mask.Shape[0] != batchSize) + throw new ArgumentException("Attention mask must have the same batch size as Q when used in Execute()"); + } + else + { + if (mask.Shape[0] != maskBatchModulo) + throw new ArgumentException("Attention mask batch dimension must match the provided maskBatchModulo"); + } + } + + var result = new Tensor(new[] { batchSize, seqLenQ, dV }); + + // Process each batch in parallel + Parallel.For(0, batchSize, b => + { + ProcessBatch(q, k, v, mask, result, b, seqLenQ, seqLenK, dK, dV, maskBatchModulo); + }); + + return result; + } + + private void ProcessBatch( + Tensor q, Tensor k, Tensor v, + Tensor? mask, Tensor result, + int batchIdx, int seqLenQ, int seqLenK, int dK, int dV, + int maskBatchModulo) + { + float scale = 1.0f / MathF.Sqrt(dK); + + // Extract batch slices + int qOffset = batchIdx * seqLenQ * dK; + int kOffset = batchIdx * seqLenK * dK; + int vOffset = batchIdx * seqLenK * dV; + int outOffset = batchIdx * seqLenQ * dV; + + // Compute attention scores: QK^T + var scores = new float[seqLenQ * seqLenK]; + + for (int i = 0; i < seqLenQ; i++) + { + int qRowOffset = qOffset + i * dK; + var qRow = q.Data.AsSpan(qRowOffset, dK); + + for (int j = 0; j < seqLenK; j++) + { + int kRowOffset = kOffset + j * dK; + var kRow = k.Data.AsSpan(kRowOffset, dK); + float score = SimdKernels.DotProduct(qRow, kRow) * scale; + + // Apply mask if provided + if (mask != null) + { + int effectiveMaskBatch = maskBatchModulo > 0 ? (batchIdx % maskBatchModulo) : batchIdx; + int maskIdx = effectiveMaskBatch * seqLenQ * seqLenK + i * seqLenK + j; + // Use epsilon-based comparison for floating point equality + if (MathF.Abs(mask.Data[maskIdx]) < 1e-6f) + { + score = float.NegativeInfinity; + } + } + + scores[i * seqLenK + j] = score; + } + } + + // Apply softmax over each row + ApplySoftmax(scores, seqLenQ, seqLenK); + + // Compute weighted sum: attention_weights * V + for (int i = 0; i < seqLenQ; i++) + { + var outRow = result.Data.AsSpan(outOffset + i * dV, dV); + outRow.Clear(); + + // Accumulate weighted values + for (int j = 0; j < seqLenK; j++) + { + float weight = scores[i * seqLenK + j]; + if (weight <= 0f) + { + continue; + } + + var vRow = v.Data.AsSpan(vOffset + j * dV, dV); + SimdKernels.ScalarMultiplyAdd(outRow, vRow, weight, outRow); + } + } + } + + private void ApplySoftmax(float[] data, int rows, int cols) + { + for (int i = 0; i < rows; i++) + { + int rowOffset = i * cols; + + // Find max for numerical stability + float maxVal = float.NegativeInfinity; + for (int j = 0; j < cols; j++) + { + float v = data[rowOffset + j]; + if (v > maxVal) + { + maxVal = v; + } + } + + // Compute exp and sum + float sum = 0.0f; + for (int j = 0; j < cols; j++) + { + int idx = rowOffset + j; + float v = data[idx]; + if (float.IsNegativeInfinity(v)) + { + data[idx] = 0.0f; + continue; + } + + float ev = MathF.Exp(v - maxVal); + data[idx] = ev; + sum += ev; + } + + // Normalize + if (sum > 0.0f) + { + float invSum = 1.0f / sum; + for (int j = 0; j < cols; j++) + { + data[rowOffset + j] *= invSum; + } + } + } + } + + /// + /// Multi-head attention variant + /// + public Tensor MultiHeadAttention( + Tensor q, Tensor k, Tensor v, + int numHeads, Tensor? mask = null) + { + if (q.Shape.Length != 3 || k.Shape.Length != 3 || v.Shape.Length != 3) + throw new ArgumentException("Multi-head attention requires 3D tensors"); + + int batchSize = q.Shape[0]; + int dModel = q.Shape[2]; + + if (k.Shape[0] != batchSize || v.Shape[0] != batchSize) + throw new ArgumentException("Q, K, and V must have the same batch size"); + + if (dModel % numHeads != 0) + throw new ArgumentException("d_model must be divisible by num_heads"); + + int dK = dModel / numHeads; + + if (k.Shape[2] != dModel || v.Shape[2] != dModel) + throw new ArgumentException("Q, K, and V must have the same feature dimension (d_model)"); + + if (v.Shape[1] != k.Shape[1]) + throw new ArgumentException("K and V must have the same sequence length"); + + // Reshape to [batch * num_heads, seq_len, d_k] + var qReshaped = ReshapeForMultiHead(q, numHeads, dK); + var kReshaped = ReshapeForMultiHead(k, numHeads, dK); + var vReshaped = ReshapeForMultiHead(v, numHeads, dK); + + // Apply attention + Tensor attended; + if (mask is null) + { + attended = ExecuteInternal(qReshaped, kReshaped, vReshaped, mask: null, maskBatchModulo: 0); + } + else + { + int expectedPerHeadBatch = batchSize * numHeads; + if (mask.Shape.Length != 3) + throw new ArgumentException("Multi-head attention mask must be a 3D tensor"); + + if (mask.Shape[1] != q.Shape[1] || mask.Shape[2] != k.Shape[1]) + throw new ArgumentException("Multi-head attention mask must match [batch, seq_len_q, seq_len_k]"); + + // Accept either per-batch mask [B, SQ, SK] (broadcast across heads) or per-head mask [B*H, SQ, SK]. + int maskBatchModulo = mask.Shape[0] switch + { + int b when b == expectedPerHeadBatch => 0, + int b when b == batchSize => batchSize, + _ => throw new ArgumentException("Multi-head attention mask must have batch dimension B or B*numHeads"), + }; + + attended = ExecuteInternal(qReshaped, kReshaped, vReshaped, mask, maskBatchModulo); + } + + // Reshape back to [batch, seq_len, d_model] + return ReshapeFromMultiHead(attended, batchSize, q.Shape[1], dModel); + } + + private Tensor ReshapeForMultiHead(Tensor input, int numHeads, int dK) + { + int batchSize = input.Shape[0]; + int seqLen = input.Shape[1]; + var reshaped = new Tensor(new[] { batchSize * numHeads, seqLen, dK }); + + for (int b = 0; b < batchSize; b++) + { + for (int h = 0; h < numHeads; h++) + { + for (int s = 0; s < seqLen; s++) + { + for (int d = 0; d < dK; d++) + { + int srcIdx = b * seqLen * numHeads * dK + s * numHeads * dK + h * dK + d; + int dstIdx = (b * numHeads + h) * seqLen * dK + s * dK + d; + reshaped.Data[dstIdx] = input.Data[srcIdx]; + } + } + } + } + + return reshaped; + } + + private Tensor ReshapeFromMultiHead(Tensor input, int batchSize, int seqLen, int dModel) + { + var reshaped = new Tensor(new[] { batchSize, seqLen, dModel }); + int numHeads = input.Shape[0] / batchSize; + int dK = input.Shape[2]; + + for (int b = 0; b < batchSize; b++) + { + for (int h = 0; h < numHeads; h++) + { + for (int s = 0; s < seqLen; s++) + { + for (int d = 0; d < dK; d++) + { + int srcIdx = (b * numHeads + h) * seqLen * dK + s * dK + d; + int dstIdx = b * seqLen * dModel + s * dModel + h * dK + d; + reshaped.Data[dstIdx] = input.Data[srcIdx]; + } + } + } + } + + return reshaped; + } + } +} diff --git a/src/InferenceOptimization/Kernels/ConvolutionKernel.cs b/src/InferenceOptimization/Kernels/ConvolutionKernel.cs new file mode 100644 index 0000000000..7762de502e --- /dev/null +++ b/src/InferenceOptimization/Kernels/ConvolutionKernel.cs @@ -0,0 +1,384 @@ +using System; +using System.Threading.Tasks; +using AiDotNet.LinearAlgebra; +using AiDotNet.Tensors.Engines; + +namespace AiDotNet.InferenceOptimization.Kernels +{ + /// + /// Optimized convolution kernels including depthwise and group convolutions + /// + public class ConvolutionKernel : ICustomOperator + { + public string Name => "Convolution"; + public string Version => "1.0.0"; + public int Priority => 100; + + public bool IsSupported() + { + return true; + } + + public double EstimatedSpeedup() + { + var caps = PlatformDetector.Capabilities; + if (caps.HasAVX2) return 2.5; + if (caps.HasNeon) return 2.0; + return 1.5; + } + + /// + /// Executes convolution on the provided inputs. + /// Expects 2-3 inputs: input tensor, kernel tensor, and optional config tensor. + /// Config tensor format: [stride, padding] (defaults to stride=1, padding=0) + /// + public Tensor Execute(params Tensor[] inputs) + { + if (inputs == null || inputs.Length < 2) + { + throw new ArgumentException( + "ConvolutionKernel requires at least 2 inputs: input tensor and kernel tensor. " + + "Optional 3rd input for config [stride, padding]."); + } + + var input = inputs[0]; + var kernel = inputs[1]; + + // Extract stride and padding from optional config tensor or use defaults + int stride = 1; + int padding = 0; + + if (inputs.Length >= 3 && inputs[2] != null && inputs[2].Data.Length >= 2) + { + stride = Math.Max(1, (int)inputs[2].Data[0]); + padding = Math.Max(0, (int)inputs[2].Data[1]); + } + + // Determine convolution type based on kernel shape + // Standard: kernel[out_channels, in_channels, kH, kW] + // Depthwise: kernel[channels, 1, kH, kW] + if (kernel.Shape.Length == 4 && kernel.Shape[1] == 1) + { + // Depthwise convolution (kernel has 1 in_channel dimension) + return DepthwiseConv2D(input, kernel, stride, padding); + } + + // Default to standard 2D convolution + return Conv2D(input, kernel, stride, padding); + } + + /// + /// Standard 2D convolution + /// + public Tensor Conv2D( + Tensor input, + Tensor kernel, + int stride = 1, + int padding = 0) + { + // Input: [batch, in_channels, height, width] + // Kernel: [out_channels, in_channels, kernel_h, kernel_w] + + if (input.Shape.Length != 4 || kernel.Shape.Length != 4) + throw new ArgumentException("Conv2D requires 4D tensors"); + + if (stride <= 0) + throw new ArgumentOutOfRangeException(nameof(stride), $"stride must be positive, but got {stride}"); + + if (padding < 0) + throw new ArgumentOutOfRangeException(nameof(padding), $"padding must be non-negative, but got {padding}"); + + int batchSize = input.Shape[0]; + int inChannels = input.Shape[1]; + int inHeight = input.Shape[2]; + int inWidth = input.Shape[3]; + + int outChannels = kernel.Shape[0]; + int kernelH = kernel.Shape[2]; + int kernelW = kernel.Shape[3]; + + if (kernelH <= 0 || kernelW <= 0) + throw new ArgumentException($"Kernel dimensions must be positive, but got {kernelH}x{kernelW}"); + + if (kernel.Shape[1] != inChannels) + throw new ArgumentException($"Conv2D requires kernel.Shape[1] == inChannels ({inChannels}), but got {kernel.Shape[1]}"); + + int outHeight = (inHeight + 2 * padding - kernelH) / stride + 1; + int outWidth = (inWidth + 2 * padding - kernelW) / stride + 1; + + if (outHeight <= 0 || outWidth <= 0) + throw new ArgumentException( + $"Invalid output dimensions ({outHeight}x{outWidth}). " + + $"Check stride ({stride}), padding ({padding}), and kernel size ({kernelH}x{kernelW})."); + var output = new Tensor(new[] { batchSize, outChannels, outHeight, outWidth }); + + // Parallelize over batch and output channels + Parallel.For(0, batchSize * outChannels, idx => + { + int b = idx / outChannels; + int oc = idx % outChannels; + + Conv2DSingleOutput(input, kernel, output, b, oc, + inChannels, inHeight, inWidth, + kernelH, kernelW, stride, padding, + outHeight, outWidth); + }); + + return output; + } + + private void Conv2DSingleOutput( + Tensor input, Tensor kernel, Tensor output, + int batch, int outChannel, + int inChannels, int inHeight, int inWidth, + int kernelH, int kernelW, int stride, int padding, + int outHeight, int outWidth) + { + var inputData = input.Data; + var kernelData = kernel.Data; + var outputData = output.Data; + + for (int oh = 0; oh < outHeight; oh++) + { + for (int ow = 0; ow < outWidth; ow++) + { + float sum = 0.0f; + + for (int ic = 0; ic < inChannels; ic++) + { + for (int kh = 0; kh < kernelH; kh++) + { + for (int kw = 0; kw < kernelW; kw++) + { + int ih = oh * stride - padding + kh; + int iw = ow * stride - padding + kw; + + if (ih >= 0 && ih < inHeight && iw >= 0 && iw < inWidth) + { + int inputIdx = ((batch * inChannels + ic) * inHeight + ih) * inWidth + iw; + int kernelIdx = ((outChannel * inChannels + ic) * kernelH + kh) * kernelW + kw; + sum += inputData[inputIdx] * kernelData[kernelIdx]; + } + } + } + } + + int outputIdx = ((batch * output.Shape[1] + outChannel) * outHeight + oh) * outWidth + ow; + outputData[outputIdx] = sum; + } + } + } + + /// + /// Depthwise separable convolution (more efficient for mobile architectures) + /// + public Tensor DepthwiseConv2D( + Tensor input, + Tensor kernel, + int stride = 1, + int padding = 0) + { + // Input: [batch, channels, height, width] + // Kernel: [channels, 1, kernel_h, kernel_w] + + if (input.Shape.Length != 4 || kernel.Shape.Length != 4) + throw new ArgumentException("DepthwiseConv2D requires 4D tensors"); + + if (stride <= 0) + throw new ArgumentOutOfRangeException(nameof(stride), $"stride must be positive, but got {stride}"); + + if (padding < 0) + throw new ArgumentOutOfRangeException(nameof(padding), $"padding must be non-negative, but got {padding}"); + + int batchSize = input.Shape[0]; + int channels = input.Shape[1]; + int inHeight = input.Shape[2]; + int inWidth = input.Shape[3]; + + int kernelH = kernel.Shape[2]; + int kernelW = kernel.Shape[3]; + + if (kernelH <= 0 || kernelW <= 0) + throw new ArgumentException($"Kernel dimensions must be positive, but got {kernelH}x{kernelW}"); + + int outHeight = (inHeight + 2 * padding - kernelH) / stride + 1; + int outWidth = (inWidth + 2 * padding - kernelW) / stride + 1; + + if (outHeight <= 0 || outWidth <= 0) + throw new ArgumentException( + $"Invalid output dimensions ({outHeight}x{outWidth}). " + + $"Check stride ({stride}), padding ({padding}), and kernel size ({kernelH}x{kernelW})."); + + if (kernel.Shape[1] != 1) + throw new ArgumentException( + $"Depthwise convolution requires kernel.Shape[1] == 1, but got {kernel.Shape[1]}"); + + if (kernel.Shape[0] != channels) + throw new ArgumentException( + $"Depthwise convolution requires kernel.Shape[0] == channels ({channels}), but got {kernel.Shape[0]}"); + + var output = new Tensor(new[] { batchSize, channels, outHeight, outWidth }); + + Parallel.For(0, batchSize * channels, idx => + { + int b = idx / channels; + int c = idx % channels; + + DepthwiseConv2DSingleChannel(input, kernel, output, b, c, + inHeight, inWidth, kernelH, kernelW, + stride, padding, outHeight, outWidth); + }); + + return output; + } + + private void DepthwiseConv2DSingleChannel( + Tensor input, Tensor kernel, Tensor output, + int batch, int channel, + int inHeight, int inWidth, int kernelH, int kernelW, + int stride, int padding, int outHeight, int outWidth) + { + var inputData = input.Data; + var kernelData = kernel.Data; + var outputData = output.Data; + + for (int oh = 0; oh < outHeight; oh++) + { + for (int ow = 0; ow < outWidth; ow++) + { + float sum = 0.0f; + + for (int kh = 0; kh < kernelH; kh++) + { + for (int kw = 0; kw < kernelW; kw++) + { + int ih = oh * stride - padding + kh; + int iw = ow * stride - padding + kw; + + if (ih >= 0 && ih < inHeight && iw >= 0 && iw < inWidth) + { + int inputIdx = ((batch * input.Shape[1] + channel) * inHeight + ih) * inWidth + iw; + int kernelIdx = (channel * kernelH + kh) * kernelW + kw; + sum += inputData[inputIdx] * kernelData[kernelIdx]; + } + } + } + + int outputIdx = ((batch * output.Shape[1] + channel) * outHeight + oh) * outWidth + ow; + outputData[outputIdx] = sum; + } + } + } + + /// + /// Group convolution (reduces parameters and computation) + /// + public Tensor GroupConv2D( + Tensor input, + Tensor kernel, + int groups, + int stride = 1, + int padding = 0) + { + if (input.Shape.Length != 4 || kernel.Shape.Length != 4) + throw new ArgumentException("GroupConv2D requires 4D tensors"); + + int batchSize = input.Shape[0]; + int inChannels = input.Shape[1]; + int inHeight = input.Shape[2]; + int inWidth = input.Shape[3]; + + int outChannels = kernel.Shape[0]; + int kernelH = kernel.Shape[2]; + int kernelW = kernel.Shape[3]; + + if (groups <= 0) + throw new ArgumentOutOfRangeException(nameof(groups), "groups must be positive."); + + if (inChannels % groups != 0 || outChannels % groups != 0) + throw new ArgumentException("Channels must be divisible by groups"); + + int inChannelsPerGroup = inChannels / groups; + int outChannelsPerGroup = outChannels / groups; + + if (kernel.Shape[1] != inChannelsPerGroup) + throw new ArgumentException( + $"Group convolution requires kernel.Shape[1] == inChannelsPerGroup ({inChannelsPerGroup}), " + + $"but got {kernel.Shape[1]}"); + + int outHeight = (inHeight + 2 * padding - kernelH) / stride + 1; + int outWidth = (inWidth + 2 * padding - kernelW) / stride + 1; + + if (outHeight <= 0 || outWidth <= 0) + throw new ArgumentException( + $"Invalid output dimensions ({outHeight}x{outWidth}). " + + $"Check stride ({stride}), padding ({padding}), and kernel size ({kernelH}x{kernelW})."); + + var output = new Tensor(new[] { batchSize, outChannels, outHeight, outWidth }); + + // Process each group independently + Parallel.For(0, groups, g => + { + for (int b = 0; b < batchSize; b++) + { + for (int oc = 0; oc < outChannelsPerGroup; oc++) + { + int globalOutChannel = g * outChannelsPerGroup + oc; + + GroupConv2DSingleOutput(input, kernel, output, b, globalOutChannel, g, + inChannelsPerGroup, inHeight, inWidth, + kernelH, kernelW, stride, padding, + outHeight, outWidth); + } + } + }); + + return output; + } + + private void GroupConv2DSingleOutput( + Tensor input, Tensor kernel, Tensor output, + int batch, int outChannel, int group, + int inChannelsPerGroup, int inHeight, int inWidth, + int kernelH, int kernelW, int stride, int padding, + int outHeight, int outWidth) + { + int inChannelStart = group * inChannelsPerGroup; + var inputData = input.Data; + var kernelData = kernel.Data; + var outputData = output.Data; + + for (int oh = 0; oh < outHeight; oh++) + { + for (int ow = 0; ow < outWidth; ow++) + { + float sum = 0.0f; + + for (int ic = 0; ic < inChannelsPerGroup; ic++) + { + int globalInChannel = inChannelStart + ic; + + for (int kh = 0; kh < kernelH; kh++) + { + for (int kw = 0; kw < kernelW; kw++) + { + int ih = oh * stride - padding + kh; + int iw = ow * stride - padding + kw; + + if (ih >= 0 && ih < inHeight && iw >= 0 && iw < inWidth) + { + int inputIdx = ((batch * input.Shape[1] + globalInChannel) * inHeight + ih) * inWidth + iw; + int kernelIdx = ((outChannel * inChannelsPerGroup + ic) * kernelH + kh) * kernelW + kw; + sum += inputData[inputIdx] * kernelData[kernelIdx]; + } + } + } + } + + int outputIdx = ((batch * output.Shape[1] + outChannel) * outHeight + oh) * outWidth + ow; + outputData[outputIdx] = sum; + } + } + } + } +} diff --git a/src/InferenceOptimization/Kernels/GemmKernel.cs b/src/InferenceOptimization/Kernels/GemmKernel.cs new file mode 100644 index 0000000000..67fd771a90 --- /dev/null +++ b/src/InferenceOptimization/Kernels/GemmKernel.cs @@ -0,0 +1,176 @@ +using System; +using System.Runtime.CompilerServices; +using System.Threading.Tasks; +using AiDotNet.LinearAlgebra; +using AiDotNet.Tensors.Engines; +using AiDotNet.Tensors.Engines.Simd; + +namespace AiDotNet.InferenceOptimization.Kernels +{ + /// + /// Optimized General Matrix Multiplication (GEMM) kernel + /// Implements cache-aware blocked matrix multiplication with SIMD + /// + public class GemmKernel : ICustomOperator + { + private const int BlockSize = 64; // Tuned for typical L1 cache + private const int MinParallelSize = 256; // Minimum size for parallel execution + + public string Name => "GEMM"; + public string Version => "1.0.0"; + public int Priority => 100; + + public bool IsSupported() + { + // GEMM is always supported, but performance varies by platform + return true; + } + + public double EstimatedSpeedup() + { + var caps = PlatformDetector.Capabilities; + if (caps.HasAVX2) return 3.0; + if (caps.HasSSE42) return 2.0; + if (caps.HasNeon) return 2.5; + return 1.5; + } + + public Tensor Execute(params Tensor[] inputs) + { + if (inputs == null || inputs.Length < 2) + throw new ArgumentException("GEMM requires at least 2 input tensors"); + + var a = inputs[0]; + var b = inputs[1]; + + if (a.Shape.Length != 2 || b.Shape.Length != 2) + throw new ArgumentException("GEMM requires 2D tensors (matrices)"); + + int m = a.Shape[0]; + int k = a.Shape[1]; + int n = b.Shape[1]; + + if (k != b.Shape[0]) + throw new ArgumentException($"Matrix dimensions incompatible: ({m}x{k}) * ({b.Shape[0]}x{n})"); + + var result = new Tensor(new[] { m, n }); + + // Choose strategy based on matrix size + if (m * n * k < MinParallelSize * MinParallelSize) + { + GemmBlocked(a.Data, b.Data, result.Data, m, n, k); + } + else + { + GemmParallel(a.Data, b.Data, result.Data, m, n, k); + } + + return result; + } + + /// + /// Cache-blocked GEMM implementation + /// + [MethodImpl(MethodImplOptions.AggressiveInlining)] + private void GemmBlocked(float[] A, float[] B, float[] C, int M, int N, int K) + { + // Blocked algorithm for cache efficiency + for (int i = 0; i < M; i += BlockSize) + { + int iMax = Math.Min(i + BlockSize, M); + + for (int j = 0; j < N; j += BlockSize) + { + int jMax = Math.Min(j + BlockSize, N); + int spanLen = jMax - j; + + for (int k = 0; k < K; k += BlockSize) + { + int kMax = Math.Min(k + BlockSize, K); + + // Process block + for (int ii = i; ii < iMax; ii++) + { + for (int kk = k; kk < kMax; kk++) + { + float aVal = A[ii * K + kk]; + var bRow = B.AsSpan(kk * N + j, spanLen); + var cRow = C.AsSpan(ii * N + j, spanLen); + + // SIMD-optimized inner loop: cRow = cRow + aVal * bRow + SimdKernels.ScalarMultiplyAdd(cRow, bRow, aVal, cRow); + } + } + } + } + } + } + + /// + /// Parallel GEMM implementation for large matrices + /// + [MethodImpl(MethodImplOptions.AggressiveInlining)] + private void GemmParallel(float[] A, float[] B, float[] C, int M, int N, int K) + { + // Parallelize over rows of A + Parallel.For(0, (M + BlockSize - 1) / BlockSize, iBlock => + { + int i = iBlock * BlockSize; + int iMax = Math.Min(i + BlockSize, M); + + for (int j = 0; j < N; j += BlockSize) + { + int jMax = Math.Min(j + BlockSize, N); + int spanLen = jMax - j; + + for (int k = 0; k < K; k += BlockSize) + { + int kMax = Math.Min(k + BlockSize, K); + + for (int ii = i; ii < iMax; ii++) + { + for (int kk = k; kk < kMax; kk++) + { + float aVal = A[ii * K + kk]; + var bRow = B.AsSpan(kk * N + j, spanLen); + var cRow = C.AsSpan(ii * N + j, spanLen); + + SimdKernels.ScalarMultiplyAdd(cRow, bRow, aVal, cRow); + } + } + } + } + }); + } + + /// + /// Matrix multiplication with transpose B optimization (C = A * B^T) + /// + public Tensor GemmTransposeB(Tensor a, Tensor b) + { + if (a.Shape.Length != 2 || b.Shape.Length != 2) + throw new ArgumentException("GemmTransposeB requires 2D tensors"); + + int m = a.Shape[0]; + int k = a.Shape[1]; + int n = b.Shape[0]; // Note: B is transposed + + if (k != b.Shape[1]) + throw new ArgumentException("Matrix dimensions incompatible for transpose"); + + var result = new Tensor(new[] { m, n }); + + Parallel.For(0, m, i => + { + var rowA = a.Data.AsSpan(i * k, k); + for (int j = 0; j < n; j++) + { + var rowB = b.Data.AsSpan(j * k, k); + result.Data[i * n + j] = SimdKernels.DotProduct(rowA, rowB); + } + }); + + return result; + } + } +} diff --git a/src/InferenceOptimization/OptimizationInitializer.cs b/src/InferenceOptimization/OptimizationInitializer.cs new file mode 100644 index 0000000000..2ae7068aa5 --- /dev/null +++ b/src/InferenceOptimization/OptimizationInitializer.cs @@ -0,0 +1,108 @@ +using System; +using AiDotNet.InferenceOptimization.Kernels; +using AiDotNet.Tensors.Engines; +using AiDotNet.Tensors.Engines.Optimization; + +namespace AiDotNet.InferenceOptimization +{ + /// + /// Initializes and registers all optimized kernels and operators + /// + public static class OptimizationInitializer + { + private static bool _initialized = false; + private static readonly object _lock = new object(); + + /// + /// Initializes the inference optimization system + /// + public static void Initialize(bool enableProfiling = false) + { + lock (_lock) + { + if (_initialized) + return; + + // Enable profiling if requested + PerformanceProfiler.Instance.Enabled = enableProfiling; + + // Register optimized kernels + RegisterKernels(); + + // Print platform capabilities + LogPlatformInfo(); + + _initialized = true; + } + } + + private static void RegisterKernels() + { + var registry = CustomOperatorRegistry.Instance; + + // Register GEMM kernel + registry.Register(new GemmKernel()); + + // Register Attention kernel + registry.Register(new AttentionKernel()); + + // Register Convolution kernel + registry.Register(new ConvolutionKernel()); + + // Future: Register GPU kernels when available + // if (PlatformDetector.Capabilities.HasCudaSupport) + // { + // registry.Register(new CudaGemmKernel()); + // registry.Register(new CudaConvolutionKernel()); + // } + } + + private static void LogPlatformInfo() + { + Console.WriteLine("=== AiDotNet Inference Optimization ==="); + Console.WriteLine(PlatformDetector.GetCapabilitiesDescription()); + Console.WriteLine(); + Console.WriteLine("Registered Operators:"); + + var operatorInfo = CustomOperatorRegistry.Instance.GetOperatorInfo(); + foreach (var kvp in operatorInfo) + { + Console.WriteLine($" {kvp.Key}:"); + foreach (var info in kvp.Value) + { + var status = info.IsSupported ? "✓" : "✗"; + Console.WriteLine($" {status} {info.Version} - Priority: {info.Priority}, Speedup: {info.EstimatedSpeedup:F1}x"); + } + } + Console.WriteLine(); + } + + /// + /// Gets a performance summary + /// + public static string GetPerformanceSummary() + { + if (!_initialized) + return "Optimization system not initialized."; + + var report = PerformanceProfiler.Instance.GenerateReport(); + return report; + } + + /// + /// Resets all profiling statistics + /// + public static void ResetStatistics() + { + PerformanceProfiler.Instance.Clear(); + } + + /// + /// Enables or disables profiling at runtime + /// + public static void SetProfilingEnabled(bool enabled) + { + PerformanceProfiler.Instance.Enabled = enabled; + } + } +} diff --git a/src/InferenceOptimization/README.md b/src/InferenceOptimization/README.md index 910f806827..7275098b80 100644 --- a/src/InferenceOptimization/README.md +++ b/src/InferenceOptimization/README.md @@ -1,340 +1,253 @@ -# Inference Optimization +# AiDotNet Inference Optimization -This module provides graph-level optimizations for neural network inference in AiDotNet. It implements operator fusion, graph transformations, memory optimizations, and computation optimizations to achieve 2-5x inference speedup. +This module provides low-level kernel optimization for critical operations, enabling hardware-specific acceleration for efficient AI model inference. ## Features -### Operator Fusion (Critical for Performance) - -Combines multiple operations into single optimized kernels: - -- **Conv + BatchNorm + ReLU**: Fuses the most common CNN pattern (ResNet, VGG, etc.) -- **Conv + BatchNorm**: Folds batch normalization into convolution weights -- **MatMul + Bias + Activation**: Optimizes transformer feed-forward networks -- **MatMul + Bias**: Gemm operation for fully connected layers -- **Elementwise Fusion**: Chains multiple elementwise operations -- **Multi-Head Attention**: Optimized attention computation - -**Expected Speedup**: 2-3x for CNN models, 1.5-2x for transformers - -### Graph Optimization - -Structural optimizations to simplify computation graphs: - -- **Constant Folding**: Pre-computes constant expressions -- **Dead Code Elimination**: Removes unused operations -- **Common Subexpression Elimination (CSE)**: Shares identical computations -- **Layout Optimization**: Optimizes NCHW vs NHWC for target hardware - -**Expected Speedup**: 1.2-1.5x additional speedup - -### Memory Optimization - -Reduces memory footprint during inference: - -- **In-Place Operations**: ReLU, Dropout, and other operations modify tensors in-place -- **Memory Reuse**: Shares memory buffers across non-overlapping lifetimes -- **Activation Memory Planning**: Optimal memory allocation strategy - -**Memory Reduction**: 30-50% for typical models - -### Computation Optimization - -Replaces expensive operations with cheaper equivalents: - -- **Algebraic Simplification**: x*1=x, x+0=x, x*0=0, etc. -- **Strength Reduction**: x^2 → x*x, x/2 → x*0.5 - -**Expected Speedup**: 1.1-1.3x additional speedup - -## Usage - -### Basic Example +### 1. Custom Operator Registration System +- Thread-safe operator registry with automatic fallback +- Priority-based operator selection +- Support for multiple implementations per operation +- Runtime operator switching based on platform capabilities + +### 2. Platform Detection +- Automatic detection of CPU architecture (x86/x64, ARM) +- SIMD instruction set detection (SSE, AVX, AVX2, AVX-512, NEON) +- Cache size estimation +- GPU capability detection (CUDA, OpenCL) + +### 3. SIMD Vectorization +- AVX2/AVX-512 optimized kernels for x86/x64 +- ARM NEON optimized kernels +- Automatic fallback to scalar implementations +- Optimized operations: + - Vector addition/multiplication + - Dot product with FMA support + - ReLU activation + - Sum reduction + - Scalar multiply-add + +### 4. Optimized Kernels + +#### GEMM (General Matrix Multiplication) +- Cache-blocked algorithm for L1 cache efficiency +- Parallel execution for large matrices +- SIMD-optimized inner loops +- Transpose optimization for better memory access patterns +- Expected speedup: 2-3x on AVX2, 2.5x on NEON + +#### Fused Attention Kernel +- Scaled dot-product attention: `softmax(QK^T/sqrt(d_k))V` +- Multi-head attention support +- Memory-efficient implementation +- Mask support for causal attention +- Expected speedup: 2.5x + +#### Convolution Kernels +- Standard 2D convolution +- Depthwise separable convolution +- Group convolution +- Parallel batch processing +- Expected speedup: 2-2.5x + +### 5. CPU Optimizations + +#### Cache Optimizer +- L1/L2/L3 cache-aware algorithms +- Automatic tiling parameter computation +- Prefetching for reduced latency +- Cache-aware transpose +- Z-order (Morton) indexing for 2D access patterns +- Cache miss estimation + +#### Loop Optimizer +- 2D and 3D loop tiling +- Loop unrolling (4x, 8x) +- Strip mining for cache utilization +- Loop fusion +- Loop interchange optimization +- Parallel tiling with work stealing + +### 6. Performance Profiling +- Thread-safe operation tracking +- Timing and memory usage statistics +- Per-operation metrics (min/avg/max/total) +- Performance report generation +- Runtime enable/disable capability + +### 7. GPU Optimization Infrastructure +- Base classes for GPU kernel implementations +- Memory management abstractions +- CUDA kernel base (ready for ILGPU/ManagedCuda integration) +- Device capability querying + +## Quick Start ```csharp -using AiDotNet.InferenceOptimization.Core; -using AiDotNet.InferenceOptimization.Passes; - -// Build a computation graph from your layers -var graphBuilder = new GraphBuilder(); -var graph = graphBuilder.BuildFromLayers(myLayers); - -// Create an optimizer with standard optimizations -var optimizer = new GraphOptimizer(); - -// Optimize the graph -var optimizedGraph = optimizer.Optimize(graph); - -// Use optimized graph for inference -// (Integration with NeuralNetworkBase is automatic) +using AiDotNet.InferenceOptimization; +using AiDotNet.InferenceOptimization.Kernels; +using AiDotNet.Tensors.Engines.Simd; // SimdKernels location +using AiDotNet.Tensors.LinearAlgebra; + +// Initialize the optimization system +OptimizationInitializer.Initialize(enableProfiling: true); + +// Use optimized GEMM +var gemmKernel = new GemmKernel(); +var a = new Tensor(new[] { 1000, 500 }); +var b = new Tensor(new[] { 500, 1000 }); +var result = gemmKernel.Execute(a, b); + +// Use fused attention +var attentionKernel = new AttentionKernel(); +var q = new Tensor(new[] { 1, 128, 64 }); // [batch, seq_len, d_k] +var k = new Tensor(new[] { 1, 128, 64 }); +var v = new Tensor(new[] { 1, 128, 64 }); +var attended = attentionKernel.Execute(q, k, v); + +// Get performance report +var report = OptimizationInitializer.GetPerformanceSummary(); +Console.WriteLine(report); ``` -### Advanced Example with Custom Options +## Platform Capabilities + +Check what optimizations are available on your platform: ```csharp -// Configure optimization options -var options = new OptimizationOptions -{ - Level = OptimizationLevel.Aggressive, - EnableOperatorFusion = true, - EnableMemoryReuse = true, - EnableCSE = true, - TargetLayout = "NCHW", // For GPU inference - PrintStatistics = true, - MaxIterations = 10 -}; - -// Create optimizer with options -var optimizer = new GraphOptimizer(options); - -// Add custom optimization pass -optimizer.AddPass(new MyCustomPass()); - -// Optimize -var optimizedGraph = optimizer.Optimize(graph); +var caps = PlatformDetector.Capabilities; +Console.WriteLine($"Best SIMD: {caps.GetBestSimdSet()}"); +Console.WriteLine($"Has AVX2: {caps.HasAVX2}"); +Console.WriteLine($"Has NEON: {caps.HasNeon}"); +Console.WriteLine($"Processor Count: {caps.ProcessorCount}"); ``` -### Optimization Levels - -#### None -No optimizations applied. Use for debugging. - -#### Basic -- Constant Folding -- Dead Code Elimination - -**Use when**: Fast compilation is critical, minimal speedup needed - -#### Standard (Recommended) -- All Basic optimizations -- Operator Fusion -- Algebraic Simplification - -**Use when**: Balanced performance and compilation time -**Expected Speedup**: 2-3x - -#### Aggressive -- All Standard optimizations -- Common Subexpression Elimination -- Strength Reduction -- In-Place Optimization -- Memory Reuse - -**Use when**: Production deployments, inference is performance-critical -**Expected Speedup**: 3-4x +## Custom Operators -#### Maximum -- All optimizations enabled -- Layout Optimization -- Maximum fusion opportunities +Register your own optimized operators: -**Use when**: Critical inference paths, compilation time not important -**Expected Speedup**: 4-5x +```csharp +public class MyCustomKernel : ICustomOperator +{ + public string Name => "MyOperation"; + public string Version => "1.0.0"; + public int Priority => 100; -## Optimization Passes + public bool IsSupported() + { + return PlatformDetector.Capabilities.HasAVX2; + } -### Operator Fusion Passes + public double EstimatedSpeedup() + { + return 3.0; // Expected 3x speedup + } -1. **ConvBatchNormReLUFusionPass**: Fuses Conv → BatchNorm → ReLU -2. **ConvBatchNormFusionPass**: Fuses Conv → BatchNorm -3. **MatMulBiasActivationFusionPass**: Fuses MatMul → Bias → Activation -4. **MatMulBiasFusionPass**: Fuses MatMul → Bias into Gemm -5. **ElementwiseFusionPass**: Fuses chains of elementwise operations -6. **MultiHeadAttentionFusionPass**: Optimizes attention mechanisms + public Tensor Execute(params Tensor[] inputs) + { + // Your optimized implementation + // ... + } +} -### Graph Structure Passes +// Register the operator +CustomOperatorRegistry.Instance.Register(new MyCustomKernel()); -1. **ConstantFoldingPass**: Evaluates constant expressions -2. **DeadCodeEliminationPass**: Removes unreachable nodes -3. **CommonSubexpressionEliminationPass**: Shares common computations -4. **LayoutOptimizationPass**: Optimizes tensor layout +// Use the operator +var kernel = CustomOperatorRegistry.Instance.GetOperator("MyOperation"); +var result = kernel.Execute(input1, input2); +``` -### Memory Passes +## Performance Profiling -1. **InPlaceOptimizationPass**: Enables in-place operations -2. **MemoryReuseOptimizationPass**: Optimizes buffer allocation +Enable profiling to track performance: -### Computation Passes +```csharp +// Enable profiling +OptimizationInitializer.Initialize(enableProfiling: true); -1. **AlgebraicSimplificationPass**: Applies algebraic identities -2. **StrengthReductionPass**: Replaces expensive operations +// Operations are automatically profiled +// ... -## Performance Benchmarks +// Get report +var report = OptimizationInitializer.GetPerformanceSummary(); +Console.WriteLine(report); -### CNN Models (ResNet-50) +// Reset statistics +OptimizationInitializer.ResetStatistics(); +``` -| Optimization Level | Inference Time | Memory Usage | Compile Time | -|-------------------|----------------|--------------|--------------| -| None | 100 ms | 1000 MB | 0 ms | -| Basic | 85 ms | 950 MB | 50 ms | -| Standard | 40 ms | 800 MB | 150 ms | -| Aggressive | 30 ms | 600 MB | 300 ms | -| Maximum | 25 ms | 550 MB | 500 ms | +## CPU Optimization Utilities -**Speedup**: 4x (None → Maximum) +Use cache-aware and loop optimization utilities: -### Transformer Models (BERT-Base) +```csharp +using AiDotNet.Tensors.Engines.Optimization; -| Optimization Level | Inference Time | Memory Usage | Compile Time | -|-------------------|----------------|--------------|--------------| -| None | 200 ms | 2000 MB | 0 ms | -| Basic | 180 ms | 1900 MB | 75 ms | -| Standard | 120 ms | 1600 MB | 200 ms | -| Aggressive | 90 ms | 1300 MB | 400 ms | -| Maximum | 75 ms | 1200 MB | 600 ms | +// Determine optimal tile size +int tileSize = LoopOptimizer.DetermineOptimalTileSize(matrixSize); -**Speedup**: 2.7x (None → Maximum) +// Use tiled loops +LoopOptimizer.Tile2D(rows, cols, tileSize, (iStart, iEnd, jStart, jEnd) => +{ + // Process tile +}); -## Architecture +// Use parallel tiling +LoopOptimizer.ParallelTile2D(rows, cols, tileSize, (iStart, iEnd, jStart, jEnd) => +{ + // Process tile in parallel +}); -``` -InferenceOptimization/ -├── Core/ -│ ├── ComputationGraph.cs # Graph data structure -│ ├── ComputationNode.cs # Graph node -│ ├── GraphOptimizer.cs # Main optimization engine -│ ├── GraphBuilder.cs # Build graphs from layers -│ ├── OptimizationOptions.cs # Configuration -│ └── OptimizationLevel.cs # Optimization levels -│ -└── Passes/ - ├── IOptimizationPass.cs # Pass interface - ├── OptimizationPassBase.cs # Pass base class - │ - ├── Fusion/ - │ ├── ConvBatchNormFusionPass.cs - │ ├── ConvBatchNormReLUFusionPass.cs - │ ├── MatMulBiasFusionPass.cs - │ ├── MatMulBiasActivationFusionPass.cs - │ ├── ElementwiseFusionPass.cs - │ └── MultiHeadAttentionFusionPass.cs - │ - ├── Graph/ - │ ├── ConstantFoldingPass.cs - │ ├── DeadCodeEliminationPass.cs - │ ├── CommonSubexpressionEliminationPass.cs - │ └── LayoutOptimizationPass.cs - │ - ├── Memory/ - │ ├── InPlaceOptimizationPass.cs - │ └── MemoryReuseOptimizationPass.cs - │ - └── Computation/ - ├── AlgebraicSimplificationPass.cs - └── StrengthReductionPass.cs +// Cache-aware transpose +CacheOptimizer.TransposeBlocked(sourceArray, destArray, rows, cols); ``` -## Creating Custom Optimization Passes +## Benchmarking -```csharp -using AiDotNet.InferenceOptimization.Passes; -using AiDotNet.InferenceOptimization.Core; -using AiDotNet.Enums; +See `AiDotNetBenchmarkTests/InferenceOptimization/` for benchmark examples. -public class MyCustomFusionPass : OptimizationPassBase where T : struct -{ - public override OptimizationPassType PassType => OptimizationPassType.Custom; - public override string Name => "My Custom Fusion"; +## Future Enhancements - public override bool Apply(IComputationGraph graph) - { - bool modified = false; - - // Find pattern: LayerNorm → Attention - var candidates = FindFusionCandidates( - graph, - OperationType.LayerNormalization, - OperationType.Attention - ); - - foreach (var sequence in candidates) - { - // Create fused node - var fusedNode = FuseNodes( - graph, - sequence, - OperationType.FusedLayerNormAttention - ); - - modified = true; - } - - return modified; - } +- GPU kernel implementations using ILGPU or ManagedCuda +- Quantization support (INT8, FP16) +- Model graph optimization +- Operator fusion +- Dynamic batching optimization +- Memory pooling - public override bool CanApply(IComputationGraph graph) - { - return graph.Nodes.Any(n => n.OperationType == OperationType.LayerNormalization); - } -} +## Integration with Existing Codebase -// Use the custom pass -var optimizer = new GraphOptimizer(); -optimizer.AddPass(new MyCustomFusionPass()); -var optimizedGraph = optimizer.Optimize(graph); -``` +The optimization module integrates with existing AiDotNet components: -## Integration with Existing Models +- **Tensor Operations**: Optimized kernels work with `AiDotNet.LinearAlgebra.Tensor` +- **Neural Networks**: Can be used to accelerate layer operations in `NeuralNetworkBase` +- **Serving**: Integrates with `RequestBatcher` for optimized inference -The optimizer automatically integrates with existing AiDotNet models: +## Requirements -```csharp -// Your existing model -var cnn = new ConvolutionalNeuralNetwork(); -cnn.AddLayer(new ConvolutionalLayer()); -cnn.AddLayer(new BatchNormalizationLayer()); -cnn.AddLayer(new ReLUActivationLayer()); -// ... more layers - -// Optimize for inference -var optimizer = new GraphOptimizer( - OptimizationOptions.FromLevel(OptimizationLevel.Aggressive) -); - -// Build and optimize graph from the model -var graphBuilder = new GraphBuilder(); -var graph = graphBuilder.BuildFromLayers(cnn.Layers); -var optimizedGraph = optimizer.Optimize(graph); - -// The optimized graph can now be used for inference -// (Future: ExecuteOptimized method on NeuralNetworkBase) -``` +- .NET 8.0 or later +- x86/x64 or ARM64 processor +- For GPU support: CUDA-capable GPU (future implementation) -## Comparison with Other Frameworks - -| Feature | AiDotNet | TensorRT | ONNX Runtime | TorchScript | -|----------------------------------|----------|----------|--------------|-------------| -| Conv+BN+ReLU Fusion | ✓ | ✓ | ✓ | ✓ | -| MatMul+Bias+Activation Fusion | ✓ | ✓ | ✓ | ✓ | -| Constant Folding | ✓ | ✓ | ✓ | ✓ | -| Dead Code Elimination | ✓ | ✓ | ✓ | ✓ | -| Common Subexpression Elimination | ✓ | ✓ | ✓ | ✓ | -| Memory Reuse Optimization | ✓ | ✓ | ✓ | ✓ | -| In-Place Operations | ✓ | ✓ | ✓ | ✓ | -| Layout Optimization | ✓ | ✓ | ✓ | ✓ | -| Algebraic Simplification | ✓ | ✓ | ✓ | ✓ | -| Native .NET Integration | ✓ | ✗ | Partial | ✗ | +## Performance Targets -## Future Enhancements +- 2-5x speedup on critical operations (achieved through SIMD and cache optimization) +- Hardware-specific optimizations (AVX2, AVX-512, NEON) +- Graceful fallback behavior (automatic platform detection) +- Benchmarking against MKL and cuBLAS (future work) + +## Contributing -- [ ] Quantization (Int8, Float16) -- [ ] Kernel auto-tuning -- [ ] Multi-GPU support -- [ ] ONNX export with optimizations -- [ ] TensorRT backend integration -- [ ] Flash Attention implementation -- [ ] Dynamic batching optimization -- [ ] Automatic mixed precision +To add new optimizations: -## Related Issues +1. Implement `ICustomOperator` interface +2. Override `IsSupported()` to check platform compatibility +3. Implement optimized `Execute()` method +4. Register operator with `CustomOperatorRegistry` +5. Add benchmarks in `AiDotNetBenchmarkTests/` -- Issue #409: Graph Optimization and Operator Fusion (this implementation) -- Issue #280: ONNX Export -- Issue #277: Inference Optimizations +## License -## References +Same as parent AiDotNet project. -- [TensorRT Optimization Guide](https://docs.nvidia.com/deeplearning/tensorrt/) -- [ONNX Runtime Performance Tuning](https://onnxruntime.ai/docs/performance/) -- [PyTorch JIT and TorchScript](https://pytorch.org/docs/stable/jit.html) -- [TVM: End-to-End Deep Learning Compiler](https://tvm.apache.org/) diff --git a/src/LoRA/Adapters/MultiLoRAAdapter.cs b/src/LoRA/Adapters/MultiLoRAAdapter.cs index ce1ac90930..66734d1952 100644 --- a/src/LoRA/Adapters/MultiLoRAAdapter.cs +++ b/src/LoRA/Adapters/MultiLoRAAdapter.cs @@ -1,4 +1,7 @@ using AiDotNet.Interfaces; +using AiDotNet.NeuralNetworks.Layers; +using AiDotNet.Tensors.LinearAlgebra; +using System.Globalization; namespace AiDotNet.LoRA.Adapters; @@ -48,7 +51,7 @@ namespace AiDotNet.LoRA.Adapters; /// You can switch between tasks at runtime, and each task only trains its specific LoRA weights! /// /// -public class MultiLoRAAdapter : LoRAAdapterBase +public class MultiLoRAAdapter : LoRAAdapterBase, ILayerSerializationExtras { /// /// Dictionary mapping task names to their specific LoRA layers. @@ -443,12 +446,13 @@ public override Vector GetParameters() } } - // All task adapters' parameters + // All task adapters' parameters (stable ordering for deterministic serialization) // Guard against null _taskAdapters during base constructor calls if (_taskAdapters != null) { - foreach (var adapter in _taskAdapters.Values) + foreach (var taskName in _taskAdapters.Keys.OrderBy(k => k, StringComparer.Ordinal)) { + var adapter = _taskAdapters[taskName]; Vector taskParams = adapter.GetParameters(); for (int i = 0; i < taskParams.Length; i++) { @@ -485,12 +489,13 @@ public override void SetParameters(Vector parameters) _baseLayer.SetParameters(baseParams); } - // All task adapters' parameters + // All task adapters' parameters (stable ordering for deterministic serialization) // Guard against null _taskAdapters during construction or early calls if (_taskAdapters != null) { - foreach (var adapter in _taskAdapters.Values) + foreach (var taskName in _taskAdapters.Keys.OrderBy(k => k, StringComparer.Ordinal)) { + var adapter = _taskAdapters[taskName]; int taskParamCount = adapter.ParameterCount; Vector taskParams = new Vector(taskParamCount); for (int i = 0; i < taskParamCount; i++) @@ -630,8 +635,9 @@ private void UpdateParameterGradientsFromLayers() currentAdapter = _taskAdapters[_currentTask]; } - foreach (var adapter in _taskAdapters.Values) + foreach (var taskName in _taskAdapters.Keys.OrderBy(k => k, StringComparer.Ordinal)) { + var adapter = _taskAdapters[taskName]; Vector? grads = (adapter == currentAdapter && currentAdapter != null) ? adapter.GetParameterGradients() : null; @@ -643,6 +649,94 @@ private void UpdateParameterGradientsFromLayers() } } + int ILayerSerializationExtras.ExtraParameterCount => _freezeBaseLayer && _baseLayer != null ? _baseLayer.ParameterCount : 0; + + Vector ILayerSerializationExtras.GetExtraParameters() + { + if (!_freezeBaseLayer || _baseLayer == null) + { + return new Vector(0); + } + + return _baseLayer.GetParameters(); + } + + void ILayerSerializationExtras.SetExtraParameters(Vector extraParameters) + { + if (!_freezeBaseLayer || _baseLayer == null) + { + return; + } + + if (extraParameters.Length != _baseLayer.ParameterCount) + { + throw new ArgumentException( + $"Expected {_baseLayer.ParameterCount} extra parameters for frozen base layer, got {extraParameters.Length}", + nameof(extraParameters)); + } + + _baseLayer.SetParameters(extraParameters); + } + + internal override Dictionary GetMetadata() + { + var meta = new Dictionary(StringComparer.Ordinal) + { + ["FreezeBaseLayer"] = _freezeBaseLayer.ToString(CultureInfo.InvariantCulture), + ["BaseLayerTypeId"] = Uri.EscapeDataString(BuildLayerTypeIdentifier(_baseLayer)) + }; + + if (_taskAdapters != null) + { + var ordered = _taskAdapters.Keys.OrderBy(k => k, StringComparer.Ordinal).ToArray(); + meta["Tasks"] = string.Join("|", ordered.Select(Uri.EscapeDataString)); + meta["TaskRanks"] = string.Join("|", ordered.Select(t => _taskAdapters[t].Rank.ToString(CultureInfo.InvariantCulture))); + meta["TaskAlphas"] = string.Join("|", ordered.Select(t => Convert.ToDouble(_taskAdapters[t].Alpha).ToString(CultureInfo.InvariantCulture))); + } + + if (!string.IsNullOrWhiteSpace(_currentTask)) + { + meta["CurrentTask"] = Uri.EscapeDataString(_currentTask); + } + + return meta; + } + + private static string BuildLayerTypeIdentifier(ILayer layer) + { + string typeName = layer.GetType().Name; + var metadata = new Dictionary(StringComparer.Ordinal); + + if (layer is LayerBase layerBase) + { + foreach (var kvp in layerBase.GetMetadata()) + { + metadata[kvp.Key] = kvp.Value; + } + + if (layerBase.VectorActivation != null) + { + metadata["VectorActivationType"] = layerBase.VectorActivation.GetType().AssemblyQualifiedName ?? layerBase.VectorActivation.GetType().FullName ?? string.Empty; + } + else if (layerBase.ScalarActivation != null) + { + metadata["ScalarActivationType"] = layerBase.ScalarActivation.GetType().AssemblyQualifiedName ?? layerBase.ScalarActivation.GetType().FullName ?? string.Empty; + } + } + + if (metadata.Count == 0) + { + return typeName; + } + + foreach (var kvp in metadata.OrderBy(k => k.Key, StringComparer.Ordinal)) + { + typeName += $";{kvp.Key}={kvp.Value}"; + } + + return typeName; + } + /// /// Resets the internal state of all layers. /// diff --git a/src/Models/Results/PredictionModelResult.cs b/src/Models/Results/PredictionModelResult.cs index 3d8131673e..d3fabefaf0 100644 --- a/src/Models/Results/PredictionModelResult.cs +++ b/src/Models/Results/PredictionModelResult.cs @@ -10,6 +10,8 @@ using AiDotNet.Deployment.Runtime; using AiDotNet.Deployment.TensorRT; using AiDotNet.Enums; +using AiDotNet.Inference; +using AiDotNet.NeuralNetworks; using AiDotNet.Helpers; using AiDotNet.Interfaces; using AiDotNet.Interpretability; @@ -455,8 +457,26 @@ public class PredictionModelResult : IFullModel [JsonIgnore] // Don't serialize - will need to be recompiled after deserialization private Func[], Tensor[]>? JitCompiledFunction { get; set; } + + [JsonProperty] private AiDotNet.Configuration.InferenceOptimizationConfig? InferenceOptimizationConfig { get; set; } + [JsonIgnore] + private readonly object _inferenceOptimizationLock = new(); + + [JsonIgnore] + private InferenceOptimizer? _inferenceOptimizer; + + [JsonIgnore] + private NeuralNetworkBase? _inferenceOptimizedNeuralModel; + + [JsonIgnore] + private bool _inferenceOptimizationsInitialized; + + // Serving assembly uses InternalsVisibleTo; keep this internal to avoid expanding user-facing API surface. + internal AiDotNet.Configuration.InferenceOptimizationConfig? GetInferenceOptimizationConfigForServing() + => InferenceOptimizationConfig; + /// /// Gets the reasoning configuration for advanced Chain-of-Thought, Tree-of-Thoughts, and Self-Consistency reasoning. /// @@ -939,10 +959,34 @@ public TOutput Predict(TInput newData) // Use JIT-compiled function if available for 5-10x faster predictions TOutput normalizedPredictions; - if (JitCompiledFunction != null && normalizedNewData is Tensor inputTensor) + + // INFERENCE OPTIMIZATION PATH: apply configured inference optimizations for neural network models + if (InferenceOptimizationConfig != null && + Model is NeuralNetworkBase neuralModel && + normalizedNewData is Tensor inputTensor) + { + var optimizedNeuralModel = EnsureStatelessInferenceOptimizationsInitialized(neuralModel); + if (optimizedNeuralModel != null) + { + var optimizedOutput = optimizedNeuralModel.Predict(inputTensor); + if ((object)optimizedOutput is TOutput output) + { + normalizedPredictions = output; + } + else + { + // Fallback to the wrapped model if type mismatch occurs + normalizedPredictions = Model.Predict(normalizedNewData); + } + + return NormalizationInfo.Normalizer.Denormalize(normalizedPredictions, NormalizationInfo.YParams); + } + } + + if (JitCompiledFunction != null && normalizedNewData is Tensor inputTensor2) { // JIT PATH: Use compiled function for accelerated inference - var jitResult = JitCompiledFunction(new[] { inputTensor }); + var jitResult = JitCompiledFunction(new[] { inputTensor2 }); if (jitResult != null && jitResult.Length > 0 && jitResult[0] is TOutput output) { normalizedPredictions = output; @@ -962,10 +1006,390 @@ public TOutput Predict(TInput newData) return NormalizationInfo.Normalizer.Denormalize(normalizedPredictions, NormalizationInfo.YParams); } + /// + /// Begins an inference session for stateful inference features (e.g., KV-cache). + /// + /// + /// + /// Sessions are intended for serving-style workloads where you run many sequential inference steps. + /// A session can create multiple independent sequences, each maintaining its own state (like KV-cache). + /// + /// + /// For Beginners: Use a session when you are doing "token-by-token" inference. + /// + /// - Use for one-off, stateless predictions. + /// - Use when you need the model to remember prior calls in the same sequence. + /// + /// + public InferenceSession BeginInferenceSession() + { + return new InferenceSession(this, InferenceOptimizationConfig); + } + + private NeuralNetworkBase? EnsureStatelessInferenceOptimizationsInitialized(NeuralNetworkBase model) + { + if (_inferenceOptimizationsInitialized) + { + return _inferenceOptimizedNeuralModel; + } + + lock (_inferenceOptimizationLock) + { + if (_inferenceOptimizationsInitialized) + { + return _inferenceOptimizedNeuralModel; + } + + try + { + if (InferenceOptimizationConfig != null) + { + // Stateless-only optimizations for plain Predict(): avoid stateful features that can leak across calls. + var statelessConfig = CreateStatelessInferenceConfig(InferenceOptimizationConfig); + var optimizer = new InferenceOptimizer(statelessConfig); + var (optimizedModel, anyApplied) = optimizer.OptimizeForInference(model, cloneModel: true); + + _inferenceOptimizer = optimizer; + _inferenceOptimizedNeuralModel = anyApplied ? optimizedModel : null; + } + } + catch (Exception ex) + { + Console.WriteLine($"Warning: inference optimizations failed: {ex.Message}"); + _inferenceOptimizer = null; + _inferenceOptimizedNeuralModel = null; + } + finally + { + _inferenceOptimizationsInitialized = true; + } + + return _inferenceOptimizedNeuralModel; + } + } + + private static AiDotNet.Configuration.InferenceOptimizationConfig CreateStatelessInferenceConfig( + AiDotNet.Configuration.InferenceOptimizationConfig config) + { + return new AiDotNet.Configuration.InferenceOptimizationConfig + { + EnableFlashAttention = config.EnableFlashAttention, + AttentionMasking = config.AttentionMasking, + + // Disable stateful/session-centric features for plain Predict(). + EnableKVCache = false, + EnablePagedKVCache = false, + EnableBatching = false, + EnableSpeculativeDecoding = false + }; + } + + /// + /// Facade-friendly inference session that owns stateful inference internals. + /// + /// + /// + /// This type intentionally keeps inference internals behind the facade. Users create sequences via + /// and run inference via . + /// + /// + public sealed class InferenceSession : IDisposable + { + private readonly PredictionModelResult _result; + private readonly AiDotNet.Configuration.InferenceOptimizationConfig? _config; + private bool _disposed; + + internal InferenceSession( + PredictionModelResult result, + AiDotNet.Configuration.InferenceOptimizationConfig? config) + { + _result = result ?? throw new ArgumentNullException(nameof(result)); + _config = config; + } + + /// + /// Creates an independent sequence within this session. + /// + /// + /// + /// Each sequence represents an independent stream (e.g., one chat) and owns its own state. + /// + /// + public InferenceSequence CreateSequence() + { + ThrowIfDisposed(); + return new InferenceSequence(_result, _config, multiLoRATask: null); + } + + // Internal (serving/tests): allow selecting a Multi-LoRA task per sequence without expanding public API surface. + internal InferenceSequence CreateSequence(string? multiLoRATask) + { + ThrowIfDisposed(); + return new InferenceSequence(_result, _config, multiLoRATask); + } + + public void Dispose() + { + _disposed = true; + } + + private void ThrowIfDisposed() + { + if (_disposed) + { + throw new ObjectDisposedException(nameof(InferenceSession)); + } + } + } + + /// + /// Represents one independent, stateful inference sequence (e.g., one chat/generation stream). + /// + /// + /// + /// A sequence may keep internal state across calls when inference optimizations are enabled (e.g., KV-cache). + /// Call to start a new logical sequence on the same object. + /// + /// + public sealed class InferenceSequence : IDisposable + { + private readonly PredictionModelResult _result; + private readonly AiDotNet.Configuration.InferenceOptimizationConfig? _config; + private bool _disposed; + + // Session-local inference state (populated lazily when used). + private InferenceOptimizer? _sequenceOptimizer; + private NeuralNetworkBase? _sequenceOptimizedNeuralModel; + private bool _sequenceInitialized; + private readonly object _sequenceLock = new(); + + internal InferenceSequence( + PredictionModelResult result, + AiDotNet.Configuration.InferenceOptimizationConfig? config, + string? multiLoRATask) + { + _result = result ?? throw new ArgumentNullException(nameof(result)); + _config = config; + _multiLoRATask = multiLoRATask; + } + + private string? _multiLoRATask; + + public TOutput Predict(TInput newData) + { + ThrowIfDisposed(); + + if (_result.Model == null) + { + throw new InvalidOperationException("Model is not initialized."); + } + + if (_result.NormalizationInfo.Normalizer == null) + { + throw new InvalidOperationException("Normalizer is not initialized."); + } + + var (normalizedNewData, _) = _result.NormalizationInfo.Normalizer.NormalizeInput(newData); + + // Session inference: use configured inference optimizations, including stateful ones, if applicable. + if (_config != null && + _result.Model is NeuralNetworkBase neuralModel && + normalizedNewData is Tensor inputTensor) + { + var optimized = EnsureSequenceOptimizationsInitialized(neuralModel); + if (optimized != null) + { + var optimizedOutput = optimized.Predict(inputTensor); + if ((object)optimizedOutput is TOutput output) + { + return _result.NormalizationInfo.Normalizer.Denormalize(output, _result.NormalizationInfo.YParams); + } + } + } + + // Fallback: normal predict path (no JIT inside a session to keep behavior consistent). + var normalizedPredictions = _result.Model.Predict(normalizedNewData); + return _result.NormalizationInfo.Normalizer.Denormalize(normalizedPredictions, _result.NormalizationInfo.YParams); + } + + public void Reset() + { + ThrowIfDisposed(); + lock (_sequenceLock) + { + _sequenceOptimizer?.ClearCache(); + } + } + + // Internal: switch Multi-LoRA task for this sequence, resetting state to avoid cache leakage. + internal void SetMultiLoRATask(string? taskName) + { + ThrowIfDisposed(); + lock (_sequenceLock) + { + if (string.Equals(_multiLoRATask, taskName, StringComparison.Ordinal)) + return; + + _multiLoRATask = taskName; + + try + { + _sequenceOptimizer?.ClearCache(); + } + catch + { + // Best-effort. + } + + _sequenceOptimizer = null; + _sequenceOptimizedNeuralModel = null; + _sequenceInitialized = false; + } + } + + public void Dispose() + { + if (_disposed) + { + return; + } + + try + { + _sequenceOptimizer?.ClearCache(); + } + catch + { + // Best-effort cleanup; disposal must not throw. + } + + _disposed = true; + } + + // Exposed to AiDotNetTests via InternalsVisibleTo for integration verification without expanding the public API surface. + internal Dictionary GetInferenceStatistics() + { + ThrowIfDisposed(); + lock (_sequenceLock) + { + return _sequenceOptimizer?.GetStatistics() ?? new Dictionary(); + } + } + + private NeuralNetworkBase? EnsureSequenceOptimizationsInitialized(NeuralNetworkBase model) + { + if (_sequenceInitialized) + { + return _sequenceOptimizedNeuralModel; + } + + lock (_sequenceLock) + { + if (_sequenceInitialized) + { + return _sequenceOptimizedNeuralModel; + } + + try + { + if (_config != null) + { + // If Multi-LoRA is in use, isolate per-sequence task selection by cloning and selecting task + // before applying any further inference optimizations. + NeuralNetworkBase modelForSequence = model; + bool hasMultiLoRATask = !string.IsNullOrWhiteSpace(_multiLoRATask); + if (hasMultiLoRATask) + { + try + { + modelForSequence = (NeuralNetworkBase)model.Clone(); + + int appliedCount = 0; + foreach (var layer in modelForSequence.Layers) + { + if (layer is AiDotNet.LoRA.Adapters.MultiLoRAAdapter multi) + { + multi.SetCurrentTask(_multiLoRATask!); + appliedCount++; + } + } + + InferenceDiagnostics.RecordDecision( + area: "InferenceSession", + feature: "MultiLoRA", + enabled: appliedCount > 0, + reason: appliedCount > 0 ? $"Task={_multiLoRATask}" : $"NoMultiLoRAAdapters(Task={_multiLoRATask})"); + } + catch (Exception ex) + { + InferenceDiagnostics.RecordException("InferenceSession", "MultiLoRA", ex, $"Task={_multiLoRATask};FallbackToBaseModel"); + modelForSequence = model; + } + } + + // In a session, prefer causal masking defaults when user left it as Auto. + var sessionConfig = _config.AttentionMasking == AiDotNet.Configuration.AttentionMaskingMode.Auto + ? new AiDotNet.Configuration.InferenceOptimizationConfig + { + EnableFlashAttention = _config.EnableFlashAttention, + EnableKVCache = _config.EnableKVCache, + EnablePagedKVCache = _config.EnablePagedKVCache, + PagedKVCacheBlockSize = _config.PagedKVCacheBlockSize, + MaxBatchSize = _config.MaxBatchSize, + KVCacheMaxSizeMB = _config.KVCacheMaxSizeMB, + KVCachePrecision = _config.KVCachePrecision, + KVCacheQuantization = _config.KVCacheQuantization, + UseSlidingWindowKVCache = _config.UseSlidingWindowKVCache, + KVCacheWindowSize = _config.KVCacheWindowSize, + EnableBatching = _config.EnableBatching, + EnableSpeculativeDecoding = _config.EnableSpeculativeDecoding, + SpeculationPolicy = _config.SpeculationPolicy, + SpeculativeMethod = _config.SpeculativeMethod, + DraftModelType = _config.DraftModelType, + SpeculationDepth = _config.SpeculationDepth, + UseTreeSpeculation = _config.UseTreeSpeculation, + EnableWeightOnlyQuantization = _config.EnableWeightOnlyQuantization, + AttentionMasking = AiDotNet.Configuration.AttentionMaskingMode.Causal + } + : _config; + + var optimizer = new InferenceOptimizer(sessionConfig); + var (optimizedModel, anyApplied) = optimizer.OptimizeForInference(modelForSequence, cloneModel: ReferenceEquals(modelForSequence, model)); + + _sequenceOptimizer = optimizer; + // If Multi-LoRA was requested, keep the per-sequence model even when no other optimizations apply. + _sequenceOptimizedNeuralModel = anyApplied || !ReferenceEquals(modelForSequence, model) ? optimizedModel : null; + } + } + catch (Exception ex) + { + Console.WriteLine($"Warning: inference session optimizations failed: {ex.Message}"); + _sequenceOptimizer = null; + _sequenceOptimizedNeuralModel = null; + } + finally + { + _sequenceInitialized = true; + } + + return _sequenceOptimizedNeuralModel; + } + } + + private void ThrowIfDisposed() + { + if (_disposed) + { + throw new ObjectDisposedException(nameof(InferenceSequence)); + } + } + } + /// /// Gets the default loss function used by this model for gradient computation. /// /// If Model is not initialized. + [JsonIgnore] public ILossFunction DefaultLossFunction { get @@ -1680,6 +2104,7 @@ public void Deserialize(byte[] data) ModelMetaData = deserializedObject.ModelMetaData; BiasDetector = deserializedObject.BiasDetector; FairnessEvaluator = deserializedObject.FairnessEvaluator; + InferenceOptimizationConfig = deserializedObject.InferenceOptimizationConfig; // Preserve RAG components and all configuration properties RagRetriever = deserializedObject.RagRetriever; @@ -1691,6 +2116,12 @@ public void Deserialize(byte[] data) AgentConfig = deserializedObject.AgentConfig; AgentRecommendation = deserializedObject.AgentRecommendation; DeploymentConfiguration = deserializedObject.DeploymentConfiguration; + + // Reset transient runtime state (will be reinitialized lazily) + JitCompiledFunction = null; + _inferenceOptimizer = null; + _inferenceOptimizedNeuralModel = null; + _inferenceOptimizationsInitialized = false; } else { diff --git a/src/NeuralNetworks/Attention/FlashAttention.cs b/src/NeuralNetworks/Attention/FlashAttention.cs index 9974e1a31b..dbfe879636 100644 --- a/src/NeuralNetworks/Attention/FlashAttention.cs +++ b/src/NeuralNetworks/Attention/FlashAttention.cs @@ -32,7 +32,7 @@ namespace AiDotNet.NeuralNetworks.Attention; /// /// /// The numeric type for computations (typically float or double). -public static class FlashAttention +internal static class FlashAttention { private static readonly INumericOperations NumOps = MathHelper.GetNumericOperations(); @@ -43,12 +43,17 @@ public static class FlashAttention /// Key tensor of shape [batch, seqLen, headDim] or [batch, heads, seqLen, headDim]. /// Value tensor of shape [batch, seqLen, headDim] or [batch, heads, seqLen, headDim]. /// Flash Attention configuration. + /// + /// Optional offset for causal masking when represents a window into a longer KV sequence. + /// Use this for KV-cached decoding where Q is the newly appended tokens and K/V contain the full cached sequence. + /// /// Output tensor of same shape as query, and optionally attention weights if configured. public static (Tensor Output, Tensor? AttentionWeights) Forward( Tensor query, Tensor key, Tensor value, - FlashAttentionConfig? config = null) + FlashAttentionConfig? config = null, + int queryOffset = 0) { config ??= FlashAttentionConfig.Default; @@ -58,14 +63,18 @@ public static (Tensor Output, Tensor? AttentionWeights) Forward( // Determine if inputs are 3D [batch, seq, dim] or 4D [batch, heads, seq, dim] bool is4D = query.Shape.Length == 4; - if (is4D) - { - return Forward4D(query, key, value, config); - } - else + int seqLenQ = is4D ? query.Shape[2] : query.Shape[1]; + int seqLenKV = is4D ? key.Shape[2] : key.Shape[1]; + if (queryOffset < 0 || queryOffset + seqLenQ > seqLenKV) { - return Forward3D(query, key, value, config); + throw new ArgumentOutOfRangeException( + nameof(queryOffset), + $"queryOffset ({queryOffset}) must satisfy 0 <= queryOffset and queryOffset + seqLenQ ({seqLenQ}) <= seqLenKV ({seqLenKV})."); } + + return is4D + ? Forward4D(query, key, value, config, queryOffset) + : Forward3D(query, key, value, config, queryOffset); } /// @@ -75,7 +84,8 @@ private static (Tensor Output, Tensor? AttentionWeights) Forward3D( Tensor query, Tensor key, Tensor value, - FlashAttentionConfig config) + FlashAttentionConfig config, + int queryOffset) { int batchSize = query.Shape[0]; int seqLenQ = query.Shape[1]; @@ -100,7 +110,7 @@ private static (Tensor Output, Tensor? AttentionWeights) Forward3D( { FlashAttentionCore( query, key, value, output, attentionWeights, - b, 0, seqLenQ, seqLenKV, headDim, scale, config); + b, 0, seqLenQ, seqLenKV, headDim, scale, config, queryOffset); } return (output, attentionWeights); @@ -113,7 +123,8 @@ private static (Tensor Output, Tensor? AttentionWeights) Forward4D( Tensor query, Tensor key, Tensor value, - FlashAttentionConfig config) + FlashAttentionConfig config, + int queryOffset) { int batchSize = query.Shape[0]; int numHeads = query.Shape[1]; @@ -141,7 +152,7 @@ private static (Tensor Output, Tensor? AttentionWeights) Forward4D( { FlashAttentionCore4D( query, key, value, output, attentionWeights, - b, h, seqLenQ, seqLenKV, headDim, scale, config); + b, h, seqLenQ, seqLenKV, headDim, scale, config, queryOffset); } } @@ -173,7 +184,8 @@ private static void FlashAttentionCore( int seqLenKV, int headDim, T scale, - FlashAttentionConfig config) + FlashAttentionConfig config, + int queryOffset) { int blockSizeQ = Math.Min(config.BlockSizeQ, seqLenQ); int blockSizeKV = Math.Min(config.BlockSizeKV, seqLenKV); @@ -211,7 +223,7 @@ private static void FlashAttentionCore( int kvBlockSize = kvEnd - kvStart; // Apply causal mask: skip blocks that are entirely masked - if (config.UseCausalMask && kvStart > qEnd - 1) + if (config.UseCausalMask && kvStart > queryOffset + qEnd - 1) { continue; } @@ -228,7 +240,7 @@ private static void FlashAttentionCore( int kIdx = kvStart + kj; // Apply causal mask - if (config.UseCausalMask && kIdx > qIdx) + if (config.UseCausalMask && kIdx > queryOffset + qIdx) { scores[qi, kj] = negInf; continue; @@ -349,7 +361,8 @@ private static void FlashAttentionCore4D( int seqLenKV, int headDim, T scale, - FlashAttentionConfig config) + FlashAttentionConfig config, + int queryOffset) { int blockSizeQ = Math.Min(config.BlockSizeQ, seqLenQ); int blockSizeKV = Math.Min(config.BlockSizeKV, seqLenKV); @@ -381,7 +394,7 @@ private static void FlashAttentionCore4D( int kvEnd = Math.Min(kvStart + blockSizeKV, seqLenKV); int kvBlockSize = kvEnd - kvStart; - if (config.UseCausalMask && kvStart > qEnd - 1) + if (config.UseCausalMask && kvStart > queryOffset + qEnd - 1) { continue; } @@ -397,7 +410,7 @@ private static void FlashAttentionCore4D( { int kIdx = kvStart + kj; - if (config.UseCausalMask && kIdx > qIdx) + if (config.UseCausalMask && kIdx > queryOffset + qIdx) { scores[qi, kj] = negInf; continue; diff --git a/src/NeuralNetworks/Attention/FlashAttentionLayer.cs b/src/NeuralNetworks/Attention/FlashAttentionLayer.cs index 4a3c869248..fd417c724e 100644 --- a/src/NeuralNetworks/Attention/FlashAttentionLayer.cs +++ b/src/NeuralNetworks/Attention/FlashAttentionLayer.cs @@ -27,7 +27,7 @@ namespace AiDotNet.NeuralNetworks.Attention; /// /// /// The numeric type for computations (typically float or double). -public class FlashAttentionLayer : LayerBase +internal class FlashAttentionLayer : LayerBase { private readonly int _headCount; private readonly int _headDimension; @@ -519,4 +519,13 @@ public override Dictionary GetDiagnostics() /// Gets the output projection weights. /// public Matrix GetOutputWeights() => _outputWeights; + + internal override Dictionary GetMetadata() + { + return new Dictionary + { + ["HeadCount"] = _headCount.ToString(), + ["UseCausalMask"] = _config.UseCausalMask.ToString() + }; + } } diff --git a/src/NeuralNetworks/Layers/DropoutLayer.cs b/src/NeuralNetworks/Layers/DropoutLayer.cs index 96837c3b15..a9f3b2f7e3 100644 --- a/src/NeuralNetworks/Layers/DropoutLayer.cs +++ b/src/NeuralNetworks/Layers/DropoutLayer.cs @@ -554,4 +554,12 @@ public override ComputationNode ExportComputationGraph(List /// public override bool SupportsJitCompilation => true; + + internal override Dictionary GetMetadata() + { + return new Dictionary + { + ["DropoutRate"] = Convert.ToDouble(_dropoutRate).ToString(System.Globalization.CultureInfo.InvariantCulture) + }; + } } diff --git a/src/NeuralNetworks/Layers/EmbeddingLayer.cs b/src/NeuralNetworks/Layers/EmbeddingLayer.cs index c6ad2df2e3..6095758c95 100644 --- a/src/NeuralNetworks/Layers/EmbeddingLayer.cs +++ b/src/NeuralNetworks/Layers/EmbeddingLayer.cs @@ -755,4 +755,13 @@ public override Autodiff.ComputationNode ExportComputationGraph(List.EmbeddingLookup(embeddingNode, inputNode); } + + internal override Dictionary GetMetadata() + { + return new Dictionary + { + ["VocabularySize"] = _embeddingTensor.Shape[0].ToString(System.Globalization.CultureInfo.InvariantCulture), + ["EmbeddingDimension"] = _embeddingTensor.Shape[1].ToString(System.Globalization.CultureInfo.InvariantCulture) + }; + } } diff --git a/src/NeuralNetworks/Layers/GraphAttentionLayer.cs b/src/NeuralNetworks/Layers/GraphAttentionLayer.cs index 626c5c4fde..f22378a1f8 100644 --- a/src/NeuralNetworks/Layers/GraphAttentionLayer.cs +++ b/src/NeuralNetworks/Layers/GraphAttentionLayer.cs @@ -1207,4 +1207,14 @@ public override ComputationNode ExportComputationGraph(List GetMetadata() + { + return new Dictionary + { + ["NumHeads"] = _numHeads.ToString(System.Globalization.CultureInfo.InvariantCulture), + ["Alpha"] = Convert.ToDouble(_alpha).ToString(System.Globalization.CultureInfo.InvariantCulture), + ["DropoutRate"] = _dropoutRate.ToString(System.Globalization.CultureInfo.InvariantCulture) + }; + } } diff --git a/src/NeuralNetworks/Layers/ILayerSerializationExtras.cs b/src/NeuralNetworks/Layers/ILayerSerializationExtras.cs new file mode 100644 index 0000000000..9b494c3827 --- /dev/null +++ b/src/NeuralNetworks/Layers/ILayerSerializationExtras.cs @@ -0,0 +1,21 @@ +using AiDotNet.Tensors.LinearAlgebra; + +namespace AiDotNet.NeuralNetworks.Layers; + +/// +/// Provides additional, optional parameter blocks for serialization that are not part of . +/// +/// Numeric type for the layer. +/// +/// This exists to support layers where intentionally reflects trainable parameters +/// (e.g., frozen base weights in LoRA adapters) but full model serialization/cloning must still preserve non-trainable state. +/// +internal interface ILayerSerializationExtras +{ + int ExtraParameterCount { get; } + + Vector GetExtraParameters(); + + void SetExtraParameters(Vector extraParameters); +} + diff --git a/src/NeuralNetworks/Layers/LayerBase.cs b/src/NeuralNetworks/Layers/LayerBase.cs index 032f9d4531..f4718f0372 100644 --- a/src/NeuralNetworks/Layers/LayerBase.cs +++ b/src/NeuralNetworks/Layers/LayerBase.cs @@ -1381,7 +1381,9 @@ public virtual void UpdateParameters(Vector parameters) throw new ArgumentException($"Expected {ParameterCount} parameters, but got {parameters.Length}"); } - Parameters = parameters; + // Delegate to SetParameters so derived layers that manage structured weights/biases + // can correctly materialize the provided flat parameter vector. + SetParameters(parameters); } /// @@ -1651,6 +1653,19 @@ public virtual Dictionary GetDiagnostics() return diagnostics; } + /// + /// Gets layer metadata required to reliably round-trip this layer via serialization. + /// + /// + /// This is intentionally internal to avoid expanding the public API surface area. Derived layers can + /// override to provide constructor-level settings that are not inferable from shapes/parameters alone + /// (e.g., attention head count, masking mode, configuration flags). + /// + internal virtual Dictionary GetMetadata() + { + return new Dictionary(StringComparer.Ordinal); + } + /// /// Applies the layer's configured activation function to a computation graph node. /// diff --git a/src/NeuralNetworks/Layers/LayerNormalizationLayer.cs b/src/NeuralNetworks/Layers/LayerNormalizationLayer.cs index 48c918bc83..7ab097be17 100644 --- a/src/NeuralNetworks/Layers/LayerNormalizationLayer.cs +++ b/src/NeuralNetworks/Layers/LayerNormalizationLayer.cs @@ -594,4 +594,12 @@ public override bool SupportsJitCompilation return _gamma != null && _beta != null; } } + + internal override Dictionary GetMetadata() + { + return new Dictionary + { + ["Epsilon"] = Convert.ToDouble(_epsilon).ToString(System.Globalization.CultureInfo.InvariantCulture) + }; + } } diff --git a/src/NeuralNetworks/Layers/MultiHeadAttentionLayer.cs b/src/NeuralNetworks/Layers/MultiHeadAttentionLayer.cs index 38ae550d91..183ce22f18 100644 --- a/src/NeuralNetworks/Layers/MultiHeadAttentionLayer.cs +++ b/src/NeuralNetworks/Layers/MultiHeadAttentionLayer.cs @@ -477,6 +477,14 @@ public Dictionary GetAuxiliaryLossDiagnostics() }; } + internal override Dictionary GetMetadata() + { + return new Dictionary + { + ["HeadCount"] = _headCount.ToString() + }; + } + /// /// Gets diagnostic information about this component's state and behavior. /// Overrides to include auxiliary loss diagnostics. diff --git a/src/NeuralNetworks/Layers/PositionalEncodingLayer.cs b/src/NeuralNetworks/Layers/PositionalEncodingLayer.cs index c89243957f..1a24bc544c 100644 --- a/src/NeuralNetworks/Layers/PositionalEncodingLayer.cs +++ b/src/NeuralNetworks/Layers/PositionalEncodingLayer.cs @@ -429,4 +429,13 @@ public override ComputationNode ExportComputationGraph(List true; + + internal override Dictionary GetMetadata() + { + return new Dictionary + { + ["MaxSequenceLength"] = maxSequenceLength.ToString(System.Globalization.CultureInfo.InvariantCulture), + ["EmbeddingSize"] = embeddingSize.ToString(System.Globalization.CultureInfo.InvariantCulture) + }; + } } diff --git a/src/NeuralNetworks/Layers/SelfAttentionLayer.cs b/src/NeuralNetworks/Layers/SelfAttentionLayer.cs index f31be7f4d9..758271224c 100644 --- a/src/NeuralNetworks/Layers/SelfAttentionLayer.cs +++ b/src/NeuralNetworks/Layers/SelfAttentionLayer.cs @@ -1371,4 +1371,12 @@ public override bool SupportsJitCompilation _valueWeights.Shape.Length >= 2 && _valueWeights.Shape[0] > 0; } } + + internal override Dictionary GetMetadata() + { + return new Dictionary + { + ["HeadCount"] = _headCount.ToString(System.Globalization.CultureInfo.InvariantCulture) + }; + } } diff --git a/src/NeuralNetworks/NeuralNetworkBase.cs b/src/NeuralNetworks/NeuralNetworkBase.cs index c7e150d8d8..9c41dd8186 100644 --- a/src/NeuralNetworks/NeuralNetworkBase.cs +++ b/src/NeuralNetworks/NeuralNetworkBase.cs @@ -1263,6 +1263,12 @@ public virtual byte[] Serialize() using var ms = new MemoryStream(); using var writer = new BinaryWriter(ms); + // Serialization format: + // - V1: [layerCount:int32] ... + // - V2+: [-version:int32][layerCount:int32] ... (supports per-layer extra parameter blocks) + const int serializationVersion = 2; + writer.Write(-serializationVersion); + // Write the number of layers writer.Write(Layers.Count); @@ -1270,7 +1276,7 @@ public virtual byte[] Serialize() foreach (var layer in Layers) { // Write layer type - writer.Write(layer.GetType().Name); + writer.Write(GetSerializedLayerTypeIdentifier(layer)); // Write input shape var inputShape = layer.GetInputShape(); @@ -1300,6 +1306,25 @@ public virtual byte[] Serialize() writer.Write(Convert.ToDouble(param)); } } + + // Write any extra parameter blocks (V2+). + int extraCount = 0; + AiDotNet.Tensors.LinearAlgebra.Vector? extras = null; + if (layer is AiDotNet.NeuralNetworks.Layers.ILayerSerializationExtras extraProvider && + extraProvider.ExtraParameterCount > 0) + { + extras = extraProvider.GetExtraParameters(); + extraCount = extras.Length; + } + + writer.Write(extraCount); + if (extraCount > 0 && extras != null) + { + for (int i = 0; i < extras.Length; i++) + { + writer.Write(Convert.ToDouble(extras[i])); + } + } } // Write network-specific data @@ -1308,6 +1333,44 @@ public virtual byte[] Serialize() return ms.ToArray(); } + private static string GetSerializedLayerTypeIdentifier(ILayer layer) + { + string typeName = layer.GetType().Name; + + var metadata = new Dictionary(StringComparer.Ordinal); + + // Persist activation types for LayerBase-derived layers so Clone/DeepCopy round-trips behavior. + if (layer is AiDotNet.NeuralNetworks.Layers.LayerBase layerBase) + { + foreach (var kvp in layerBase.GetMetadata()) + { + metadata[kvp.Key] = kvp.Value; + } + + if (layerBase.VectorActivation != null) + { + metadata["VectorActivationType"] = layerBase.VectorActivation.GetType().AssemblyQualifiedName ?? layerBase.VectorActivation.GetType().FullName ?? string.Empty; + } + else if (layerBase.ScalarActivation != null) + { + metadata["ScalarActivationType"] = layerBase.ScalarActivation.GetType().AssemblyQualifiedName ?? layerBase.ScalarActivation.GetType().FullName ?? string.Empty; + } + } + + if (metadata.Count == 0) + { + return typeName; + } + + // Stable ordering for deterministic serialization. + foreach (var kvp in metadata.OrderBy(k => k.Key, StringComparer.Ordinal)) + { + typeName += $";{kvp.Key}={kvp.Value}"; + } + + return typeName; + } + /// /// Deserializes the neural network from a byte array. /// @@ -1320,8 +1383,20 @@ public virtual void Deserialize(byte[] data) // Clear existing layers ClearLayers(); - // Read the number of layers - int layerCount = reader.ReadInt32(); + // Read the number of layers (support both V1 and V2+ formats). + int first = reader.ReadInt32(); + int serializationVersion; + int layerCount; + if (first < 0) + { + serializationVersion = -first; + layerCount = reader.ReadInt32(); + } + else + { + serializationVersion = 1; + layerCount = first; + } // Read and recreate each layer for (int i = 0; i < layerCount; i++) @@ -1363,6 +1438,24 @@ public virtual void Deserialize(byte[] data) layer.UpdateParameters(parameters); } + if (serializationVersion >= 2) + { + int extraCount = reader.ReadInt32(); + if (extraCount > 0) + { + var extraParams = new Vector(extraCount); + for (int j = 0; j < extraCount; j++) + { + extraParams[j] = NumOps.FromDouble(reader.ReadDouble()); + } + + if (layer is AiDotNet.NeuralNetworks.Layers.ILayerSerializationExtras extraProvider) + { + extraProvider.SetExtraParameters(extraParams); + } + } + } + // Add the layer to the network _layers.Add(layer); } diff --git a/src/NeuralNetworks/Transformer.cs b/src/NeuralNetworks/Transformer.cs index bb862a54ad..dd570025e2 100644 --- a/src/NeuralNetworks/Transformer.cs +++ b/src/NeuralNetworks/Transformer.cs @@ -660,7 +660,11 @@ protected override void DeserializeNetworkSpecificData(BinaryReader reader) T dropoutRate = NumOps.FromDouble(reader.ReadDouble()); // Read and reconstruct loss function and optimizer - _optimizer = DeserializationHelper.DeserializeInterface, Tensor>>(reader) ?? new GradientDescentOptimizer, Tensor>(this); + LossFunction = DeserializationHelper.DeserializeInterface>(reader) + ?? NeuralNetworkHelper.GetDefaultLossFunction(_transformerArchitecture.TaskType); + + _optimizer = DeserializationHelper.DeserializeInterface, Tensor>>(reader) + ?? new GradientDescentOptimizer, Tensor>(this); } /// diff --git a/src/Normalizers/NoNormalizer.cs b/src/Normalizers/NoNormalizer.cs index d2ec9c362d..fc5e492306 100644 --- a/src/Normalizers/NoNormalizer.cs +++ b/src/Normalizers/NoNormalizer.cs @@ -122,16 +122,17 @@ public override (TInput, List>) NormalizeInput(TInput var parameters = Enumerable.Repeat(new NormalizationParameters { Method = NormalizationMethod.None }, matrix.Columns).ToList(); return (data, parameters); } - else if (data is Tensor tensor && tensor.Shape.Length == 2) + else if (data is Tensor tensor) { - int columns = tensor.Shape[1]; - var parameters = Enumerable.Repeat(new NormalizationParameters { Method = NormalizationMethod.None }, columns).ToList(); + // Treat the last dimension as the "feature" dimension for parameter bookkeeping. + int featureCount = tensor.Shape.Length == 0 ? 1 : tensor.Shape[^1]; + var parameters = Enumerable.Repeat(new NormalizationParameters { Method = NormalizationMethod.None }, featureCount).ToList(); return (data, parameters); } throw new InvalidOperationException( $"Unsupported data type {typeof(TInput).Name}. " + - $"Supported types are Matrix<{typeof(T).Name}> and 2D Tensor<{typeof(T).Name}>."); + $"Supported types are Matrix<{typeof(T).Name}> and Tensor<{typeof(T).Name}>."); } /// diff --git a/src/PredictionModelBuilder.cs b/src/PredictionModelBuilder.cs index a80d27cde0..a52ba749fb 100644 --- a/src/PredictionModelBuilder.cs +++ b/src/PredictionModelBuilder.cs @@ -1090,6 +1090,7 @@ private PredictionModelResult BuildMetaLearningInternalAsync QueryProcessors = _queryProcessors, AgentConfig = _agentConfig, DeploymentConfiguration = deploymentConfig, + InferenceOptimizationConfig = _inferenceOptimizationConfig, ReasoningConfig = _reasoningConfig, KnowledgeGraph = _knowledgeGraph, GraphStore = _graphStore, diff --git a/src/Serving/ContinuousBatching/BatchScheduler.cs b/src/Serving/ContinuousBatching/BatchScheduler.cs index 313e10ccdd..3540758c55 100644 --- a/src/Serving/ContinuousBatching/BatchScheduler.cs +++ b/src/Serving/ContinuousBatching/BatchScheduler.cs @@ -108,7 +108,19 @@ public List> ScheduleNextBatch() lock (_lock) { var batch = new List>(); - int availableSlots = _config.MaxBatchSize - _runningSequences.Count; + + // Always include already-running sequences (continuous batching). + // These sequences must be processed every iteration until they complete or are preempted. + foreach (var seq in _runningSequences) + { + if (batch.Count >= _config.MaxBatchSize) + break; + + if (seq.Status is SequenceStatus.Generating or SequenceStatus.Prefilling) + batch.Add(seq); + } + + int availableSlots = _config.MaxBatchSize - batch.Count; long availableMemory = _config.MaxMemoryBytes - _usedMemoryBytes; // First, try to resume preempted sequences (FIFO order) diff --git a/src/Serving/ContinuousBatching/ContinuousBatcher.cs b/src/Serving/ContinuousBatching/ContinuousBatcher.cs index 82de7cc366..bf369b26ca 100644 --- a/src/Serving/ContinuousBatching/ContinuousBatcher.cs +++ b/src/Serving/ContinuousBatching/ContinuousBatcher.cs @@ -1,5 +1,7 @@ using System.Collections.Concurrent; using AiDotNet.Inference; +using AiDotNet.Inference.SpeculativeDecoding; +using AiDotNet.Helpers; using AiDotNet.Tensors.Helpers; namespace AiDotNet.Serving.ContinuousBatching; @@ -33,12 +35,22 @@ namespace AiDotNet.Serving.ContinuousBatching; /// /// /// The numeric type for tensor computations. -public class ContinuousBatcher : IDisposable +internal class ContinuousBatcher : IDisposable { private readonly ContinuousBatcherConfig _config; private readonly BatchScheduler _scheduler; private readonly KVCache? _kvCache; private readonly Func, Tensor>? _model; + private readonly IDraftModel? _draftModelOverride; + + private SpeculativeDecoder? _speculativeDecoder; + private readonly object _speculativeLock = new(); + private volatile bool _speculationDisabledDueToFailure; + private long _speculationDisabledUntilIteration; + + internal bool LastStepUsedSpeculation { get; private set; } + internal int LastStepSpeculationTokens { get; private set; } + internal string LastStepSpeculationReason { get; private set; } = string.Empty; private readonly ConcurrentDictionary>> _pendingResults; private readonly ConcurrentQueue> _incomingRequests; @@ -88,11 +100,13 @@ public class ContinuousBatcher : IDisposable public ContinuousBatcher( ContinuousBatcherConfig config, Func, Tensor>? model = null, - KVCache? kvCache = null) + KVCache? kvCache = null, + IDraftModel? draftModel = null) { _config = config ?? throw new ArgumentNullException(nameof(config)); _model = model; _kvCache = kvCache; + _draftModelOverride = draftModel; _scheduler = new BatchScheduler(config.SchedulerConfig); _pendingResults = new ConcurrentDictionary>>(); @@ -202,6 +216,11 @@ public int Step() if (batch.Count == 0) return 0; + bool useSpeculation = ShouldUseSpeculativeDecoding(batch, out var speculationReason); + LastStepUsedSpeculation = useSpeculation; + LastStepSpeculationTokens = 0; + LastStepSpeculationReason = speculationReason; + _totalIterations++; int tokensGenerated = 0; @@ -220,26 +239,32 @@ public int Step() { if (seq.Status == SequenceStatus.Generating) { - int newToken = RunDecodeStep(seq); - if (newToken >= 0) + var newTokens = useSpeculation ? RunDecodeStepSpeculative(seq) : RunDecodeStep(seq); + if (newTokens.Count > 0) { - tokensGenerated++; - _totalTokensGenerated++; - - // Fire token generated event - TokenGenerated?.Invoke(this, new TokenGeneratedEventArgs + foreach (var newToken in newTokens) { - Sequence = seq, - TokenId = newToken - }); - - // Invoke callback if provided - seq.Request.OnTokenGenerated?.Invoke(newToken); - - // Check for completion - if (seq.ShouldStop(_config.EosTokenId, seq.Request.StopTokenIds)) - { - CompleteSequence(seq); + tokensGenerated++; + _totalTokensGenerated++; + if (useSpeculation) + LastStepSpeculationTokens++; + + // Fire token generated event + TokenGenerated?.Invoke(this, new TokenGeneratedEventArgs + { + Sequence = seq, + TokenId = newToken + }); + + // Invoke callback if provided + seq.Request.OnTokenGenerated?.Invoke(newToken); + + // Check for completion after each appended token + if (seq.ShouldStop(_config.EosTokenId, seq.Request.StopTokenIds)) + { + CompleteSequence(seq); + break; + } } } } @@ -325,9 +350,9 @@ private void RunPrefill(SequenceState sequence) sequence.Status = SequenceStatus.Generating; } - private int RunDecodeStep(SequenceState sequence) + private IReadOnlyList RunDecodeStep(SequenceState sequence) { - if (_model == null) return -1; + if (_model == null) return Array.Empty(); // Create input tensor from last token only (incremental decoding) int lastToken = sequence.TokenIds[^1]; @@ -340,7 +365,292 @@ private int RunDecodeStep(SequenceState sequence) int nextToken = SampleFromLogits(logits, sequence.Request); sequence.AppendToken(nextToken); - return nextToken; + return new[] { nextToken }; + } + + private IReadOnlyList RunDecodeStepSpeculative(SequenceState sequence) + { + if (_model == null) return Array.Empty(); + if (!ShouldSpeculateForThisIteration()) return RunDecodeStep(sequence); + + int remaining = sequence.MaxNewTokens - sequence.GeneratedLength; + if (remaining <= 0) return Array.Empty(); + + var decoder = EnsureSpeculativeDecoder(); + if (decoder == null) return RunDecodeStep(sequence); + + var numOps = MathHelper.GetNumericOperations(); + T temperature = numOps.FromDouble(sequence.Request.Temperature); + + var inputTokens = new Vector(sequence.TokenIds.ToArray()); + int maxNew = Math.Min(remaining, Math.Max(1, _config.SpeculationDepth + 1)); + + SpeculativeResult result; + try + { + result = decoder.Generate( + inputTokens, + maxNewTokens: maxNew, + temperature: temperature, + eosToken: _config.EosTokenId); + } + catch (Exception ex) + { + _speculationDisabledDueToFailure = true; + InferenceDiagnostics.RecordException( + area: "Serving.ContinuousBatching", + feature: "SpeculativeDecoding", + ex: ex, + reason: "Speculative decoder execution failed; falling back to baseline decode."); + InferenceDiagnostics.RecordDecision( + area: "Serving.ContinuousBatching", + feature: "SpeculativeDecoding", + enabled: false, + reason: "DisabledDueToFailure"); + return RunDecodeStep(sequence); + } + + if (result.NewTokens.Length == 0) + return Array.Empty(); + + var tokens = new List(result.NewTokens.Length); + for (int i = 0; i < result.NewTokens.Length; i++) + { + int token = result.NewTokens[i]; + sequence.AppendToken(token); + tokens.Add(token); + + // Prevent appending beyond stop conditions (e.g., EOS in the speculative batch). + if (sequence.ShouldStop(_config.EosTokenId, sequence.Request.StopTokenIds)) + { + break; + } + } + + return tokens; + } + + private bool ShouldUseSpeculativeDecoding(IReadOnlyCollection> batch, out string reason) + { + if (_speculationDisabledDueToFailure) + { + reason = "DisabledDueToFailure"; + InferenceDiagnostics.RecordDecision("Serving.ContinuousBatching", "SpeculativeDecoding", enabled: false, reason: reason); + return false; + } + + if (!_config.EnableSpeculativeDecoding) + { + reason = "DisabledByConfig"; + InferenceDiagnostics.RecordDecision("Serving.ContinuousBatching", "SpeculativeDecoding", enabled: false, reason: reason); + return false; + } + + if (_config.SpeculationPolicy == AiDotNet.Configuration.SpeculationPolicy.ForceOff) + { + reason = "ForceOff"; + InferenceDiagnostics.RecordDecision("Serving.ContinuousBatching", "SpeculativeDecoding", enabled: false, reason: reason); + return false; + } + + if (_config.SpeculationPolicy == AiDotNet.Configuration.SpeculationPolicy.ForceOn) + { + reason = "ForceOn"; + InferenceDiagnostics.RecordDecision("Serving.ContinuousBatching", "SpeculativeDecoding", enabled: true, reason: reason); + return true; + } + + if (_config.SpeculationPolicy == AiDotNet.Configuration.SpeculationPolicy.ThroughputFirst) + { + // Extremely conservative: only speculate when there is no queue pressure and batches are tiny. + bool ok = batch.Count == 1 && _scheduler.WaitingCount == 0 && _speculationDisabledUntilIteration <= _totalIterations; + reason = ok ? "ThroughputFirst(Enabled)" : "ThroughputFirst(Backoff)"; + InferenceDiagnostics.RecordDecision("Serving.ContinuousBatching", "SpeculativeDecoding", enabled: ok, reason: reason); + return ok; + } + + // Auto policy: back off under load and when draft acceptance is too low. + if (_speculationDisabledUntilIteration > _totalIterations) + { + reason = "AutoBackoff(Cooldown)"; + InferenceDiagnostics.RecordDecision("Serving.ContinuousBatching", "SpeculativeDecoding", enabled: false, reason: reason); + return false; + } + + int maxBatchForSpeculation = _config.SchedulerConfig.MaxBatchSize / 2; + if (_config.SpeculationPolicy == AiDotNet.Configuration.SpeculationPolicy.LatencyFirst) + { + // Allow more speculation under load, but still avoid it when the queue is growing. + maxBatchForSpeculation = Math.Max(1, _config.SchedulerConfig.MaxBatchSize); + } + + bool enabled = batch.Count <= Math.Max(1, maxBatchForSpeculation) && _scheduler.WaitingCount == 0; + if (!enabled) + { + reason = _config.SpeculationPolicy == AiDotNet.Configuration.SpeculationPolicy.LatencyFirst + ? "LatencyFirst(Backoff:LoadOrQueue)" + : "AutoBackoff(LoadOrQueue)"; + InferenceDiagnostics.RecordDecision("Serving.ContinuousBatching", "SpeculativeDecoding", enabled: false, reason: reason); + return false; + } + + // If we have enough evidence that the draft model is low-quality, disable speculation for a short cooldown. + var decoder = _speculativeDecoder; + if (decoder != null && decoder.TotalDraftTokens >= 32 && decoder.AcceptanceRate < 0.25) + { + _speculationDisabledUntilIteration = _totalIterations + 25; + reason = $"AutoBackoff(LowAcceptanceRate={decoder.AcceptanceRate:0.00})"; + InferenceDiagnostics.RecordDecision("Serving.ContinuousBatching", "SpeculativeDecoding", enabled: false, reason: reason); + return false; + } + + reason = "AutoEnabled"; + InferenceDiagnostics.RecordDecision("Serving.ContinuousBatching", "SpeculativeDecoding", enabled: true, reason: reason); + return true; + } + + private bool ShouldSpeculateForThisIteration() + { + // Defensive: if speculation is enabled but we don't have a model forward, we can't speculate. + return !_speculationDisabledDueToFailure && + _model != null && + _config.EnableSpeculativeDecoding && + _config.SpeculationPolicy != AiDotNet.Configuration.SpeculationPolicy.ForceOff; + } + + private SpeculativeDecoder? EnsureSpeculativeDecoder() + { + if (_speculationDisabledDueToFailure) + return null; + + if (_speculativeDecoder != null) + return _speculativeDecoder; + + lock (_speculativeLock) + { + if (_speculativeDecoder != null) + return _speculativeDecoder; + + if (_speculationDisabledDueToFailure) + return null; + + if (_model == null) + return null; + + int vocabSize; + try + { + vocabSize = DetectVocabSize(); + } + catch (Exception ex) + { + _speculationDisabledDueToFailure = true; + InferenceDiagnostics.RecordException("Serving.ContinuousBatching", "SpeculativeDecoding", ex, "Vocab size detection failed; disabling speculation."); + return null; + } + + if (vocabSize <= 0) + { + _speculationDisabledDueToFailure = true; + InferenceDiagnostics.RecordDecision("Serving.ContinuousBatching", "SpeculativeDecoding", enabled: false, reason: "DisabledDueToFailure(VocabSizeInvalid)"); + return null; + } + + IDraftModel draft; + try + { + draft = _draftModelOverride ?? new NGramDraftModel(ngramSize: 3, vocabSize: vocabSize, seed: 42); + } + catch (Exception ex) + { + _speculationDisabledDueToFailure = true; + InferenceDiagnostics.RecordException("Serving.ContinuousBatching", "SpeculativeDecoding", ex, "Draft model init failed; disabling speculation."); + return null; + } + + Matrix TargetForward(Vector tokens) + { + // Run the target model over the full sequence and return per-position probabilities. + var input = CreateInputTensor(tokens.ToArray()); + var logits = _model(input); + + int seqLen = logits.Shape.Length > 2 ? logits.Shape[^2] : 1; + int localVocabSize = logits.Shape[^1]; + + var numOps = MathHelper.GetNumericOperations(); + var probs = new Matrix(seqLen, localVocabSize); + for (int pos = 0; pos < seqLen; pos++) + { + // Extract logits for this position + var row = new double[localVocabSize]; + double max = double.NegativeInfinity; + for (int v = 0; v < localVocabSize; v++) + { + double val = Convert.ToDouble(logits[logits.Shape.Length > 2 ? new[] { 0, pos, v } : new[] { 0, v }]); + row[v] = val; + if (val > max) max = val; + } + + // Softmax + double sum = 0.0; + for (int v = 0; v < localVocabSize; v++) + { + row[v] = Math.Exp(row[v] - max); + sum += row[v]; + } + if (sum <= 0) sum = 1; + + for (int v = 0; v < localVocabSize; v++) + { + probs[pos, v] = numOps.FromDouble(row[v] / sum); + } + } + + return probs; + } + + var config = new SpeculativeDecodingConfig + { + NumDraftTokens = Math.Max(1, _config.SpeculationDepth), + Seed = 42, + AdaptiveDraftLength = _config.SpeculationPolicy == AiDotNet.Configuration.SpeculationPolicy.Auto, + MinAcceptanceRate = MathHelper.GetNumericOperations().FromDouble(0.5), + UseTreeSpeculation = _config.UseTreeSpeculation || + _config.SpeculativeMethod == AiDotNet.Configuration.SpeculativeMethod.Medusa || + _config.SpeculativeMethod == AiDotNet.Configuration.SpeculativeMethod.Eagle, + TreeBranchFactor = _config.SpeculativeMethod == AiDotNet.Configuration.SpeculativeMethod.Medusa ? 4 : 2, + MaxTreeDepth = Math.Max(1, _config.SpeculationDepth) + }; + + try + { + _speculativeDecoder = new SpeculativeDecoder(draft, TargetForward, config); + InferenceDiagnostics.RecordDecision("Serving.ContinuousBatching", "SpeculativeDecoding", enabled: true, reason: "DecoderInitialized"); + return _speculativeDecoder; + } + catch (Exception ex) + { + _speculationDisabledDueToFailure = true; + InferenceDiagnostics.RecordException("Serving.ContinuousBatching", "SpeculativeDecoding", ex, "Decoder init failed; disabling speculation."); + return null; + } + } + } + + private int DetectVocabSize() + { + try + { + // Probe the model with a minimal input to infer the vocabulary dimension. + var probe = CreateInputTensor([0]); + var logits = _model!(probe); + return logits.Shape.Length >= 1 ? logits.Shape[^1] : 0; + } + catch + { + // Let the caller handle vocab detection failure. + return 0; + } } private Tensor CreateInputTensor(int[] tokenIds) diff --git a/src/Serving/ContinuousBatching/ContinuousBatcherConfig.cs b/src/Serving/ContinuousBatching/ContinuousBatcherConfig.cs index acaffdccea..fb91ecf939 100644 --- a/src/Serving/ContinuousBatching/ContinuousBatcherConfig.cs +++ b/src/Serving/ContinuousBatching/ContinuousBatcherConfig.cs @@ -35,6 +35,34 @@ public class ContinuousBatcherConfig /// public bool EnableSpeculativeDecoding { get; set; } = false; + /// + /// Policy for when speculative decoding should run (default: Auto). + /// + public AiDotNet.Configuration.SpeculationPolicy SpeculationPolicy { get; set; } = AiDotNet.Configuration.SpeculationPolicy.Auto; + + /// + /// Number of tokens to draft ahead when speculative decoding is enabled. + /// + public int SpeculationDepth { get; set; } = 4; + + /// + /// Speculative decoding method to use (default: Auto). + /// + /// + /// This keeps the public serving surface compact while enabling internal selection of + /// classic draft-model speculation vs tree-based alternatives (Medusa/EAGLE). + /// + public AiDotNet.Configuration.SpeculativeMethod SpeculativeMethod { get; set; } = AiDotNet.Configuration.SpeculativeMethod.Auto; + + /// + /// Whether to use tree-based speculation (multiple draft continuations). + /// + /// + /// This is an advanced option; when false the batcher uses classic speculative decoding. + /// Some speculative methods may implicitly enable this internally. + /// + public bool UseTreeSpeculation { get; set; } = false; + /// /// Creates config for a specific model. /// diff --git a/tests/AiDotNet.Serving.Tests/ServingIntegrationTests.cs b/tests/AiDotNet.Serving.Tests/ServingIntegrationTests.cs index 96d9182e4e..3bf5280ede 100644 --- a/tests/AiDotNet.Serving.Tests/ServingIntegrationTests.cs +++ b/tests/AiDotNet.Serving.Tests/ServingIntegrationTests.cs @@ -225,6 +225,69 @@ public async Task Predict_WithValidInput_ReturnsResults() repository.UnloadModel("test-model-4"); } + /// + /// Verifies that serving can route to a pre-loaded model variant via an adapter header (Multi-LoRA MVP). + /// + [Fact] + public async Task Predict_WithAdapterHeader_RoutesToModelVariant() + { + // Arrange + using var scope = _factory.Services.CreateScope(); + var repository = scope.ServiceProvider.GetRequiredService(); + + var baseName = "test-model-variant"; + var adapterId = "adapterA"; + var variantName = $"{baseName}__{adapterId}"; + + repository.LoadModel(baseName, CreateSimpleTestModel(baseName)); + + // Variant model returns (sum + 100) so we can detect routing. + var numOps = MathHelper.GetNumericOperations(); + var variant = new ServableModelWrapper( + modelName: variantName, + inputDimension: 3, + outputDimension: 1, + predictFunc: input => + { + var sum = numOps.Zero; + for (int i = 0; i < input.Length; i++) + { + sum = numOps.Add(sum, input[i]); + } + return new Vector(new[] { sum + 100.0 }); + }); + repository.LoadModel(variantName, variant); + + var request = new PredictionRequest + { + Features = new[] { new[] { 1.0, 2.0, 3.0 } }, + RequestId = "test-request-variant" + }; + + // Act + var message = new HttpRequestMessage(HttpMethod.Post, $"/api/inference/predict/{baseName}") + { + Content = JsonContent.Create(request) + }; + message.Headers.Add("X-AiDotNet-Lora", adapterId); + + var response = await _client.SendAsync(message); + + // Assert + response.EnsureSuccessStatusCode(); + var result = await response.Content.ReadFromJsonAsync(); + Assert.NotNull(result); + Assert.Equal("test-request-variant", result.RequestId); + Assert.NotNull(result.Predictions); + Assert.Single(result.Predictions); + Assert.Single(result.Predictions[0]); + Assert.Equal(106.0, result.Predictions[0][0], 5); + + // Cleanup + repository.UnloadModel(baseName); + repository.UnloadModel(variantName); + } + /// /// Critical test: Verifies that batch processing works correctly. /// This test ensures that multiple concurrent requests are batched together diff --git a/tests/AiDotNet.Tensors.Tests/AiDotNet.Tensors.Tests.csproj b/tests/AiDotNet.Tensors.Tests/AiDotNet.Tensors.Tests.csproj index 6b29872e95..118f5e070f 100644 --- a/tests/AiDotNet.Tensors.Tests/AiDotNet.Tensors.Tests.csproj +++ b/tests/AiDotNet.Tensors.Tests/AiDotNet.Tensors.Tests.csproj @@ -1,6 +1,6 @@ - net8.0;net471;net462 + net8.0;net471 enable enable false diff --git a/tests/AiDotNet.Tensors.Tests/Engines/Optimization/CacheOptimizerTests.cs b/tests/AiDotNet.Tensors.Tests/Engines/Optimization/CacheOptimizerTests.cs new file mode 100644 index 0000000000..bb6083afda --- /dev/null +++ b/tests/AiDotNet.Tensors.Tests/Engines/Optimization/CacheOptimizerTests.cs @@ -0,0 +1,68 @@ +using System; +using AiDotNet.Tensors.Engines.Optimization; +using Xunit; + +namespace AiDotNet.Tensors.Tests.Engines.Optimization; + +public class CacheOptimizerTests +{ + [Fact] + public void ComputeOptimalTiling_ReturnsPositiveTiles_WithinDimensions() + { + var (tileM, tileN, tileK) = CacheOptimizer.ComputeOptimalTiling(m: 128, n: 256, k: 64); + + Assert.InRange(tileM, 1, 128); + Assert.InRange(tileN, 1, 256); + Assert.InRange(tileK, 1, 64); + } + + [Fact] + public void TransposeBlocked_TransposesCorrectly() + { + const int rows = 3; + const int cols = 4; + + var src = new float[rows * cols]; + for (int i = 0; i < src.Length; i++) + { + src[i] = i + 1; + } + + var dst = new float[rows * cols]; + + CacheOptimizer.TransposeBlocked(src, dst, rows, cols); + + for (int r = 0; r < rows; r++) + { + for (int c = 0; c < cols; c++) + { + Assert.Equal(src[r * cols + c], dst[c * rows + r]); + } + } + } + + [Fact] + public void CopyWithPrefetch_CopiesAllElements() + { + var src = new float[] { 1f, 2f, 3f, 4f }; + var dst = new float[src.Length]; + + CacheOptimizer.CopyWithPrefetch(src, dst, src.Length); + + Assert.Equal(src, dst); + } + + [Fact] + public void MortonEncodeDecode_RoundTrips() + { + const int x = 123; + const int y = 456; + + int code = CacheOptimizer.MortonEncode(x, y); + var (rx, ry) = CacheOptimizer.MortonDecode(code); + + Assert.Equal(x & 0x0000ffff, rx); + Assert.Equal(y & 0x0000ffff, ry); + } +} + diff --git a/tests/AiDotNet.Tensors.Tests/Engines/Optimization/LoopOptimizerTests.cs b/tests/AiDotNet.Tensors.Tests/Engines/Optimization/LoopOptimizerTests.cs new file mode 100644 index 0000000000..909d7a7ee5 --- /dev/null +++ b/tests/AiDotNet.Tensors.Tests/Engines/Optimization/LoopOptimizerTests.cs @@ -0,0 +1,121 @@ +using System; +using System.Collections.Concurrent; +using AiDotNet.Tensors.Engines.Optimization; +using Xunit; + +namespace AiDotNet.Tensors.Tests.Engines.Optimization; + +public class LoopOptimizerTests +{ + [Fact] + public void Tile2D_VisitsAllTiles() + { + int rows = 10; + int cols = 9; + int tileSize = 4; + int count = 0; + + LoopOptimizer.Tile2D(rows, cols, tileSize, (iStart, iEnd, jStart, jEnd) => + { + Assert.InRange(iStart, 0, rows - 1); + Assert.InRange(iEnd, 1, rows); + Assert.InRange(jStart, 0, cols - 1); + Assert.InRange(jEnd, 1, cols); + count++; + }); + + int expectedTilesI = (rows + tileSize - 1) / tileSize; + int expectedTilesJ = (cols + tileSize - 1) / tileSize; + Assert.Equal(expectedTilesI * expectedTilesJ, count); + } + + [Fact] + public void UnrollBy4_InvokesActionForAllIndices() + { + const int length = 17; + int seen = 0; + + LoopOptimizer.UnrollBy4(length, _ => seen++); + + Assert.Equal(length, seen); + } + + [Fact] + public void UnrollBy8_InvokesActionForAllIndices() + { + const int length = 17; + int seen = 0; + + LoopOptimizer.UnrollBy8(length, _ => seen++); + + Assert.Equal(length, seen); + } + + [Fact] + public void StripMine_CoversFullRange() + { + const int total = 10; + const int strip = 4; + int covered = 0; + + LoopOptimizer.StripMine(total, strip, (start, end) => covered += (end - start)); + + Assert.Equal(total, covered); + } + + [Fact] + public void Fuse_RunsAllActionsEachIteration() + { + const int length = 5; + int a = 0; + int b = 0; + + LoopOptimizer.Fuse(length, _ => a++, _ => b++); + + Assert.Equal(length, a); + Assert.Equal(length, b); + } + + [Fact] + public void OptimalOrder2D_RowMajorAndColumnMajorVisitAll() + { + const int rows = 3; + const int cols = 4; + + int rowMajor = 0; + LoopOptimizer.OptimalOrder2D(rows, cols, rowMajorAccess: true, (_, _) => rowMajor++); + Assert.Equal(rows * cols, rowMajor); + + int colMajor = 0; + LoopOptimizer.OptimalOrder2D(rows, cols, rowMajorAccess: false, (_, _) => colMajor++); + Assert.Equal(rows * cols, colMajor); + } + + [Fact] + public void ParallelTile2D_VisitsAllTiles() + { + int rows = 9; + int cols = 9; + int tileSize = 4; + + var tiles = new ConcurrentBag<(int, int, int, int)>(); + + LoopOptimizer.ParallelTile2D(rows, cols, tileSize, (iStart, iEnd, jStart, jEnd) => + { + tiles.Add((iStart, iEnd, jStart, jEnd)); + }); + + int expectedTilesI = (rows + tileSize - 1) / tileSize; + int expectedTilesJ = (cols + tileSize - 1) / tileSize; + Assert.Equal(expectedTilesI * expectedTilesJ, tiles.Count); + } + + [Fact] + public void DetermineOptimalTileSize_ReturnsAtLeastOne_AndAtMostDimension() + { + int tileSize = LoopOptimizer.DetermineOptimalTileSize(dimension: 128); + + Assert.InRange(tileSize, 1, 128); + } +} + diff --git a/tests/AiDotNet.Tensors.Tests/Engines/Optimization/PerformanceProfilerTests.cs b/tests/AiDotNet.Tensors.Tests/Engines/Optimization/PerformanceProfilerTests.cs new file mode 100644 index 0000000000..3df0b0f81d --- /dev/null +++ b/tests/AiDotNet.Tensors.Tests/Engines/Optimization/PerformanceProfilerTests.cs @@ -0,0 +1,63 @@ +using AiDotNet.Tensors.Engines.Optimization; +using Xunit; + +namespace AiDotNet.Tensors.Tests.Engines.Optimization; + +public class PerformanceProfilerTests +{ + [Fact] + public void Profile_WhenEnabled_RecordsStats() + { + var profiler = PerformanceProfiler.Instance; + string operationName = $"op-{Guid.NewGuid():N}"; + bool wasEnabled = profiler.Enabled; + + profiler.Clear(); + profiler.Enabled = true; + + try + { + using (profiler.Profile(operationName)) + { + _ = 1 + 1; + } + + var stats = profiler.GetStats(operationName); + Assert.NotNull(stats); + Assert.Equal(operationName, stats!.OperationName); + Assert.Equal(1, stats.CallCount); + Assert.True(stats.TotalTicks > 0); + } + finally + { + profiler.Enabled = wasEnabled; + profiler.Clear(); + } + } + + [Fact] + public void Profile_WhenDisabled_ReturnsEmptyDisposable() + { + var profiler = PerformanceProfiler.Instance; + string operationName = $"op-disabled-{Guid.NewGuid():N}"; + bool wasEnabled = profiler.Enabled; + + profiler.Clear(); + profiler.Enabled = false; + + try + { + using (profiler.Profile(operationName)) + { + _ = 1 + 1; + } + + Assert.Null(profiler.GetStats(operationName)); + } + finally + { + profiler.Enabled = wasEnabled; + profiler.Clear(); + } + } +} diff --git a/tests/AiDotNet.Tensors.Tests/Engines/PlatformDetectorTests.cs b/tests/AiDotNet.Tensors.Tests/Engines/PlatformDetectorTests.cs new file mode 100644 index 0000000000..14f8500b91 --- /dev/null +++ b/tests/AiDotNet.Tensors.Tests/Engines/PlatformDetectorTests.cs @@ -0,0 +1,21 @@ +using AiDotNet.Tensors.Engines; +using Xunit; + +namespace AiDotNet.Tensors.Tests.Engines; + +public class PlatformDetectorTests +{ + [Fact] + public void Capabilities_IsPopulated_AndDoesNotThrow() + { + var caps = PlatformDetector.Capabilities; + + Assert.True(caps.ProcessorCount > 0); + Assert.NotNull(caps.OSDescription); + Assert.NotNull(caps.FrameworkDescription); + Assert.True(caps.L1CacheSize > 0); + Assert.True(caps.L2CacheSize > 0); + Assert.True(caps.L3CacheSize > 0); + } +} + diff --git a/tests/AiDotNet.Tensors.Tests/Engines/Simd/SimdKernelsTests.cs b/tests/AiDotNet.Tensors.Tests/Engines/Simd/SimdKernelsTests.cs new file mode 100644 index 0000000000..554c16a38f --- /dev/null +++ b/tests/AiDotNet.Tensors.Tests/Engines/Simd/SimdKernelsTests.cs @@ -0,0 +1,100 @@ +using System; +using AiDotNet.Tensors.Engines.Simd; +using Xunit; + +namespace AiDotNet.Tensors.Tests.Engines.Simd; + +public class SimdKernelsTests +{ + [Fact] + public void VectorAdd_MatchesScalar() + { + var a = new float[] { 1, 2, 3, 4, 5 }; + var b = new float[] { 10, 20, 30, 40, 50 }; + var result = new float[a.Length]; + + SimdKernels.VectorAdd(a, b, result); + + for (int i = 0; i < a.Length; i++) + { + Assert.Equal(a[i] + b[i], result[i]); + } + } + + [Fact] + public void VectorMultiply_MatchesScalar() + { + var a = new float[] { 1, 2, 3, 4, 5 }; + var b = new float[] { 10, 20, 30, 40, 50 }; + var result = new float[a.Length]; + + SimdKernels.VectorMultiply(a, b, result); + + for (int i = 0; i < a.Length; i++) + { + Assert.Equal(a[i] * b[i], result[i]); + } + } + + [Fact] + public void DotProduct_MatchesScalar() + { + var a = new float[] { 1, 2, 3, 4 }; + var b = new float[] { 10, 20, 30, 40 }; + + float dot = SimdKernels.DotProduct(a, b); + + Assert.Equal(1 * 10 + 2 * 20 + 3 * 30 + 4 * 40, dot); + } + + [Fact] + public void ScalarMultiplyAdd_MatchesScalar() + { + var a = new float[] { 1, 2, 3, 4 }; + var b = new float[] { 10, 20, 30, 40 }; + var result = new float[a.Length]; + + SimdKernels.ScalarMultiplyAdd(a, b, scalar: 0.5f, result); + + for (int i = 0; i < a.Length; i++) + { + Assert.Equal(a[i] + (b[i] * 0.5f), result[i]); + } + } + + [Fact] + public void ReLU_ZeroesNegatives() + { + var input = new float[] { -1, 0, 2, -3, 4 }; + var output = new float[input.Length]; + + SimdKernels.ReLU(input, output); + + Assert.Equal(new float[] { 0, 0, 2, 0, 4 }, output); + } + + [Fact] + public void Exp_MatchesMathExpWithinTolerance() + { + var input = new float[] { -1f, 0f, 1f }; + var output = new float[input.Length]; + + SimdKernels.Exp(input, output); + + for (int i = 0; i < input.Length; i++) + { + Assert.InRange(output[i], (float)Math.Exp(input[i]) * 0.999f, (float)Math.Exp(input[i]) * 1.001f); + } + } + + [Fact] + public void Sum_MatchesScalar() + { + var input = new float[] { 1, 2, 3, 4, 5 }; + + float sum = SimdKernels.Sum(input); + + Assert.Equal(15f, sum); + } +} + diff --git a/tests/AiDotNet.Tests/InferenceOptimization/AttentionKernelValidationTests.cs b/tests/AiDotNet.Tests/InferenceOptimization/AttentionKernelValidationTests.cs new file mode 100644 index 0000000000..20319aa445 --- /dev/null +++ b/tests/AiDotNet.Tests/InferenceOptimization/AttentionKernelValidationTests.cs @@ -0,0 +1,120 @@ +using AiDotNet.InferenceOptimization.Kernels; +using AiDotNet.LinearAlgebra; +using System; +using Xunit; + +namespace AiDotNet.Tests.InferenceOptimization; + +public class AttentionKernelValidationTests +{ + [Fact] + public void Execute_MatchesNaiveAttention() + { + var kernel = new AttentionKernel(); + + // [batch=1, seq=2, d=4] + var q = CreateTensor(new[] { 1, 2, 4 }, new float[] { 1, 0, 0, 0, 0, 1, 0, 0 }); + var k = CreateTensor(new[] { 1, 2, 4 }, new float[] { 1, 0, 0, 0, 0, 1, 0, 0 }); + var v = CreateTensor(new[] { 1, 2, 4 }, new float[] { 10, 11, 12, 13, 20, 21, 22, 23 }); + + var actual = kernel.Execute(q, k, v); + var expected = NaiveAttention(q, k, v); + + Assert.Equal(expected.Shape, actual.Shape); + for (int i = 0; i < expected.Data.Length; i++) + { + Assert.Equal(expected.Data[i], actual.Data[i], 5); + } + } + + [Fact] + public void Execute_WithMask_RespectsMaskZeros() + { + var kernel = new AttentionKernel(); + + // [batch=1, seq=2, d=2] + var q = CreateTensor(new[] { 1, 2, 2 }, new float[] { 1, 0, 0, 1 }); + var k = CreateTensor(new[] { 1, 2, 2 }, new float[] { 1, 0, 0, 1 }); + var v = CreateTensor(new[] { 1, 2, 2 }, new float[] { 1, 2, 100, 200 }); + + // Mask: allow only attending to token 0 (mask value 1), disallow token 1 (mask value 0) + var mask = CreateTensor(new[] { 1, 2, 2 }, new float[] { 1, 0, 1, 0 }); + + var actual = kernel.Execute(q, k, v, mask); + + // With token 1 masked out, both rows should match v0. + Assert.Equal(1f, actual.Data[0], 5); + Assert.Equal(2f, actual.Data[1], 5); + Assert.Equal(1f, actual.Data[2], 5); + Assert.Equal(2f, actual.Data[3], 5); + } + + private static Tensor NaiveAttention(Tensor q, Tensor k, Tensor v) + { + int seq = q.Shape[1]; + int d = q.Shape[2]; + float scale = 1f / MathF.Sqrt(d); + + var scores = new float[seq * seq]; + for (int i = 0; i < seq; i++) + { + for (int j = 0; j < seq; j++) + { + float dot = 0f; + for (int t = 0; t < d; t++) + { + dot += q.Data[i * d + t] * k.Data[j * d + t]; + } + + scores[i * seq + j] = dot * scale; + } + } + + for (int i = 0; i < seq; i++) + { + float max = float.NegativeInfinity; + for (int j = 0; j < seq; j++) + { + max = Math.Max(max, scores[i * seq + j]); + } + + float sum = 0f; + for (int j = 0; j < seq; j++) + { + scores[i * seq + j] = MathF.Exp(scores[i * seq + j] - max); + sum += scores[i * seq + j]; + } + + for (int j = 0; j < seq; j++) + { + scores[i * seq + j] /= sum; + } + } + + var result = new Tensor(new[] { 1, seq, v.Shape[2] }); + int dV = v.Shape[2]; + for (int i = 0; i < seq; i++) + { + for (int j = 0; j < dV; j++) + { + float sum = 0f; + for (int t = 0; t < seq; t++) + { + sum += scores[i * seq + t] * v.Data[t * dV + j]; + } + + result.Data[i * dV + j] = sum; + } + } + + return result; + } + + private static Tensor CreateTensor(int[] shape, float[] data) + { + var t = new Tensor(shape); + Assert.Equal(t.Data.Length, data.Length); + Array.Copy(data, t.Data, data.Length); + return t; + } +} diff --git a/tests/AiDotNet.Tests/InferenceOptimization/CacheOptimizerTests.cs b/tests/AiDotNet.Tests/InferenceOptimization/CacheOptimizerTests.cs new file mode 100644 index 0000000000..e8d456d659 --- /dev/null +++ b/tests/AiDotNet.Tests/InferenceOptimization/CacheOptimizerTests.cs @@ -0,0 +1,36 @@ +using AiDotNet.Tensors.Engines.Optimization; +using Xunit; + +namespace AiDotNet.Tests.InferenceOptimization; + +public class CacheOptimizerTests +{ + [Fact] + public void TransposeBlocked_Transposes2DMatrix() + { + // 2x3 + float[] src = new float[] { 1f, 2f, 3f, 4f, 5f, 6f }; + float[] dst = new float[src.Length]; + + CacheOptimizer.TransposeBlocked(src, dst, rows: 2, cols: 3); + + // 3x2 (row-major): [ [1,4], [2,5], [3,6] ] + float[] expected = new float[] { 1f, 4f, 2f, 5f, 3f, 6f }; + Assert.Equal(expected, dst); + } + + [Fact] + public void CopyWithPrefetch_CopiesPrefix() + { + float[] src = new float[] { 1f, 2f, 3f, 4f, 5f }; + float[] dst = new float[] { 0f, 0f, 0f, 0f, 0f }; + + CacheOptimizer.CopyWithPrefetch(src, dst, length: 3); + + Assert.Equal(1f, dst[0]); + Assert.Equal(2f, dst[1]); + Assert.Equal(3f, dst[2]); + Assert.Equal(0f, dst[3]); + Assert.Equal(0f, dst[4]); + } +} diff --git a/tests/AiDotNet.Tests/InferenceOptimization/ConvolutionKernelValidationTests.cs b/tests/AiDotNet.Tests/InferenceOptimization/ConvolutionKernelValidationTests.cs new file mode 100644 index 0000000000..ecbd2b81f5 --- /dev/null +++ b/tests/AiDotNet.Tests/InferenceOptimization/ConvolutionKernelValidationTests.cs @@ -0,0 +1,22 @@ +using System; +using AiDotNet.InferenceOptimization.Kernels; +using AiDotNet.LinearAlgebra; +using Xunit; + +namespace AiDotNet.Tests.InferenceOptimization; + +public class ConvolutionKernelValidationTests +{ + [Fact] + public void Conv2D_Throws_WhenKernelInChannelsMismatch() + { + var kernel = new ConvolutionKernel(); + + var input = new Tensor(new[] { 1, 3, 5, 5 }); + var badKernel = new Tensor(new[] { 2, 2, 3, 3 }); + + var ex = Assert.Throws(() => kernel.Conv2D(input, badKernel)); + Assert.Contains("kernel.Shape[1] == inChannels", ex.Message, StringComparison.OrdinalIgnoreCase); + } +} + diff --git a/tests/AiDotNet.Tests/InferenceOptimization/GemmKernelValidationTests.cs b/tests/AiDotNet.Tests/InferenceOptimization/GemmKernelValidationTests.cs new file mode 100644 index 0000000000..12caf6b441 --- /dev/null +++ b/tests/AiDotNet.Tests/InferenceOptimization/GemmKernelValidationTests.cs @@ -0,0 +1,100 @@ +using AiDotNet.InferenceOptimization.Kernels; +using AiDotNet.LinearAlgebra; +using Xunit; + +namespace AiDotNet.Tests.InferenceOptimization; + +public class GemmKernelValidationTests +{ + [Fact] + public void Execute_MatchesNaiveGemm() + { + var kernel = new GemmKernel(); + + // A: 2x3 + var a = CreateTensor(new[] { 2, 3 }, new float[] { 1, 2, 3, 4, 5, 6 }); + // B: 3x2 + var b = CreateTensor(new[] { 3, 2 }, new float[] { 7, 8, 9, 10, 11, 12 }); + + var actual = kernel.Execute(a, b); + var expected = NaiveGemm(a, b); + + Assert.Equal(expected.Shape, actual.Shape); + Assert.Equal(expected.Data, actual.Data); + } + + [Fact] + public void GemmTransposeB_MatchesNaive() + { + var kernel = new GemmKernel(); + + // A: 2x3 + var a = CreateTensor(new[] { 2, 3 }, new float[] { 1, 2, 3, 4, 5, 6 }); + // B: 2x3 (represents B^T; result is 2x2) + var b = CreateTensor(new[] { 2, 3 }, new float[] { 7, 8, 9, 10, 11, 12 }); + + var actual = kernel.GemmTransposeB(a, b); + var expected = NaiveGemmTransposeB(a, b); + + Assert.Equal(expected.Shape, actual.Shape); + Assert.Equal(expected.Data, actual.Data); + } + + private static Tensor NaiveGemm(Tensor a, Tensor b) + { + int m = a.Shape[0]; + int k = a.Shape[1]; + int n = b.Shape[1]; + + var c = new Tensor(new[] { m, n }); + + for (int i = 0; i < m; i++) + { + for (int j = 0; j < n; j++) + { + float sum = 0f; + for (int t = 0; t < k; t++) + { + sum += a.Data[i * k + t] * b.Data[t * n + j]; + } + + c.Data[i * n + j] = sum; + } + } + + return c; + } + + private static Tensor NaiveGemmTransposeB(Tensor a, Tensor b) + { + int m = a.Shape[0]; + int k = a.Shape[1]; + int n = b.Shape[0]; + + var c = new Tensor(new[] { m, n }); + + for (int i = 0; i < m; i++) + { + for (int j = 0; j < n; j++) + { + float sum = 0f; + for (int t = 0; t < k; t++) + { + sum += a.Data[i * k + t] * b.Data[j * k + t]; + } + + c.Data[i * n + j] = sum; + } + } + + return c; + } + + private static Tensor CreateTensor(int[] shape, float[] data) + { + var t = new Tensor(shape); + Assert.Equal(t.Data.Length, data.Length); + Array.Copy(data, t.Data, data.Length); + return t; + } +} diff --git a/tests/AiDotNet.Tests/InferenceOptimization/SimdKernelsTests.cs b/tests/AiDotNet.Tests/InferenceOptimization/SimdKernelsTests.cs new file mode 100644 index 0000000000..c86ab40cf6 --- /dev/null +++ b/tests/AiDotNet.Tests/InferenceOptimization/SimdKernelsTests.cs @@ -0,0 +1,152 @@ +using AiDotNet.Tensors.Engines.Simd; +using System; +using Xunit; + +namespace AiDotNet.Tests.InferenceOptimization; + +public class SimdKernelsTests +{ + [Fact] + public void VectorAdd_MatchesScalar() + { + var a = CreateInput(32, 1); + var b = CreateInput(32, 17); + var result = new float[a.Length]; + var expected = new float[a.Length]; + + for (int i = 0; i < a.Length; i++) + { + expected[i] = a[i] + b[i]; + } + + SimdKernels.VectorAdd(a, b, result); + + AssertEqual(expected, result); + } + + [Fact] + public void VectorMultiply_MatchesScalar() + { + var a = CreateInput(32, 3); + var b = CreateInput(32, 9); + var result = new float[a.Length]; + var expected = new float[a.Length]; + + for (int i = 0; i < a.Length; i++) + { + expected[i] = a[i] * b[i]; + } + + SimdKernels.VectorMultiply(a, b, result); + + AssertEqual(expected, result); + } + + [Fact] + public void DotProduct_MatchesScalar() + { + var a = CreateInput(37, 5); + var b = CreateInput(37, 11); + + float expected = 0f; + for (int i = 0; i < a.Length; i++) + { + expected += a[i] * b[i]; + } + + float actual = SimdKernels.DotProduct(a, b); + Assert.Equal(expected, actual, 5); + } + + [Fact] + public void ScalarMultiplyAdd_MatchesScalar() + { + var a = CreateInput(31, 7); + var b = CreateInput(31, 13); + var result = new float[a.Length]; + var expected = new float[a.Length]; + + float scalar = 0.25f; + for (int i = 0; i < a.Length; i++) + { + expected[i] = a[i] + scalar * b[i]; + } + + SimdKernels.ScalarMultiplyAdd(a, b, scalar, result); + + AssertEqual(expected, result); + } + + [Fact] + public void ReLU_MatchesScalar() + { + var input = CreateSignedInput(33); + var output = new float[input.Length]; + var expected = new float[input.Length]; + + for (int i = 0; i < input.Length; i++) + { + expected[i] = Math.Max(0f, input[i]); + } + + SimdKernels.ReLU(input, output); + + AssertEqual(expected, output); + } + + [Fact] + public void Sum_MatchesScalar() + { + var input = CreateInput(100, 23); + float expected = 0f; + for (int i = 0; i < input.Length; i++) + { + expected += input[i]; + } + + float actual = SimdKernels.Sum(input); + Assert.Equal(expected, actual, 5); + } + + private static float[] CreateInput(int length, int seed) + { + var data = new float[length]; + for (int i = 0; i < length; i++) + { + data[i] = DeterministicValue(i + seed); + } + + return data; + } + + private static float[] CreateSignedInput(int length) + { + var data = new float[length]; + for (int i = 0; i < length; i++) + { + float v = DeterministicValue(i); + data[i] = (i % 2 == 0) ? v : -v; + } + + return data; + } + + private static float DeterministicValue(int i) + { + unchecked + { + uint x = (uint)(i * 1664525 + 1013904223); + return (x & 0x00FFFFFF) / 16777216f; + } + } + + private static void AssertEqual(float[] expected, float[] actual) + { + Assert.Equal(expected.Length, actual.Length); + for (int i = 0; i < expected.Length; i++) + { + Assert.Equal(expected[i], actual[i], 5); + } + } +} + diff --git a/tests/AiDotNet.Tests/IntegrationTests/Inference/InferenceSessionIntegrationTests.cs b/tests/AiDotNet.Tests/IntegrationTests/Inference/InferenceSessionIntegrationTests.cs new file mode 100644 index 0000000000..3f09c7ff70 --- /dev/null +++ b/tests/AiDotNet.Tests/IntegrationTests/Inference/InferenceSessionIntegrationTests.cs @@ -0,0 +1,644 @@ +using AiDotNet.Configuration; +using AiDotNet.Enums; +using AiDotNet.Interfaces; +using AiDotNet.Models; +using AiDotNet.Models.Options; +using AiDotNet.Models.Results; +using AiDotNet.NeuralNetworks; +using AiDotNet.NeuralNetworks.Layers; +using AiDotNet.Normalizers; +using AiDotNet.Tensors.LinearAlgebra; +using System.Linq; +using System.Threading.Tasks; +using Xunit; + +namespace AiDotNet.Tests.IntegrationTests.Inference; + +[Collection(AiDotNet.Tests.TestInfrastructure.DiagnosticsEnvironmentCollection.Name)] +public class InferenceSessionIntegrationTests +{ + private const float Tolerance = 1e-4f; + private const int SequenceLength = 1; + private const int EmbeddingDimension = 8; + private const int HeadCount = 2; + private const int FlatSize = SequenceLength * EmbeddingDimension; + + [Fact] + public void PredictionModelResult_Predict_IsStateless_WhenInferenceOptimizationsConfigured() + { + var result = CreateDeterministicResult( + new InferenceOptimizationConfig + { + EnableFlashAttention = false, + EnableKVCache = true, + EnablePagedKVCache = false, + AttentionMasking = AttentionMaskingMode.Auto + }); + + var token = CreateTokenTensor(1.0f); + + var y1 = result.Predict(token); + var y2 = result.Predict(token); + + AssertTensorsEqual(y1, y2, Tolerance); + } + + [Fact] + public void PredictionModelResult_SerializeDeserialize_PreservesInferenceOptimizationConfig() + { + var config = new InferenceOptimizationConfig + { + EnableFlashAttention = false, + EnableKVCache = true, + EnablePagedKVCache = false, + AttentionMasking = AttentionMaskingMode.Auto + }; + + var original = CreateDeterministicResult(config); + var bytes = original.Serialize(); + + var loaded = CreateDeterministicResult( + new InferenceOptimizationConfig + { + EnableFlashAttention = true, + EnableKVCache = false, + EnablePagedKVCache = true, + AttentionMasking = AttentionMaskingMode.Causal + }); + + loaded.Deserialize(bytes); + + var loadedConfig = loaded.GetInferenceOptimizationConfigForServing(); + Assert.NotNull(loadedConfig); + Assert.Equal(config.EnableFlashAttention, loadedConfig!.EnableFlashAttention); + Assert.Equal(config.EnableKVCache, loadedConfig.EnableKVCache); + Assert.Equal(config.EnablePagedKVCache, loadedConfig.EnablePagedKVCache); + Assert.Equal(config.AttentionMasking, loadedConfig.AttentionMasking); + } + + [Fact] + public void BeginInferenceSession_SequencesAreIndependent() + { + var result = CreateDeterministicResult( + new InferenceOptimizationConfig + { + EnableFlashAttention = false, + EnableKVCache = true, + EnablePagedKVCache = false, + AttentionMasking = AttentionMaskingMode.Auto + }); + + var token = CreateTokenTensor(0.75f); + var tokenForB = CreateTokenTensor(0.75f); + var tokenFresh = CreateTokenTensor(0.75f); + + using var session = result.BeginInferenceSession(); + + var seqA = session.CreateSequence(); + var seqB = session.CreateSequence(); + var seqFresh = session.CreateSequence(); + + var a1 = seqA.Predict(token); + var statsAfterFirst = seqA.GetInferenceStatistics(); + var lengthsAfterFirst = (int[])statsAfterFirst["KVCache_SequenceLengths"]; + int lenAfterFirst = lengthsAfterFirst[0]; + + var b1 = seqB.Predict(tokenForB); + var fresh1 = seqFresh.Predict(tokenFresh); + + AssertTensorsEqual(a1, b1, Tolerance); + AssertTensorsEqual(a1, fresh1, Tolerance); + + var freshStatsAfterFirst = seqFresh.GetInferenceStatistics(); + var freshLengthsAfterFirst = (int[])freshStatsAfterFirst["KVCache_SequenceLengths"]; + int freshLenAfterFirst = freshLengthsAfterFirst[0]; + Assert.Equal(lenAfterFirst, freshLenAfterFirst); + + _ = seqA.Predict(CreateTokenTensor(-0.25f)); + + var statsAfterSecond = seqA.GetInferenceStatistics(); + var lengthsAfterSecond = (int[])statsAfterSecond["KVCache_SequenceLengths"]; + Assert.True(lengthsAfterSecond[0] > lenAfterFirst, $"Expected KV-cache length to grow, but got {lenAfterFirst} -> {lengthsAfterSecond[0]}"); + + // Fresh sequence should grow independently when it advances. + _ = seqFresh.Predict(CreateTokenTensor(-0.25f)); + var freshStatsAfterSecond = seqFresh.GetInferenceStatistics(); + var freshLengthsAfterSecond = (int[])freshStatsAfterSecond["KVCache_SequenceLengths"]; + Assert.True( + freshLengthsAfterSecond[0] > freshLenAfterFirst, + $"Expected fresh KV-cache length to grow, but got {freshLenAfterFirst} -> {freshLengthsAfterSecond[0]}"); + } + + [Fact] + public void BeginInferenceSession_ResetRestoresInitialSequenceState() + { + var result = CreateDeterministicResult( + new InferenceOptimizationConfig + { + EnableFlashAttention = false, + EnableKVCache = true, + EnablePagedKVCache = false, + AttentionMasking = AttentionMaskingMode.Auto + }); + + var token1 = CreateTokenTensor(0.25f); + var token2 = CreateTokenTensor(0.5f); + + using var session = result.BeginInferenceSession(); + var seq = session.CreateSequence(); + + var y1 = seq.Predict(token1); + _ = seq.Predict(token2); + + seq.Reset(); + + var y1AfterReset = seq.Predict(token1); + AssertTensorsEqual(y1, y1AfterReset, Tolerance); + } + + [Fact] + public async Task BeginInferenceSession_ConcurrentPredict_MultipleSequences_DoesNotThrow() + { + var result = CreateDeterministicResult( + new InferenceOptimizationConfig + { + EnableFlashAttention = false, + EnableKVCache = true, + EnablePagedKVCache = true, + AttentionMasking = AttentionMaskingMode.Auto + }); + + using var session = result.BeginInferenceSession(); + var seqA = session.CreateSequence(); + var seqB = session.CreateSequence(); + + var tasks = Enumerable.Range(0, 20) + .Select(i => Task.Run(() => + { + var t = CreateTokenTensor(0.1f + (i * 0.01f)); + _ = (i % 2 == 0 ? seqA : seqB).Predict(t); + })) + .ToArray(); + + await Task.WhenAll(tasks); + + var statsA = seqA.GetInferenceStatistics(); + var statsB = seqB.GetInferenceStatistics(); + Assert.True((int)statsA["PagedAttentionLayerCount"] > 0); + Assert.True((int)statsB["PagedAttentionLayerCount"] > 0); + } + + [Fact] + public void BeginInferenceSession_KVCacheQuantization_Int8_UsesQuantizedStorage() + { + var result = CreateDeterministicResult( + new InferenceOptimizationConfig + { + EnableFlashAttention = false, + EnableKVCache = true, + EnablePagedKVCache = false, + KVCacheQuantization = KVCacheQuantizationMode.Int8, + AttentionMasking = AttentionMaskingMode.Auto + }); + + using var session = result.BeginInferenceSession(); + var seq = session.CreateSequence(); + + _ = seq.Predict(CreateTokenTensor(0.1f)); + + var stats = seq.GetInferenceStatistics(); + Assert.True(stats.TryGetValue("KVCache_DataType", out var dataType)); + Assert.Equal("Int8", dataType); + Assert.True(stats.TryGetValue("KVCache_UseInt8Storage", out var useInt8)); + Assert.True((bool)useInt8); + } + + [Fact] + public void BeginInferenceSession_KVCachePrecision_Auto_UsesFloat16Storage_ForFloatModel() + { + var result = CreateDeterministicResult( + new InferenceOptimizationConfig + { + EnableFlashAttention = false, + EnableKVCache = true, + EnablePagedKVCache = false, + KVCachePrecision = KVCachePrecisionMode.Auto, + KVCacheQuantization = KVCacheQuantizationMode.None, + AttentionMasking = AttentionMaskingMode.Auto + }); + + using var session = result.BeginInferenceSession(); + var seq = session.CreateSequence(); + + _ = seq.Predict(CreateTokenTensor(0.1f)); + + var stats = seq.GetInferenceStatistics(); + Assert.True(stats.TryGetValue("KVCache_DataType", out var dataType)); + Assert.Equal("Float16", dataType); + Assert.True(stats.TryGetValue("KVCache_UseFp16Storage", out var useFp16)); + Assert.True((bool)useFp16); + Assert.True(stats.TryGetValue("KVCache_UseInt8Storage", out var useInt8)); + Assert.False((bool)useInt8); + } + + [Fact] + public void BeginInferenceSession_SpeculativeDecoding_Configured_DoesNotRunDuringPredict() + { + var result = CreateDeterministicResult( + new InferenceOptimizationConfig + { + EnableFlashAttention = false, + EnableKVCache = false, + EnablePagedKVCache = false, + EnableSpeculativeDecoding = true, + DraftModelType = DraftModelType.NGram, + AttentionMasking = AttentionMaskingMode.Auto + }); + + using var session = result.BeginInferenceSession(); + var seq = session.CreateSequence(); + + _ = seq.Predict(CreateTokenTensor(0.1f)); + + var stats = seq.GetInferenceStatistics(); + Assert.True(stats.TryGetValue("SpeculativeDecodingEnabled", out var enabled)); + Assert.True((bool)enabled); + Assert.False(stats.ContainsKey("DraftModelType")); + Assert.False(stats.ContainsKey("SpeculationDepth")); + } + + [Fact] + public void BeginInferenceSession_PagedKVCache_IsInitialized_WhenEnabled() + { + var result = CreateDeterministicResult( + new InferenceOptimizationConfig + { + EnableFlashAttention = false, + EnableKVCache = true, + EnablePagedKVCache = true, + AttentionMasking = AttentionMaskingMode.Auto + }); + + using var session = result.BeginInferenceSession(); + var seq = session.CreateSequence(); + + _ = seq.Predict(CreateTokenTensor(0.1f)); + + var stats = seq.GetInferenceStatistics(); + Assert.True(stats.TryGetValue("PagedKVCacheInitialized", out var initialized)); + Assert.True((bool)initialized); + Assert.True(stats.TryGetValue("PagedAttentionLayerCount", out var count)); + Assert.True((int)count > 0); + } + + [Fact] + public void BeginInferenceSession_PagedAttention_WOQ_IsEnabled_WhenConfigured() + { + var result = CreateDeterministicResult( + new InferenceOptimizationConfig + { + EnableFlashAttention = false, + EnableKVCache = true, + EnablePagedKVCache = true, + EnableWeightOnlyQuantization = true, + AttentionMasking = AttentionMaskingMode.Auto + }); + + using var session = result.BeginInferenceSession(); + var seq = session.CreateSequence(); + + _ = seq.Predict(CreateTokenTensor(0.2f)); + + var stats = seq.GetInferenceStatistics(); + Assert.True(stats.TryGetValue("PagedAttentionWeightOnlyQuantizationEnabled", out var enabled)); + Assert.True((bool)enabled); + } + + [Fact] + public void BeginInferenceSession_MultiLoRA_TaskSelection_IsIsolatedPerSequence() + { + var config = new InferenceOptimizationConfig + { + EnableFlashAttention = false, + EnableKVCache = false, + EnablePagedKVCache = false, + EnableSpeculativeDecoding = false, + EnableBatching = false + }; + + var model = CreateDeterministicMultiLoRAModel(); + var result = CreateDeterministicResultWithModel(config, model); + + var token = CreateTokenTensor(0.25f); + + using var session = result.BeginInferenceSession(); + var seqA = session.CreateSequence("taskA"); + var seqB = session.CreateSequence("taskB"); + + var yA = seqA.Predict(token); + var yB = seqB.Predict(token); + + AssertTensorsNotEqual(yA, yB, minAbsDiff: 1e-3f); + + seqA.SetMultiLoRATask("taskB"); + var yA2 = seqA.Predict(token); + AssertTensorsNotEqual(yA, yA2, minAbsDiff: 1e-3f); + } + + [Fact] + public void BeginInferenceSession_MultiLoRA_TaskSwitch_ResetsKVCacheState_ForSameSequence() + { + var originalDiagnostics = Environment.GetEnvironmentVariable("AIDOTNET_DIAGNOSTICS"); + + var config = new InferenceOptimizationConfig + { + EnableFlashAttention = false, + EnableKVCache = true, + EnablePagedKVCache = false, + AttentionMasking = AttentionMaskingMode.Auto + }; + + var model = CreateDeterministicAttentionWithMultiLoRAModel(); + var result = CreateDeterministicResultWithModel(config, model); + + try + { + Environment.SetEnvironmentVariable("AIDOTNET_DIAGNOSTICS", "1"); + AiDotNet.Helpers.InferenceDiagnostics.Clear(); + + using var session = result.BeginInferenceSession(); + var seq = session.CreateSequence("taskA"); + + var token1 = CreateTokenTensor(0.25f); + var token2 = CreateTokenTensor(0.5f); + + _ = seq.Predict(token1); + var statsAfterFirst = seq.GetInferenceStatistics(); + var lenAfterFirst = ((int[])statsAfterFirst["KVCache_SequenceLengths"])[0]; + + _ = seq.Predict(token2); + var statsAfterSecond = seq.GetInferenceStatistics(); + var lenAfterSecond = ((int[])statsAfterSecond["KVCache_SequenceLengths"])[0]; + Assert.True(lenAfterSecond > lenAfterFirst, $"Expected KV-cache length to grow, but got {lenAfterFirst} -> {lenAfterSecond}"); + + seq.SetMultiLoRATask("taskB"); + _ = seq.Predict(token1); + var statsAfterSwitch = seq.GetInferenceStatistics(); + var lenAfterSwitch = ((int[])statsAfterSwitch["KVCache_SequenceLengths"])[0]; + + Assert.True(lenAfterSwitch <= lenAfterFirst, $"Expected KV-cache to reset after task switch, but got {lenAfterFirst} -> {lenAfterSwitch}"); + + var entries = AiDotNet.Helpers.InferenceDiagnostics.Snapshot(); + Assert.Contains(entries, e => e.Area == "InferenceSession" && e.Feature == "MultiLoRA" && e.Reason.Contains("Task=taskB")); + } + finally + { + AiDotNet.Helpers.InferenceDiagnostics.Clear(); + Environment.SetEnvironmentVariable("AIDOTNET_DIAGNOSTICS", originalDiagnostics); + } + } + + [Fact] + public void NeuralNetworkBase_Clone_DoesNotShareParameters() + { + var model = CreateDeterministicAttentionOnlyModel(); + var clone = (NeuralNetworkBase)model.Clone(); + + // Clone should preserve parameters exactly (deep copy via serialization/deserialization). + Assert.Equal(model.GetParameters().Length, clone.GetParameters().Length); + for (int i = 0; i < model.GetParameters().Length; i++) + { + Assert.True( + Math.Abs(model.GetParameters()[i] - clone.GetParameters()[i]) <= Tolerance, + $"Parameter mismatch at {i}: {model.GetParameters()[i]} != {clone.GetParameters()[i]}"); + } + + var cloneParams = clone.GetParameters(); + cloneParams[0] += 1.0f; + clone.UpdateParameters(cloneParams); + + Assert.NotEqual(model.GetParameters()[0], clone.GetParameters()[0]); + } + + private static PredictionModelResult, Tensor> CreateDeterministicResult(InferenceOptimizationConfig config) + { + var model = CreateDeterministicAttentionOnlyModel(); + return CreateDeterministicResultWithModel(config, model); + } + + private static PredictionModelResult, Tensor> CreateDeterministicResultWithModel( + InferenceOptimizationConfig config, + NeuralNetworkBase model) + { + if (model == null) throw new ArgumentNullException(nameof(model)); + + var optimization = new OptimizationResult, Tensor> + { + BestSolution = model + }; + + var normalization = new NormalizationInfo, Tensor> + { + Normalizer = new NoNormalizer, Tensor>(), + YParams = new NormalizationParameters { Method = NormalizationMethod.None } + }; + + var options = new PredictionModelResultOptions, Tensor> + { + OptimizationResult = optimization, + NormalizationInfo = normalization, + InferenceOptimizationConfig = config + }; + + return new PredictionModelResult, Tensor>(options); + } + + private static NeuralNetworkBase CreateDeterministicMultiLoRAModel() + { + const int inputSize = FlatSize; + const int outputSize = FlatSize; + + var baseDense = new DenseLayer(inputSize, outputSize, activationFunction: new AiDotNet.ActivationFunctions.IdentityActivation()); + var multi = new AiDotNet.LoRA.Adapters.MultiLoRAAdapter(baseDense, defaultTaskName: "taskA", defaultRank: 1, alpha: 1.0, freezeBaseLayer: true); + multi.AddTask("taskB", rank: 1, alpha: 1.0); + + var layers = new System.Collections.Generic.List> + { + new InputLayer(inputSize), + multi, + new DenseLayer(outputSize, outputSize, activationFunction: new AiDotNet.ActivationFunctions.IdentityActivation()) + }; + + var architecture = new NeuralNetworkArchitecture( + inputType: InputType.OneDimensional, + taskType: NeuralNetworkTaskType.Regression, + complexity: NetworkComplexity.Simple, + inputSize: inputSize, + outputSize: outputSize, + layers: layers); + + var model = new NeuralNetwork(architecture); + + // Deterministic base weights across the whole model. + var p = model.GetParameters(); + var deterministic = new float[p.Length]; + for (int i = 0; i < deterministic.Length; i++) + { + deterministic[i] = ((i % 19) - 9) / 9.0f; + } + model.UpdateParameters(new Vector(deterministic)); + + // Make taskB differ from taskA by setting distinct LoRA parameters. + // (Both A and B must be non-zero for the low-rank delta to have an effect.) + var taskA = multi.GetTaskAdapter("taskA"); + var taskB = multi.GetTaskAdapter("taskB"); + + var aParams = taskA.GetParameters(); + var bParams = taskB.GetParameters(); + + var a = new float[aParams.Length]; // all zeros => no delta + var b = new float[bParams.Length]; + for (int i = 0; i < b.Length; i++) + { + b[i] = 0.05f; + } + + taskA.UpdateParameters(new Vector(a)); + taskB.UpdateParameters(new Vector(b)); + + return model; + } + + private static NeuralNetworkBase CreateDeterministicAttentionWithMultiLoRAModel() + { + const int inputSize = FlatSize; + const int outputSize = FlatSize; + + var baseDense = new DenseLayer(outputSize, outputSize, activationFunction: new AiDotNet.ActivationFunctions.IdentityActivation()); + var multi = new AiDotNet.LoRA.Adapters.MultiLoRAAdapter(baseDense, defaultTaskName: "taskA", defaultRank: 1, alpha: 1.0, freezeBaseLayer: true); + multi.AddTask("taskB", rank: 1, alpha: 1.0); + + var layers = new System.Collections.Generic.List> + { + new InputLayer(inputSize), + new ReshapeLayer(new[] { FlatSize }, new[] { SequenceLength, EmbeddingDimension }), + new MultiHeadAttentionLayer( + sequenceLength: SequenceLength, + embeddingDimension: EmbeddingDimension, + headCount: HeadCount, + activationFunction: new AiDotNet.ActivationFunctions.IdentityActivation()), + new FlattenLayer(new[] { SequenceLength, EmbeddingDimension }), + multi, + new DenseLayer(outputSize, outputSize, activationFunction: new AiDotNet.ActivationFunctions.IdentityActivation()) + }; + + var architecture = new NeuralNetworkArchitecture( + inputType: InputType.OneDimensional, + taskType: NeuralNetworkTaskType.TextGeneration, + complexity: NetworkComplexity.Simple, + inputSize: inputSize, + outputSize: outputSize, + layers: layers); + + var model = new NeuralNetwork(architecture); + + var p = model.GetParameters(); + var deterministic = new float[p.Length]; + for (int i = 0; i < deterministic.Length; i++) + { + deterministic[i] = ((i % 23) - 11) / 11.0f; + } + model.UpdateParameters(new Vector(deterministic)); + + var taskA = multi.GetTaskAdapter("taskA"); + var taskB = multi.GetTaskAdapter("taskB"); + + var aParams = taskA.GetParameters(); + var bParams = taskB.GetParameters(); + + var a = new float[aParams.Length]; + var b = new float[bParams.Length]; + for (int i = 0; i < b.Length; i++) + { + b[i] = 0.05f; + } + + taskA.UpdateParameters(new Vector(a)); + taskB.UpdateParameters(new Vector(b)); + + return model; + } + + private static NeuralNetworkBase CreateDeterministicAttentionOnlyModel() + { + var layers = new System.Collections.Generic.List> + { + new InputLayer(FlatSize), + new ReshapeLayer(new[] { FlatSize }, new[] { SequenceLength, EmbeddingDimension }), + new MultiHeadAttentionLayer( + sequenceLength: SequenceLength, + embeddingDimension: EmbeddingDimension, + headCount: HeadCount, + activationFunction: new AiDotNet.ActivationFunctions.IdentityActivation()) + , + new FlattenLayer(new[] { SequenceLength, EmbeddingDimension }), + new DenseLayer(FlatSize, FlatSize, activationFunction: new AiDotNet.ActivationFunctions.IdentityActivation()) + }; + + var architecture = new NeuralNetworkArchitecture( + inputType: InputType.OneDimensional, + taskType: NeuralNetworkTaskType.TextGeneration, + complexity: NetworkComplexity.Simple, + inputSize: FlatSize, + outputSize: FlatSize, + layers: layers); + + var model = new NeuralNetwork(architecture); + + var p = model.GetParameters(); + var deterministic = new float[p.Length]; + for (int i = 0; i < deterministic.Length; i++) + { + deterministic[i] = ((i % 23) - 11) / 11.0f; + } + model.UpdateParameters(new Vector(deterministic)); + + return model; + } + + private static Tensor CreateTokenTensor(float scalar) + { + var t = new Tensor(new[] { 1, FlatSize }); + for (int i = 0; i < t.Length; i++) + { + t[i] = scalar + (i * 0.01f); + } + return t; + } + + private static void AssertTensorsEqual(Tensor a, Tensor b, float tolerance) + { + Assert.Equal(a.Shape, b.Shape); + for (int i = 0; i < a.Length; i++) + { + Assert.True(Math.Abs(a[i] - b[i]) <= tolerance, $"Index {i}: {a[i]} != {b[i]}"); + } + } + + private static void AssertTensorsNotEqual(Tensor a, Tensor b, float minAbsDiff) + { + Assert.Equal(a.Shape, b.Shape); + + float maxAbs = 0f; + for (int i = 0; i < a.Length; i++) + { + float abs = Math.Abs(a[i] - b[i]); + if (abs > maxAbs) + { + maxAbs = abs; + } + } + + Assert.True(maxAbs >= minAbsDiff, $"Expected tensors to differ by at least {minAbsDiff}, but max diff was {maxAbs}"); + } +} diff --git a/tests/AiDotNet.Tests/StressTests/GpuStressTests.cs b/tests/AiDotNet.Tests/StressTests/GpuStressTests.cs index 61cde0547c..dfa28fcdc3 100644 --- a/tests/AiDotNet.Tests/StressTests/GpuStressTests.cs +++ b/tests/AiDotNet.Tests/StressTests/GpuStressTests.cs @@ -201,15 +201,16 @@ public void Conv2D_LongRun_1KIterations_StablePerformance() var lastQuartileAvg = timings.Skip(3 * MediumRunIterations / 4).Average(); // Guard against zero division on very fast hardware - double performanceDrift = 0; - if (firstQuartileAvg > 0) + // Only check for degradation (last > first), not improvement + double performanceDegradation = 0; + if (firstQuartileAvg > 0 && lastQuartileAvg > firstQuartileAvg) { - performanceDrift = Math.Abs(lastQuartileAvg - firstQuartileAvg) / firstQuartileAvg; + performanceDegradation = (lastQuartileAvg - firstQuartileAvg) / firstQuartileAvg; } - // Performance should not degrade by more than 20% - Assert.True(performanceDrift < 0.20, - $"Performance degraded by {performanceDrift * 100:F1}% (first: {firstQuartileAvg:F2}ms, last: {lastQuartileAvg:F2}ms)"); + // Performance should not degrade by more than 20% (improvement is acceptable) + Assert.True(performanceDegradation < 0.20, + $"Performance degraded by {performanceDegradation * 100:F1}% (first: {firstQuartileAvg:F2}ms, last: {lastQuartileAvg:F2}ms)"); // Memory growth should be minimal Assert.True(memoryGrowth < 20_000_000, diff --git a/tests/AiDotNet.Tests/TestInfrastructure/DiagnosticsEnvironmentCollection.cs b/tests/AiDotNet.Tests/TestInfrastructure/DiagnosticsEnvironmentCollection.cs new file mode 100644 index 0000000000..cc02f7c3f7 --- /dev/null +++ b/tests/AiDotNet.Tests/TestInfrastructure/DiagnosticsEnvironmentCollection.cs @@ -0,0 +1,28 @@ +using System; +using AiDotNet.Helpers; +using Xunit; + +namespace AiDotNet.Tests.TestInfrastructure; + +[CollectionDefinition(Name, DisableParallelization = true)] +public sealed class DiagnosticsEnvironmentCollection : ICollectionFixture +{ + public const string Name = "DiagnosticsEnv"; + + public sealed class Fixture : IDisposable + { + private readonly string? _original; + + public Fixture() + { + _original = Environment.GetEnvironmentVariable("AIDOTNET_DIAGNOSTICS"); + } + + public void Dispose() + { + InferenceDiagnostics.Clear(); + Environment.SetEnvironmentVariable("AIDOTNET_DIAGNOSTICS", _original); + } + } +} + diff --git a/tests/AiDotNet.Tests/UnitTests/Attention/FlashAttentionTests.cs b/tests/AiDotNet.Tests/UnitTests/Attention/FlashAttentionTests.cs index 1019b5e8df..5b8bd00fe9 100644 --- a/tests/AiDotNet.Tests/UnitTests/Attention/FlashAttentionTests.cs +++ b/tests/AiDotNet.Tests/UnitTests/Attention/FlashAttentionTests.cs @@ -111,6 +111,41 @@ public void FlashAttention_WithCausalMask_MasksCorrectly() } } + [Fact] + public void FlashAttention_WithCausalMask_RespectsQueryOffsetForCachedDecoding() + { + // Arrange + int batchSize = 1; + int seqLenKV = 6; + int seqLenQ = 2; + int headDim = 8; + + var query = CreateRandomTensor(batchSize, seqLenQ, headDim, seed: 42); + var key = CreateRandomTensor(batchSize, seqLenKV, headDim, seed: 43); + var value = CreateRandomTensor(batchSize, seqLenKV, headDim, seed: 44); + + int queryOffset = seqLenKV - seqLenQ; + var config = new FlashAttentionConfig { UseCausalMask = true, ReturnAttentionWeights = true }; + + // Act + var (_, attnWeights) = FlashAttention.Forward(query, key, value, config, queryOffset: queryOffset); + + // Assert - For each query row, positions beyond (queryOffset + qIdx) must be masked + Assert.NotNull(attnWeights); + for (int b = 0; b < batchSize; b++) + { + for (int q = 0; q < seqLenQ; q++) + { + int maxAllowedK = queryOffset + q; + for (int k = maxAllowedK + 1; k < seqLenKV; k++) + { + float weight = attnWeights[new[] { b, q, k }]; + Assert.True(weight < 1e-6f, $"Position (q={q}, k={k}) should be masked but has weight {weight}"); + } + } + } + } + [Fact] public void FlashAttention_AttentionWeightsRowSumToOne() { diff --git a/tests/AiDotNet.Tests/UnitTests/Helpers/DeserializationHelperTests.cs b/tests/AiDotNet.Tests/UnitTests/Helpers/DeserializationHelperTests.cs index eb7fc14de0..b44c1526b2 100644 --- a/tests/AiDotNet.Tests/UnitTests/Helpers/DeserializationHelperTests.cs +++ b/tests/AiDotNet.Tests/UnitTests/Helpers/DeserializationHelperTests.cs @@ -3,6 +3,9 @@ using System.IO; using AiDotNet.Enums; using AiDotNet.Helpers; +using AiDotNet.Inference; +using AiDotNet.LoRA.Adapters; +using AiDotNet.NeuralNetworks.Attention; using AiDotNet.NeuralNetworks.Layers; using Xunit; @@ -374,5 +377,170 @@ public void CreateLayerFromType_WithPoolingLayerAverageType_CreatesCorrectly() Assert.NotNull(layer); Assert.IsType>(layer); } + + [Fact] + public void DeserializeInterface_WhenTypeDoesNotImplementInterface_Throws() + { + // Arrange + using var ms = new MemoryStream(); + using var writer = new BinaryWriter(ms); + writer.Write(typeof(List).AssemblyQualifiedName); + ms.Position = 0; + using var reader = new BinaryReader(ms); + + // Act & Assert + Assert.Throws(() => + DeserializationHelper.DeserializeInterface(reader)); + } + + private sealed class NoDefaultCtorDisposable : IDisposable + { + public NoDefaultCtorDisposable(int _) { } + public void Dispose() { } + } + + [Fact] + public void DeserializeInterface_WhenNoParameterlessCtor_ReturnsNull() + { + // Arrange + using var ms = new MemoryStream(); + using var writer = new BinaryWriter(ms); + writer.Write(typeof(NoDefaultCtorDisposable).AssemblyQualifiedName); + ms.Position = 0; + using var reader = new BinaryReader(ms); + + // Act + var result = DeserializationHelper.DeserializeInterface(reader); + + // Assert + Assert.Null(result); + } + + [Fact] + public void CreateLayerFromType_WithAttentionLayers_CreatesCorrectly() + { + // Arrange + var inputShape = new int[] { 8, 16 }; + var outputShape = new int[] { 8, 16 }; + + // Act & Assert + Assert.IsType>( + DeserializationHelper.CreateLayerFromType(typeof(MultiHeadAttentionLayer<>).Name, inputShape, outputShape)); + + Assert.IsType>( + DeserializationHelper.CreateLayerFromType(typeof(SelfAttentionLayer<>).Name, inputShape, outputShape)); + + Assert.IsType>( + DeserializationHelper.CreateLayerFromType(typeof(FlashAttentionLayer<>).Name, inputShape, outputShape, + new Dictionary { { "UseCausalMask", true } })); + + Assert.IsType>( + DeserializationHelper.CreateLayerFromType(typeof(CachedMultiHeadAttention<>).Name, inputShape, outputShape)); + + Assert.IsType>( + DeserializationHelper.CreateLayerFromType(typeof(PagedCachedMultiHeadAttention<>).Name, inputShape, outputShape)); + } + + [Fact] + public void CreateLayerFromType_WithAttentionLayer_CreatesCorrectly() + { + // Arrange + var inputShape = new int[] { 16 }; + var outputShape = new int[] { 8 }; + + // Act + var layer = DeserializationHelper.CreateLayerFromType(typeof(AttentionLayer<>).Name, inputShape, outputShape); + + // Assert + Assert.NotNull(layer); + Assert.IsType>(layer); + } + + [Fact] + public void CreateLayerFromType_WithGraphAttentionLayer_CreatesCorrectly() + { + // Arrange + var inputShape = new int[] { 16 }; + var outputShape = new int[] { 8 }; + + // Act + var layer = DeserializationHelper.CreateLayerFromType(typeof(GraphAttentionLayer<>).Name, inputShape, outputShape, + new Dictionary { { "NumHeads", 2 } }); + + // Assert + Assert.NotNull(layer); + Assert.IsType>(layer); + } + + [Fact] + public void CreateLayerFromType_WithDropoutAndLayerNorm_CreatesCorrectly() + { + // Arrange + var inputShape = new int[] { 16 }; + var outputShape = new int[] { 16 }; + + // Act & Assert + Assert.IsType>( + DeserializationHelper.CreateLayerFromType(typeof(DropoutLayer<>).Name, inputShape, outputShape)); + + Assert.IsType>( + DeserializationHelper.CreateLayerFromType(typeof(LayerNormalizationLayer<>).Name, inputShape, outputShape)); + } + + [Fact] + public void CreateLayerFromType_WithPositionalEncoding_CreatesCorrectly() + { + // Arrange + var inputShape = new int[] { 128, 16 }; + var outputShape = new int[] { 128, 16 }; + + // Act + var layer = DeserializationHelper.CreateLayerFromType(typeof(PositionalEncodingLayer<>).Name, inputShape, outputShape); + + // Assert + Assert.NotNull(layer); + Assert.IsType>(layer); + } + + [Fact] + public void CreateLayerFromType_WithEncodedParamsInIdentifier_Works() + { + // Arrange + var inputShape = new int[] { 8, 16 }; + var outputShape = new int[] { 8, 16 }; + + // Act + var layer = DeserializationHelper.CreateLayerFromType( + $"{typeof(PagedCachedMultiHeadAttention<>).Name};HeadCount=2;UseCausalMask=false", + inputShape, + outputShape); + + // Assert + Assert.NotNull(layer); + Assert.IsType>(layer); + } + + [Fact] + public void CreateLayerFromType_WithMultiLoRAAdapter_CreatesCorrectly() + { + // Arrange + var inputShape = new int[] { 10 }; + var outputShape = new int[] { 5 }; + + var additionalParams = new Dictionary + { + { "Tasks", "taskA|taskB" }, + { "TaskRanks", "2|4" }, + { "TaskAlphas", "1.0|2.0" }, + { "CurrentTask", Uri.EscapeDataString("taskB") } + }; + + // Act + var layer = DeserializationHelper.CreateLayerFromType(typeof(MultiLoRAAdapter<>).Name, inputShape, outputShape, additionalParams); + + // Assert + Assert.NotNull(layer); + Assert.IsType>(layer); + } } } diff --git a/tests/AiDotNet.Tests/UnitTests/Helpers/InferenceDiagnosticsTests.cs b/tests/AiDotNet.Tests/UnitTests/Helpers/InferenceDiagnosticsTests.cs new file mode 100644 index 0000000000..0deb5e6990 --- /dev/null +++ b/tests/AiDotNet.Tests/UnitTests/Helpers/InferenceDiagnosticsTests.cs @@ -0,0 +1,50 @@ +using System; +using AiDotNet.Helpers; +using Xunit; + +namespace AiDotNet.Tests.UnitTests.Helpers; + +[Collection(AiDotNet.Tests.TestInfrastructure.DiagnosticsEnvironmentCollection.Name)] +public class InferenceDiagnosticsTests +{ + [Fact] + public void InferenceDiagnostics_Disabled_DoesNotRecord() + { + var original = Environment.GetEnvironmentVariable("AIDOTNET_DIAGNOSTICS"); + try + { + Environment.SetEnvironmentVariable("AIDOTNET_DIAGNOSTICS", null); + InferenceDiagnostics.Clear(); + + InferenceDiagnostics.RecordDecision("Test", "Feature", enabled: true, reason: "Reason"); + + Assert.Empty(InferenceDiagnostics.Snapshot()); + } + finally + { + InferenceDiagnostics.Clear(); + Environment.SetEnvironmentVariable("AIDOTNET_DIAGNOSTICS", original); + } + } + + [Fact] + public void InferenceDiagnostics_Enabled_Records() + { + var original = Environment.GetEnvironmentVariable("AIDOTNET_DIAGNOSTICS"); + try + { + Environment.SetEnvironmentVariable("AIDOTNET_DIAGNOSTICS", "1"); + InferenceDiagnostics.Clear(); + + InferenceDiagnostics.RecordDecision("Test", "Feature", enabled: true, reason: "Reason"); + + var entries = InferenceDiagnostics.Snapshot(); + Assert.Contains(entries, e => e.Area == "Test" && e.Feature == "Feature" && e.Enabled); + } + finally + { + InferenceDiagnostics.Clear(); + Environment.SetEnvironmentVariable("AIDOTNET_DIAGNOSTICS", original); + } + } +} diff --git a/tests/AiDotNet.Tests/UnitTests/Inference/InferenceOptimizerTests.cs b/tests/AiDotNet.Tests/UnitTests/Inference/InferenceOptimizerTests.cs new file mode 100644 index 0000000000..31a16d1b79 --- /dev/null +++ b/tests/AiDotNet.Tests/UnitTests/Inference/InferenceOptimizerTests.cs @@ -0,0 +1,384 @@ +using AiDotNet.Configuration; +using AiDotNet.Enums; +using AiDotNet.Inference; +using AiDotNet.NeuralNetworks; +using AiDotNet.NeuralNetworks.Attention; +using AiDotNet.NeuralNetworks.Layers; +using Xunit; + +namespace AiDotNet.Tests.UnitTests.Inference; + +[Collection(AiDotNet.Tests.TestInfrastructure.DiagnosticsEnvironmentCollection.Name)] +public class InferenceOptimizerTests +{ + [Fact] + public void InferenceOptimizer_WhenDiagnosticsEnabled_RecordsDecisions() + { + var original = Environment.GetEnvironmentVariable("AIDOTNET_DIAGNOSTICS"); + try + { + Environment.SetEnvironmentVariable("AIDOTNET_DIAGNOSTICS", "1"); + AiDotNet.Helpers.InferenceDiagnostics.Clear(); + + var model = CreateTinyTransformer(taskType: NeuralNetworkTaskType.TextGeneration); + var config = new InferenceOptimizationConfig + { + EnableKVCache = true, + EnableFlashAttention = false, + EnablePagedKVCache = false, + AttentionMasking = AttentionMaskingMode.Auto + }; + + var optimizer = new InferenceOptimizer(config); + _ = optimizer.OptimizeForInference(model, cloneModel: false); + + var entries = AiDotNet.Helpers.InferenceDiagnostics.Snapshot(); + Assert.Contains(entries, e => e.Area == "InferenceOptimizer" && e.Feature == "KVCachePrecision"); + } + finally + { + AiDotNet.Helpers.InferenceDiagnostics.Clear(); + Environment.SetEnvironmentVariable("AIDOTNET_DIAGNOSTICS", original); + } + } + + [Fact] + public void InferenceOptimizer_RewritesMultiHeadAttention_ToFlashAttention_WhenEnabled() + { + var model = CreateTinyTransformer(taskType: NeuralNetworkTaskType.Regression); + Assert.Contains(model.Layers, l => l is MultiHeadAttentionLayer); + + var config = new InferenceOptimizationConfig + { + EnableKVCache = false, + EnableFlashAttention = true, + AttentionMasking = AttentionMaskingMode.Disabled + }; + + var optimizer = new InferenceOptimizer(config); + // Clone relies on serialization of every layer in the graph; this test focuses on rewrite behavior. + var (optimized, anyApplied) = optimizer.OptimizeForInference(model, cloneModel: false); + + Assert.True(anyApplied); + Assert.Contains(optimized.Layers, l => l is FlashAttentionLayer); + Assert.DoesNotContain(optimized.Layers, l => l is MultiHeadAttentionLayer); + } + + [Fact] + public void InferenceOptimizer_RewritesMultiHeadAttention_ToCachedAttention_ForTextGeneration_WhenKVCacheEnabled() + { + var model = CreateTinyTransformer(taskType: NeuralNetworkTaskType.TextGeneration); + Assert.Contains(model.Layers, l => l is MultiHeadAttentionLayer); + + var config = new InferenceOptimizationConfig + { + EnableKVCache = true, + EnableFlashAttention = true, + // Paged KV-cache is industry-standard and enabled by default; keep it enabled for this test. + AttentionMasking = AttentionMaskingMode.Auto + }; + + var optimizer = new InferenceOptimizer(config); + var (optimized, anyApplied) = optimizer.OptimizeForInference(model, cloneModel: true); + + Assert.True(anyApplied); + Assert.Contains(optimized.Layers, l => l is PagedCachedMultiHeadAttention); + Assert.DoesNotContain(optimized.Layers, l => l is MultiHeadAttentionLayer); + + foreach (var layer in optimized.Layers) + { + if (layer is PagedCachedMultiHeadAttention cached) + { + Assert.True(cached.InferenceMode); + Assert.NotNull(cached.Kernel); + } + } + } + + [Fact] + public void InferenceOptimizer_RewritesSelfAttention_ToCachedAttention_WhenKVCacheEnabled() + { + var model = CreateTinySelfAttentionModel(taskType: NeuralNetworkTaskType.TextGeneration); + Assert.Contains(model.Layers, l => l is SelfAttentionLayer); + + var config = new InferenceOptimizationConfig + { + EnableKVCache = true, + EnablePagedKVCache = false, + EnableFlashAttention = false, + AttentionMasking = AttentionMaskingMode.Auto + }; + + var optimizer = new InferenceOptimizer(config); + var (optimized, anyApplied) = optimizer.OptimizeForInference(model, cloneModel: false); + + Assert.True(anyApplied); + Assert.Contains(optimized.Layers, l => l is CachedMultiHeadAttention); + Assert.DoesNotContain(optimized.Layers, l => l is SelfAttentionLayer); + + // In-place rewrite expected when cloneModel=false. + Assert.DoesNotContain(model.Layers, l => l is SelfAttentionLayer); + } + + [Fact] + public void InferenceOptimizer_SpeculativeDecoding_FallsBackToNGram_WhenSmallNeuralUnavailable() + { + var model = CreateTinyTransformer(taskType: NeuralNetworkTaskType.TextGeneration); + + var config = new InferenceOptimizationConfig + { + EnableKVCache = false, + EnableFlashAttention = false, + EnableSpeculativeDecoding = true, + DraftModelType = DraftModelType.SmallNeural + }; + + var optimizer = new InferenceOptimizer(config); + + // Should never throw: SmallNeural draft models are not available in MVP and must fall back. + var (_, anyApplied) = optimizer.OptimizeForInference(model, cloneModel: false); + + Assert.True(anyApplied); + Assert.NotNull(optimizer.DraftModel); + Assert.Contains("NGramDraftModel", optimizer.DraftModel!.GetType().Name); + Assert.True(optimizer.DraftModel!.VocabSize > 0); + } + + [Fact] + public void InferenceOptimizer_SpeculativeDecoding_FallsBackToNGram_WhenCustomNotProvided() + { + var model = CreateTinyTransformer(taskType: NeuralNetworkTaskType.TextGeneration); + + var config = new InferenceOptimizationConfig + { + EnableKVCache = false, + EnableFlashAttention = false, + EnableSpeculativeDecoding = true, + DraftModelType = DraftModelType.Custom + }; + + var optimizer = new InferenceOptimizer(config); + + // Should never throw: the public facade does not wire custom draft models in MVP. + var (_, anyApplied) = optimizer.OptimizeForInference(model, cloneModel: false); + + Assert.True(anyApplied); + Assert.NotNull(optimizer.DraftModel); + Assert.True(optimizer.DraftModel!.VocabSize > 0); + } + + [Fact] + public void InferenceOptimizer_WeightOnlyQuantization_RewritesDenseLayer_OnClonedModel_AndPreservesOutputs() + { + var model = CreateTinyDenseModel(); + + var input = new AiDotNet.Tensors.LinearAlgebra.Tensor(new[] { 1, 4 }); + for (int i = 0; i < input.Length; i++) + { + input[i] = 0.1f * (i + 1); + } + + var baseline = model.Predict(input); + + var config = new InferenceOptimizationConfig + { + EnableKVCache = false, + EnableFlashAttention = false, + EnableWeightOnlyQuantization = true + }; + + var optimizer = new InferenceOptimizer(config); + var (optimized, anyApplied) = optimizer.OptimizeForInference(model, cloneModel: true); + + Assert.True(anyApplied); + Assert.Contains(optimized.Layers, l => l.GetType().Name.Contains("QuantizedDenseLayer")); + Assert.Contains(model.Layers, l => l is DenseLayer); + + var y = optimized.Predict(input); + Assert.Equal(baseline.Shape, y.Shape); + + for (int i = 0; i < y.Length; i++) + { + Assert.True(Math.Abs(baseline[i] - y[i]) < 1e-1f, $"Mismatch at {i}: {baseline[i]} vs {y[i]}"); + } + } + + [Fact] + public void InferenceOptimizer_Skips_AttentionLayer_WhenKVCacheEnabled() + { + var model = CreateTinyAttentionLayerModel(); + + var config = new InferenceOptimizationConfig + { + EnableKVCache = true, + EnableFlashAttention = true, + AttentionMasking = AttentionMaskingMode.Auto + }; + + var optimizer = new InferenceOptimizer(config); + var (optimized, anyApplied) = optimizer.OptimizeForInference(model, cloneModel: true); + + Assert.False(anyApplied); + Assert.Same(model, optimized); + Assert.Contains(optimized.Layers, l => l is AttentionLayer); + } + + [Fact] + public void InferenceOptimizer_Skips_GraphAttentionLayer_WhenKVCacheEnabled() + { + var model = CreateTinyGraphAttentionModel(); + + var config = new InferenceOptimizationConfig + { + EnableKVCache = true, + EnableFlashAttention = true, + AttentionMasking = AttentionMaskingMode.Auto + }; + + var optimizer = new InferenceOptimizer(config); + + // Should not throw: graph attention is not part of inference-time transformer KV-cache rewriting. + var (optimized, anyApplied) = optimizer.OptimizeForInference(model, cloneModel: true); + + Assert.False(anyApplied); + Assert.Same(model, optimized); + Assert.Contains(optimized.Layers, l => l is GraphAttentionLayer); + } + + private static Transformer CreateTinyTransformer(NeuralNetworkTaskType taskType) + { + var architecture = new TransformerArchitecture( + inputType: InputType.OneDimensional, + taskType: taskType, + numEncoderLayers: 1, + numDecoderLayers: 0, + numHeads: 2, + modelDimension: 8, + feedForwardDimension: 16, + complexity: NetworkComplexity.Simple, + inputSize: 1, + outputSize: 8, + dropoutRate: 0.0, + maxSequenceLength: 4, + vocabularySize: 0, + usePositionalEncoding: false); + + return new Transformer(architecture); + } + + private static NeuralNetworkBase CreateTinySelfAttentionModel(NeuralNetworkTaskType taskType) + { + const int seqLen = 4; + const int embDim = 8; + const int headCount = 2; + const int flatSize = seqLen * embDim; + + var layers = new System.Collections.Generic.List> + { + new InputLayer(flatSize), + new ReshapeLayer(new[] { flatSize }, new[] { seqLen, embDim }), + new SelfAttentionLayer(seqLen, embDim, headCount, activationFunction: new AiDotNet.ActivationFunctions.IdentityActivation()), + new FlattenLayer(new[] { seqLen, embDim }), + new DenseLayer(flatSize, flatSize, activationFunction: new AiDotNet.ActivationFunctions.IdentityActivation()) + }; + + var architecture = new NeuralNetworkArchitecture( + inputType: InputType.OneDimensional, + taskType: taskType, + complexity: NetworkComplexity.Simple, + inputSize: flatSize, + outputSize: flatSize, + layers: layers); + + var model = new NeuralNetwork(architecture); + + // Ensure parameters are initialized deterministically for stable tests. + var p = model.GetParameters(); + var deterministic = new float[p.Length]; + for (int i = 0; i < deterministic.Length; i++) + { + deterministic[i] = ((i % 17) - 8) / 8.0f; + } + model.UpdateParameters(new AiDotNet.Tensors.LinearAlgebra.Vector(deterministic)); + + return model; + } + + private static NeuralNetworkBase CreateTinyDenseModel() + { + const int inSize = 4; + const int outSize = 3; + + var layers = new System.Collections.Generic.List> + { + new InputLayer(inSize), + new DenseLayer(inSize, outSize, activationFunction: new AiDotNet.ActivationFunctions.IdentityActivation()) + }; + + var architecture = new NeuralNetworkArchitecture( + inputType: InputType.OneDimensional, + taskType: NeuralNetworkTaskType.Regression, + complexity: NetworkComplexity.Simple, + inputSize: inSize, + outputSize: outSize, + layers: layers); + + var model = new NeuralNetwork(architecture); + + var p = model.GetParameters(); + var deterministic = new float[p.Length]; + for (int i = 0; i < deterministic.Length; i++) + { + deterministic[i] = ((i % 13) - 6) / 6.0f; + } + model.UpdateParameters(new AiDotNet.Tensors.LinearAlgebra.Vector(deterministic)); + + return model; + } + + private static NeuralNetworkBase CreateTinyAttentionLayerModel() + { + const int inputSize = 8; + const int attentionSize = 8; + + var layers = new System.Collections.Generic.List> + { + new InputLayer(inputSize), + new AttentionLayer(inputSize, attentionSize, activation: (AiDotNet.Interfaces.IActivationFunction?)null), + new DenseLayer(attentionSize, attentionSize, activationFunction: new AiDotNet.ActivationFunctions.IdentityActivation()) + }; + + var architecture = new NeuralNetworkArchitecture( + inputType: InputType.OneDimensional, + taskType: NeuralNetworkTaskType.Regression, + complexity: NetworkComplexity.Simple, + inputSize: inputSize, + outputSize: attentionSize, + layers: layers); + + return new NeuralNetwork(architecture); + } + + private static NeuralNetworkBase CreateTinyGraphAttentionModel() + { + const int inputSize = 8; + const int outputSize = 8; + + var layers = new System.Collections.Generic.List> + { + new InputLayer(inputSize), + new GraphAttentionLayer(inputSize, outputSize, numHeads: 1), + new DenseLayer(outputSize, outputSize, activationFunction: new AiDotNet.ActivationFunctions.IdentityActivation()) + }; + + var architecture = new NeuralNetworkArchitecture( + inputType: InputType.OneDimensional, + taskType: NeuralNetworkTaskType.Regression, + complexity: NetworkComplexity.Simple, + inputSize: inputSize, + outputSize: outputSize, + layers: layers); + + return new NeuralNetwork(architecture); + } +} diff --git a/tests/AiDotNet.Tests/UnitTests/Inference/KVCacheTests.cs b/tests/AiDotNet.Tests/UnitTests/Inference/KVCacheTests.cs new file mode 100644 index 0000000000..14bfdbd4be --- /dev/null +++ b/tests/AiDotNet.Tests/UnitTests/Inference/KVCacheTests.cs @@ -0,0 +1,131 @@ +using AiDotNet.Inference; +using AiDotNet.Tensors.LinearAlgebra; +using Xunit; + +namespace AiDotNet.Tests.UnitTests.Inference; + +public class KVCacheTests +{ + [Fact] + public void KVCache_AppendAcrossLayers_MaintainsIndependentLengths() + { + var config = new KVCacheConfig + { + NumLayers = 2, + NumHeads = 1, + HeadDimension = 2, + MaxSequenceLength = 8, + MaxBatchSize = 1, + PreAllocate = true + }; + + var cache = new KVCache(config); + + var keys0 = new Tensor(new[] { 1, 1, 2, 2 }); + var values0 = new Tensor(new[] { 1, 1, 2, 2 }); + keys0[new[] { 0, 0, 0, 0 }] = 1f; + keys0[new[] { 0, 0, 0, 1 }] = 2f; + keys0[new[] { 0, 0, 1, 0 }] = 3f; + keys0[new[] { 0, 0, 1, 1 }] = 4f; + values0[new[] { 0, 0, 0, 0 }] = 5f; + values0[new[] { 0, 0, 0, 1 }] = 6f; + values0[new[] { 0, 0, 1, 0 }] = 7f; + values0[new[] { 0, 0, 1, 1 }] = 8f; + + var (layer0Keys, _) = cache.Append(0, keys0, values0); + Assert.Equal(2, layer0Keys.Shape[2]); + + var keys1 = new Tensor(new[] { 1, 1, 2, 2 }); + var values1 = new Tensor(new[] { 1, 1, 2, 2 }); + keys1[new[] { 0, 0, 0, 0 }] = 10f; + keys1[new[] { 0, 0, 0, 1 }] = 11f; + keys1[new[] { 0, 0, 1, 0 }] = 12f; + keys1[new[] { 0, 0, 1, 1 }] = 13f; + values1[new[] { 0, 0, 0, 0 }] = 14f; + values1[new[] { 0, 0, 0, 1 }] = 15f; + values1[new[] { 0, 0, 1, 0 }] = 16f; + values1[new[] { 0, 0, 1, 1 }] = 17f; + + var (layer1Keys, _) = cache.Append(1, keys1, values1); + Assert.Equal(2, layer1Keys.Shape[2]); + Assert.Equal(10f, layer1Keys[new[] { 0, 0, 0, 0 }]); + Assert.Equal(13f, layer1Keys[new[] { 0, 0, 1, 1 }]); + + var (layer0KeysAfter, _) = cache.GetCached(0, batchSize: 1); + Assert.Equal(2, layer0KeysAfter.Shape[2]); + Assert.Equal(1f, layer0KeysAfter[new[] { 0, 0, 0, 0 }]); + Assert.Equal(4f, layer0KeysAfter[new[] { 0, 0, 1, 1 }]); + } + + [Fact] + public void KVCache_Float16Storage_RoundTripsValues() + { + var config = new KVCacheConfig + { + NumLayers = 1, + NumHeads = 1, + HeadDimension = 2, + MaxSequenceLength = 8, + MaxBatchSize = 1, + PreAllocate = true, + DataType = CacheDataType.Float16 + }; + + var cache = new KVCache(config); + + var keys = new Tensor(new[] { 1, 1, 2, 2 }); + var values = new Tensor(new[] { 1, 1, 2, 2 }); + keys[new[] { 0, 0, 0, 0 }] = 1f; + keys[new[] { 0, 0, 0, 1 }] = 2f; + keys[new[] { 0, 0, 1, 0 }] = 3f; + keys[new[] { 0, 0, 1, 1 }] = 4f; + values[new[] { 0, 0, 0, 0 }] = 5f; + values[new[] { 0, 0, 0, 1 }] = 6f; + values[new[] { 0, 0, 1, 0 }] = 7f; + values[new[] { 0, 0, 1, 1 }] = 8f; + + var (cachedKeys, cachedValues) = cache.Append(0, keys, values); + Assert.Equal(2, cachedKeys.Shape[2]); + Assert.Equal(1f, cachedKeys[new[] { 0, 0, 0, 0 }]); + Assert.Equal(4f, cachedKeys[new[] { 0, 0, 1, 1 }]); + Assert.Equal(5f, cachedValues[new[] { 0, 0, 0, 0 }]); + Assert.Equal(8f, cachedValues[new[] { 0, 0, 1, 1 }]); + } + + [Fact] + public void KVCache_Int8Storage_RoundTripsApproximately() + { + var config = new KVCacheConfig + { + NumLayers = 1, + NumHeads = 1, + HeadDimension = 2, + MaxSequenceLength = 8, + MaxBatchSize = 1, + PreAllocate = true, + DataType = CacheDataType.Int8 + }; + + var cache = new KVCache(config); + + var keys = new Tensor(new[] { 1, 1, 2, 2 }); + var values = new Tensor(new[] { 1, 1, 2, 2 }); + keys[new[] { 0, 0, 0, 0 }] = 1f; + keys[new[] { 0, 0, 0, 1 }] = 2f; + keys[new[] { 0, 0, 1, 0 }] = 3f; + keys[new[] { 0, 0, 1, 1 }] = 4f; + values[new[] { 0, 0, 0, 0 }] = 5f; + values[new[] { 0, 0, 0, 1 }] = 6f; + values[new[] { 0, 0, 1, 0 }] = 7f; + values[new[] { 0, 0, 1, 1 }] = 8f; + + var (cachedKeys, cachedValues) = cache.Append(0, keys, values); + Assert.Equal(2, cachedKeys.Shape[2]); + + // Int8 quantization is approximate; tolerate small error. + Assert.InRange(Math.Abs(cachedKeys[new[] { 0, 0, 0, 0 }] - 1f), 0f, 0.1f); + Assert.InRange(Math.Abs(cachedKeys[new[] { 0, 0, 1, 1 }] - 4f), 0f, 0.1f); + Assert.InRange(Math.Abs(cachedValues[new[] { 0, 0, 0, 0 }] - 5f), 0f, 0.1f); + Assert.InRange(Math.Abs(cachedValues[new[] { 0, 0, 1, 1 }] - 8f), 0f, 0.1f); + } +} diff --git a/tests/AiDotNet.Tests/UnitTests/Inference/PagedAttentionTests.cs b/tests/AiDotNet.Tests/UnitTests/Inference/PagedAttentionTests.cs index 26b352b8d6..d929552c0c 100644 --- a/tests/AiDotNet.Tests/UnitTests/Inference/PagedAttentionTests.cs +++ b/tests/AiDotNet.Tests/UnitTests/Inference/PagedAttentionTests.cs @@ -1,4 +1,5 @@ using AiDotNet.Inference.PagedAttention; +using AiDotNet.Inference.Quantization; using Xunit; namespace AiDotNet.Tests.UnitTests.Inference; @@ -777,6 +778,66 @@ public void PagedAttentionKernel_ComputeBatchedAttention_ProcessesMultiple() // Assert Assert.Contains(outputs, v => v != 0); } + + [Fact] + public void PagedAttentionKernel_ForwardQuantized_MatchesFloatWithinTolerance() + { + // Arrange + using var cacheFloat = CreateTestCache(); + cacheFloat.AllocateSequence(1, 1); + var kernelFloat = new PagedAttentionKernel(cacheFloat); + + using var cacheQ = CreateTestCache(); + cacheQ.AllocateSequence(1, 1); + var kernelQ = new PagedAttentionKernel(cacheQ); + + int hiddenDim = kernelFloat.Config.NumHeads * kernelFloat.Config.HeadDimension; + int projDim = hiddenDim; + + var rnd = new Random(42); + var hidden = new float[hiddenDim]; + for (int i = 0; i < hidden.Length; i++) + { + hidden[i] = (float)(rnd.NextDouble() * 0.2 - 0.1); + } + + float[] MakeWeights(int rows, int cols) + { + var w = new float[rows * cols]; + for (int i = 0; i < w.Length; i++) + { + w[i] = (float)(rnd.NextDouble() * 0.02 - 0.01); + } + return w; + } + + var wQ = MakeWeights(projDim, hiddenDim); + var wK = MakeWeights(projDim, hiddenDim); + var wV = MakeWeights(projDim, hiddenDim); + var wO = MakeWeights(hiddenDim, projDim); + + var qWQ = Int8WeightOnlyQuantization.QuantizePerRow(wQ, projDim, hiddenDim); + var qWK = Int8WeightOnlyQuantization.QuantizePerRow(wK, projDim, hiddenDim); + var qWV = Int8WeightOnlyQuantization.QuantizePerRow(wV, projDim, hiddenDim); + var qWO = Int8WeightOnlyQuantization.QuantizePerRow(wO, hiddenDim, projDim); + + var outFloat = new float[hiddenDim]; + var outQ = new float[hiddenDim]; + + // Act + kernelFloat.Forward(hidden, wQ, wK, wV, wO, sequenceId: 1, position: 0, layer: 0, output: outFloat); + kernelQ.ForwardQuantized(hidden, qWQ, qWK, qWV, qWO, sequenceId: 1, position: 0, layer: 0, output: outQ); + + // Assert + float maxAbsDiff = 0f; + for (int i = 0; i < hiddenDim; i++) + { + float diff = MathF.Abs(outFloat[i] - outQ[i]); + if (diff > maxAbsDiff) maxAbsDiff = diff; + } + + Assert.True(maxAbsDiff <= 1e-2f, $"Max abs diff was {maxAbsDiff}"); + } } /// @@ -850,8 +911,14 @@ public void PagedAttentionServer_ForkSequence_ForBeamSearch() Assert.Equal(4, server.GetStats().ActiveSequences); } +#if NET471 + [Fact(Skip = "4GB contiguous allocation exceeds typical .NET Framework single-object limits; validated on net8.0.")] + public void PagedAttentionServer_ForModel_CreatesValidServer() + { + } +#else [Fact] - [Trait("Category", "Integration")] // Skip on net471 - 4GB allocation exceeds .NET Framework array size limits + [Trait("Category", "Integration")] public void PagedAttentionServer_ForModel_CreatesValidServer() { // Act @@ -861,6 +928,7 @@ public void PagedAttentionServer_ForModel_CreatesValidServer() Assert.NotNull(server.KVCache); Assert.NotNull(server.Kernel); } +#endif } /// diff --git a/tests/AiDotNet.Tests/UnitTests/Inference/SpeculativeDecodingTests.cs b/tests/AiDotNet.Tests/UnitTests/Inference/SpeculativeDecodingTests.cs index a584bbbb64..0d0aa8a9e9 100644 --- a/tests/AiDotNet.Tests/UnitTests/Inference/SpeculativeDecodingTests.cs +++ b/tests/AiDotNet.Tests/UnitTests/Inference/SpeculativeDecodingTests.cs @@ -608,4 +608,122 @@ public void TreeSpeculativeConfig_DefaultValues_AreReasonable() Assert.Equal(4, config.MaxDepth); Assert.Equal(16, config.MaxNodes); } + + [Fact] + public async Task SpeculativeDecoder_GenerateAsync_TreeMode_RecordsDraftWork() + { + // Arrange + var draftModel = new NGramDraftModel(ngramSize: 2, vocabSize: 20, seed: 42); + var corpus = new List> + { + new Vector(Enumerable.Range(0, 200).Select(i => i % 10).ToArray()) + }; + draftModel.Train(corpus); + + Func, Matrix> targetForward = tokens => + { + var probs = new Matrix(tokens.Length, 20); + for (int i = 0; i < tokens.Length; i++) + { + for (int v = 0; v < 20; v++) probs[i, v] = 0.05f; + probs[i, 7] = 0.8f; + } + return probs; + }; + + var decoder = new SpeculativeDecoder( + draftModel, + targetForward, + new SpeculativeDecodingConfig + { + UseTreeSpeculation = true, + TreeBranchFactor = 3, + MaxTreeDepth = 3, + Seed = 42 + }); + + // Act + var result = await decoder.GenerateAsync(new Vector(new[] { 1 }), maxNewTokens: 8, temperature: 1.0f); + + // Assert + Assert.True(result.NumGenerated > 0); + Assert.True(decoder.TotalDraftTokens > 0); + Assert.True(result.StepStatistics.Count > 0); + Assert.True(result.TokensPerVerification > 0); + Assert.True(result.StepStatistics.Any(s => s.DraftTokens > 0)); + } + + [Fact] + public async Task SpeculativeDecoder_AdaptiveDraftLength_ReducesDraftTokens_WhenAcceptanceLow() + { + // Arrange + var config = new SpeculativeDecodingConfig + { + NumDraftTokens = 4, + AdaptiveDraftLength = true, + MinAcceptanceRate = 0.8f + }; + + // Target strongly prefers token 2. + Func, Matrix> targetForward = tokens => + { + var probs = new Matrix(tokens.Length, 10); + for (int i = 0; i < tokens.Length; i++) + { + for (int v = 0; v < 10; v++) probs[i, v] = 0.0001f; + probs[i, 2] = 0.999f; + } + return probs; + }; + + // Draft consistently proposes token 1 => low acceptance. + var draft = new DeterministicDraftModel(vocabSize: 10, tokenId: 1); + var decoder = new SpeculativeDecoder(draft, targetForward, config); + + // Act + _ = await decoder.GenerateAsync(new Vector(new[] { 0 }), maxNewTokens: 24, temperature: 1.0f); + + // Assert + Assert.True(decoder.CurrentDraftTokens < 4); + } +} + +internal sealed class DeterministicDraftModel : IDraftModel +{ + private readonly int _vocabSize; + private readonly int _tokenId; + + public DeterministicDraftModel(int vocabSize, int tokenId) + { + _vocabSize = vocabSize; + _tokenId = tokenId; + } + + public int MaxDraftTokens => 32; + + public int VocabSize => _vocabSize; + + public DraftResult GenerateDraft(Vector inputTokens, int numDraftTokens, float temperature) + { + int n = Math.Max(0, numDraftTokens); + var tokens = new Vector(Enumerable.Repeat(_tokenId, n).ToArray()); + + var probs = new Matrix(n, _vocabSize); + for (int i = 0; i < n; i++) + { + probs[i, _tokenId] = 1.0f; + } + + var tokenProbs = new Vector(Enumerable.Repeat(1.0f, n).ToArray()); + return new DraftResult + { + Tokens = tokens, + TokenProbabilities = tokenProbs, + Probabilities = probs + }; + } + + public void Reset() + { + } } diff --git a/tests/AiDotNet.Tests/UnitTests/Serving/ContinuousBatchingTests.cs b/tests/AiDotNet.Tests/UnitTests/Serving/ContinuousBatchingTests.cs index 478323f51d..374d054715 100644 --- a/tests/AiDotNet.Tests/UnitTests/Serving/ContinuousBatchingTests.cs +++ b/tests/AiDotNet.Tests/UnitTests/Serving/ContinuousBatchingTests.cs @@ -1,4 +1,5 @@ using AiDotNet.Serving.ContinuousBatching; +using AiDotNet.Inference.SpeculativeDecoding; using AiDotNet.Tensors.LinearAlgebra; using Xunit; @@ -488,6 +489,290 @@ public async Task ContinuousBatcher_StartStop_Works() Assert.False(isNowRunning); } + [Fact] + public void ContinuousBatcher_SpeculationPolicy_ForceOn_GeneratesMultipleTokensPerStep() + { + // Arrange + var config = new ContinuousBatcherConfig + { + AutoStart = false, + EosTokenId = 2, + EnableSpeculativeDecoding = true, + SpeculationPolicy = AiDotNet.Configuration.SpeculationPolicy.ForceOn, + SpeculationDepth = 3 + }; + + // Target model: always makes token 5 overwhelmingly likely for every position. + Tensor mockModel(Tensor input) + { + var vocabSize = 10; + int seqLen = input.Shape[1]; + var logits = new Tensor(new[] { 1, seqLen, vocabSize }); + for (int pos = 0; pos < seqLen; pos++) + { + for (int i = 0; i < vocabSize; i++) + { + logits[new[] { 0, pos, i }] = i == 5 ? 100f : -100f; + } + } + return logits; + } + + var draft = new DeterministicDraftModel(vocabSize: 10, tokenId: 5); + using var batcher = new ContinuousBatcher(config, mockModel, draftModel: draft); + + var request = new GenerationRequest + { + PromptTokenIds = new List { 1, 2, 3 }, + MaxNewTokens = 10, + Temperature = 1.0f + }; + + var sequence = new SequenceState(request); + + var scheduler = GetSchedulerFromBatcher(batcher); + scheduler.AddSequence(sequence); + + // Act + int tokensGenerated = batcher.Step(); + + // Assert + Assert.True(tokensGenerated > 1); + Assert.True(batcher.LastStepUsedSpeculation); + Assert.True(batcher.LastStepSpeculationTokens > 1); + } + + [Fact] + public void ContinuousBatcher_SpeculationPolicy_ForceOff_DisablesSpeculation() + { + // Arrange + var config = new ContinuousBatcherConfig + { + AutoStart = false, + EnableSpeculativeDecoding = true, + SpeculationPolicy = AiDotNet.Configuration.SpeculationPolicy.ForceOff, + SpeculationDepth = 3 + }; + + Tensor mockModel(Tensor input) + { + var vocabSize = 10; + var logits = new Tensor(new[] { 1, 1, vocabSize }); + logits[new[] { 0, 0, 5 }] = 10f; + return logits; + } + + var draft = new DeterministicDraftModel(vocabSize: 10, tokenId: 5); + using var batcher = new ContinuousBatcher(config, mockModel, draftModel: draft); + + var request = new GenerationRequest + { + PromptTokenIds = new List { 1 }, + MaxNewTokens = 10 + }; + + var sequence = new SequenceState(request); + + var scheduler = GetSchedulerFromBatcher(batcher); + scheduler.AddSequence(sequence); + + // Act + int tokensGenerated = batcher.Step(); + + // Assert - baseline path generates one token per sequence per step + Assert.Equal(1, tokensGenerated); + Assert.False(batcher.LastStepUsedSpeculation); + Assert.Equal(0, batcher.LastStepSpeculationTokens); + } + + [Fact] + public void ContinuousBatcher_SpeculationPolicy_Auto_BacksOff_WhenAcceptanceRateLow() + { + // Arrange + var config = new ContinuousBatcherConfig + { + AutoStart = false, + EosTokenId = 2, + EnableSpeculativeDecoding = true, + SpeculationPolicy = AiDotNet.Configuration.SpeculationPolicy.Auto, + SpeculationDepth = 8, + SchedulerConfig = new BatchSchedulerConfig { MaxBatchSize = 8 } + }; + + // Target model strongly prefers token 5 at every position. + Tensor mockModel(Tensor input) + { + var vocabSize = 10; + int seqLen = input.Shape[1]; + var logits = new Tensor(new[] { 1, seqLen, vocabSize }); + for (int pos = 0; pos < seqLen; pos++) + { + for (int i = 0; i < vocabSize; i++) + { + logits[new[] { 0, pos, i }] = i == 5 ? 100f : -100f; + } + } + return logits; + } + + // Draft always proposes token 4 => low acceptance. + var draft = new DeterministicDraftModel(vocabSize: 10, tokenId: 4); + using var batcher = new ContinuousBatcher(config, mockModel, draftModel: draft); + + var request = new GenerationRequest + { + PromptTokenIds = new List { 1, 2, 3 }, + MaxNewTokens = 64, + Temperature = 1.0f + }; + + var sequence = new SequenceState(request); + var scheduler = GetSchedulerFromBatcher(batcher); + scheduler.AddSequence(sequence); + + // Act: run enough steps to gather acceptance-rate evidence and trigger auto backoff. + bool sawAutoBackoff = false; + for (int i = 0; i < 12; i++) + { + batcher.Step(); + if (!batcher.LastStepUsedSpeculation && + batcher.LastStepSpeculationReason.StartsWith("AutoBackoff(LowAcceptanceRate=")) + { + sawAutoBackoff = true; + break; + } + } + + // Assert + Assert.True(sawAutoBackoff); + } + + [Fact] + public void ContinuousBatcher_SpeculationPolicy_ThroughputFirst_BacksOff_WhenBatchSizeGreaterThanOne() + { + var config = new ContinuousBatcherConfig + { + AutoStart = false, + EosTokenId = 2, + EnableSpeculativeDecoding = true, + SpeculationPolicy = AiDotNet.Configuration.SpeculationPolicy.ThroughputFirst, + SpeculationDepth = 4, + SchedulerConfig = new BatchSchedulerConfig { MaxBatchSize = 4 } + }; + + Tensor mockModel(Tensor input) + { + var vocabSize = 10; + int seqLen = input.Shape[1]; + var logits = new Tensor(new[] { 1, seqLen, vocabSize }); + for (int pos = 0; pos < seqLen; pos++) + { + logits[new[] { 0, pos, 5 }] = 10f; + } + return logits; + } + + var draft = new DeterministicDraftModel(vocabSize: 10, tokenId: 5); + using var batcher = new ContinuousBatcher(config, mockModel, draftModel: draft); + + var scheduler = GetSchedulerFromBatcher(batcher); + scheduler.AddSequence(new SequenceState(new GenerationRequest { PromptTokenIds = new List { 1 }, MaxNewTokens = 10 })); + scheduler.AddSequence(new SequenceState(new GenerationRequest { PromptTokenIds = new List { 1 }, MaxNewTokens = 10 })); + + batcher.Step(); + + Assert.False(batcher.LastStepUsedSpeculation); + Assert.Equal("ThroughputFirst(Backoff)", batcher.LastStepSpeculationReason); + } + + [Fact] + public void ContinuousBatcher_SpeculationPolicy_LatencyFirst_AllowsSpeculation_WithBatchSizeGreaterThanOne() + { + var config = new ContinuousBatcherConfig + { + AutoStart = false, + EosTokenId = 2, + EnableSpeculativeDecoding = true, + SpeculationPolicy = AiDotNet.Configuration.SpeculationPolicy.LatencyFirst, + SpeculationDepth = 4, + SchedulerConfig = new BatchSchedulerConfig { MaxBatchSize = 4 } + }; + + Tensor mockModel(Tensor input) + { + var vocabSize = 10; + int seqLen = input.Shape[1]; + var logits = new Tensor(new[] { 1, seqLen, vocabSize }); + for (int pos = 0; pos < seqLen; pos++) + { + logits[new[] { 0, pos, 5 }] = 10f; + } + return logits; + } + + var draft = new DeterministicDraftModel(vocabSize: 10, tokenId: 5); + using var batcher = new ContinuousBatcher(config, mockModel, draftModel: draft); + + var scheduler = GetSchedulerFromBatcher(batcher); + scheduler.AddSequence(new SequenceState(new GenerationRequest { PromptTokenIds = new List { 1 }, MaxNewTokens = 10 })); + scheduler.AddSequence(new SequenceState(new GenerationRequest { PromptTokenIds = new List { 1 }, MaxNewTokens = 10 })); + + batcher.Step(); + + Assert.True(batcher.LastStepUsedSpeculation); + Assert.DoesNotContain("Backoff", batcher.LastStepSpeculationReason); + } + + [Fact] + public void ContinuousBatcher_SpeculativeDecoding_DisablesAfterFailure() + { + // Arrange + var config = new ContinuousBatcherConfig + { + AutoStart = false, + EnableSpeculativeDecoding = true, + SpeculationPolicy = AiDotNet.Configuration.SpeculationPolicy.ForceOn, + SpeculationDepth = 3 + }; + + Tensor mockModel(Tensor input) + { + var vocabSize = 10; + var logits = new Tensor(new[] { 1, 1, vocabSize }); + logits[new[] { 0, 0, 5 }] = 10f; + return logits; + } + + var throwingDraft = new ThrowingDraftModel(vocabSize: 10); + using var batcher = new ContinuousBatcher(config, mockModel, draftModel: throwingDraft); + + var request = new GenerationRequest + { + PromptTokenIds = new List { 1 }, + MaxNewTokens = 10 + }; + + var sequence = new SequenceState(request); + + var scheduler = GetSchedulerFromBatcher(batcher); + scheduler.AddSequence(sequence); + + // Act + int tokensGeneratedFirst = batcher.Step(); + bool usedSpeculationFirst = batcher.LastStepUsedSpeculation; + + int tokensGeneratedSecond = batcher.Step(); + bool usedSpeculationSecond = batcher.LastStepUsedSpeculation; + + // Assert + Assert.Equal(1, tokensGeneratedFirst); // falls back to baseline + Assert.True(usedSpeculationFirst); // ForceOn decision, even though it failed internally + + Assert.Equal(1, tokensGeneratedSecond); + Assert.False(usedSpeculationSecond); // disabled after failure + Assert.Equal("DisabledDueToFailure", batcher.LastStepSpeculationReason); + } + [Fact] public async Task ContinuousBatcher_GenerateAsync_ReturnsCancellableTask() { @@ -625,4 +910,66 @@ private static BatchScheduler GetSchedulerFromBatcher(ContinuousBatcher } #endregion + + private sealed class DeterministicDraftModel : IDraftModel + { + public int MaxDraftTokens => 16; + public int VocabSize { get; } + + private readonly int _tokenId; + + public DeterministicDraftModel(int vocabSize, int tokenId) + { + VocabSize = vocabSize; + _tokenId = tokenId; + } + + public DraftResult GenerateDraft(Vector inputTokens, int numDraftTokens, float temperature) + { + var tokens = new Vector(numDraftTokens); + var tokenProbs = new Vector(numDraftTokens); + var probs = new Matrix(numDraftTokens, VocabSize); + + for (int i = 0; i < numDraftTokens; i++) + { + tokens[i] = _tokenId; + tokenProbs[i] = 1.0f; + for (int v = 0; v < VocabSize; v++) + { + probs[i, v] = v == _tokenId ? 1.0f : 0.0f; + } + } + + return new DraftResult + { + Tokens = tokens, + TokenProbabilities = tokenProbs, + Probabilities = probs + }; + } + + public void Reset() + { + } + } + + private sealed class ThrowingDraftModel : IDraftModel + { + public int MaxDraftTokens => 16; + public int VocabSize { get; } + + public ThrowingDraftModel(int vocabSize) + { + VocabSize = vocabSize; + } + + public DraftResult GenerateDraft(Vector inputTokens, int numDraftTokens, float temperature) + { + throw new InvalidOperationException("Draft model failure (test)."); + } + + public void Reset() + { + } + } }