SIMD Refactor: Merge simd-dev into dev (#55)

Co-authored-by: Claude Opus 4.5 <noreply@anthropic.com>
Co-authored-by: aider (openrouter/anthropic/claude-sonnet-4) <aider@aider.chat>
Co-authored-by: Warp <agent@warp.dev>
This commit is contained in:
Miha Kralj
2026-01-18 19:02:03 -08:00
committed by GitHub
co-authored by Claude Opus 4.5 aider Warp
parent 5bcdf8d614
commit 86fe32a682
1750 changed files with 198235 additions and 80539 deletions
+250
View File
@@ -0,0 +1,250 @@
namespace QuanTAlib.Tests;
public class RmseTests
{
[Fact]
public void Constructor_ValidatesInput()
{
Assert.Throws<ArgumentException>(() => new Rmse(0));
Assert.Throws<ArgumentException>(() => new Rmse(-1));
var rmse = new Rmse(10);
Assert.NotNull(rmse);
}
[Fact]
public void Properties_Accessible()
{
var rmse = new Rmse(10);
Assert.Equal(0, rmse.Last.Value);
Assert.False(rmse.IsHot);
Assert.Contains("Rmse", rmse.Name, StringComparison.Ordinal);
rmse.Update(100, 105);
Assert.NotEqual(0, rmse.Last.Time);
}
[Fact]
public void IsHot_BecomesTrueWhenBufferFull()
{
const int period = 5;
var rmse = new Rmse(period);
for (int i = 0; i < period - 1; i++)
{
Assert.False(rmse.IsHot);
rmse.Update(i * 10, i * 10 + 5);
}
rmse.Update((period - 1) * 10, (period - 1) * 10 + 5);
Assert.True(rmse.IsHot);
}
[Fact]
public void Rmse_CalculatesCorrectly()
{
var rmse = new Rmse(3);
// (10 - 15)² = 25, RMSE = √25 = 5
var res1 = rmse.Update(10, 15);
Assert.Equal(5.0, res1.Value, 10);
// (20 - 30)² = 100, MSE = (25 + 100) / 2 = 62.5, RMSE = √62.5
var res2 = rmse.Update(20, 30);
Assert.Equal(Math.Sqrt(62.5), res2.Value, 10);
// (30 - 25)² = 25, MSE = (25 + 100 + 25) / 3 = 50, RMSE = √50
var res3 = rmse.Update(30, 25);
Assert.Equal(Math.Sqrt(50.0), res3.Value, 10);
}
[Fact]
public void Rmse_IsSqrtOfMse()
{
var rmse = new Rmse(5);
var mse = new Mse(5);
for (int i = 0; i < 20; i++)
{
rmse.Update(i * 10, i * 10 + 7);
mse.Update(i * 10, i * 10 + 7);
}
Assert.Equal(Math.Sqrt(mse.Last.Value), rmse.Last.Value, 10);
}
[Fact]
public void Rmse_PerfectPrediction_ReturnsZero()
{
var rmse = new Rmse(5);
for (int i = 0; i < 10; i++)
{
rmse.Update(i * 10, i * 10);
}
Assert.Equal(0.0, rmse.Last.Value, 10);
}
[Fact]
public void Rmse_ConstantError_ReturnsSameAsError()
{
var rmse = new Rmse(5);
for (int i = 0; i < 10; i++)
{
rmse.Update(100, 110); // Constant error of 10
}
// MSE = 100, RMSE = √100 = 10 (same as error because error is constant)
Assert.Equal(10.0, rmse.Last.Value, 10);
}
[Fact]
public void Calc_IsNew_False_UpdatesValue()
{
var rmse = new Rmse(10);
rmse.Update(100, 110);
rmse.Update(100, 120, isNew: true);
double beforeUpdate = rmse.Last.Value;
rmse.Update(100, 130, isNew: false);
double afterUpdate = rmse.Last.Value;
Assert.NotEqual(beforeUpdate, afterUpdate);
}
[Fact]
public void IterativeCorrections_RestoreToOriginalState()
{
var rmse = new Rmse(5);
double tenthActual = 0;
double tenthPredicted = 0;
for (int i = 0; i < 10; i++)
{
tenthActual = i * 10;
tenthPredicted = i * 10 + 5;
rmse.Update(tenthActual, tenthPredicted);
}
double stateAfterTen = rmse.Last.Value;
for (int i = 0; i < 5; i++)
{
rmse.Update(100 + i, 200 + i, isNew: false);
}
rmse.Update(tenthActual, tenthPredicted, isNew: false);
Assert.Equal(stateAfterTen, rmse.Last.Value, 10);
}
[Fact]
public void Reset_ClearsState()
{
var rmse = new Rmse(5);
for (int i = 0; i < 10; i++)
{
rmse.Update(i * 10, i * 10 + 5);
}
Assert.True(rmse.IsHot);
rmse.Reset();
Assert.False(rmse.IsHot);
Assert.Equal(0, rmse.Last.Value);
}
[Fact]
public void NaN_Input_UsesLastValidValue()
{
var rmse = new Rmse(5);
rmse.Update(100, 110);
rmse.Update(110, 120);
var result = rmse.Update(double.NaN, double.NaN);
Assert.True(double.IsFinite(result.Value));
}
[Fact]
public void Rmse_Throws_On_Single_Input()
{
var rmse = new Rmse(10);
Assert.Throws<NotSupportedException>(() => rmse.Update(new TValue(DateTime.UtcNow, 1)));
Assert.Throws<NotSupportedException>(() => rmse.Update(new TSeries()));
Assert.Throws<NotSupportedException>(() => rmse.Prime([1, 2, 3]));
}
[Fact]
public void BatchSpan_MatchesStreaming()
{
int period = 5;
int count = 100;
var gbm = new GBM(startPrice: 100, mu: 0.05, sigma: 0.2, seed: 123);
double[] actual = new double[count];
double[] predicted = new double[count];
for (int i = 0; i < count; i++)
{
var bar = gbm.Next();
actual[i] = bar.Close;
predicted[i] = bar.Close * 1.05 + 2;
}
var rmse = new Rmse(period);
var streamingResults = new double[count];
for (int i = 0; i < count; i++)
{
streamingResults[i] = rmse.Update(actual[i], predicted[i]).Value;
}
double[] batchResults = new double[count];
Rmse.Batch(actual, predicted, batchResults, period);
for (int i = 0; i < count; i++)
{
Assert.Equal(streamingResults[i], batchResults[i], 9);
}
}
[Fact]
public void BatchSpan_ValidatesInput()
{
double[] actual = [1, 2, 3, 4, 5];
double[] predicted = [1, 2, 3, 4, 5];
double[] output = new double[5];
Assert.Throws<ArgumentException>(() =>
Rmse.Batch(actual.AsSpan(), predicted.AsSpan(), output.AsSpan(), 0));
Assert.Throws<ArgumentException>(() =>
Rmse.Batch(actual.AsSpan(), predicted.AsSpan(), new double[3].AsSpan(), 3));
}
[Fact]
public void Calculate_Works()
{
var actual = new TSeries();
var predicted = new TSeries();
var now = DateTime.UtcNow;
for (int i = 0; i < 10; i++)
{
actual.Add(now.AddMinutes(i), i * 10);
predicted.Add(now.AddMinutes(i), i * 10 + 5);
}
var results = Rmse.Calculate(actual, predicted, 3);
Assert.Equal(10, results.Count);
// All errors are 5, MSE = 25, RMSE = 5
Assert.Equal(5.0, results.Last.Value, 10);
}
}
+76
View File
@@ -0,0 +1,76 @@
using MathNet.Numerics;
using QuanTAlib.Tests;
namespace QuanTAlib.Validation;
public sealed class RmseValidationTests : IDisposable
{
private readonly ValidationTestData _data = new();
public void Dispose() => _data.Dispose();
[Fact]
public void Rmse_Matches_MathNet()
{
int[] periods = { 5, 10, 20, 50, 100 };
var quotes = _data.SkenderQuotes.ToList();
double[] actual = quotes.Select(q => (double)q.Close).ToArray();
double[] predicted = quotes.Select(q => (double)q.Open).ToArray();
foreach (int period in periods)
{
var rmse = new Rmse(period);
for (int i = 0; i < actual.Length; i++)
{
var val = rmse.Update(
new TValue(quotes[i].Date, actual[i]),
new TValue(quotes[i].Date, predicted[i]));
// Validate last 100 bars
if (i >= actual.Length - 100 && i >= period - 1)
{
var windowActual = actual[(i - period + 1)..(i + 1)];
var windowPredicted = predicted[(i - period + 1)..(i + 1)];
// RMSE = sqrt(MSE)
double expected = Math.Sqrt(Distance.MSE(windowActual, windowPredicted));
Assert.Equal(expected, val.Value, 1e-9);
}
}
}
}
[Fact]
public void Rmse_Batch_Matches_MathNet()
{
int[] periods = { 5, 10, 20, 50, 100 };
var quotes = _data.SkenderQuotes.ToList();
double[] actual = quotes.Select(q => (double)q.Close).ToArray();
double[] predicted = quotes.Select(q => (double)q.Open).ToArray();
foreach (int period in periods)
{
double[] output = new double[actual.Length];
Rmse.Batch(actual, predicted, output, period);
// Validate last 100 bars
for (int i = actual.Length - 100; i < actual.Length; i++)
{
if (i >= period - 1)
{
var windowActual = actual[(i - period + 1)..(i + 1)];
var windowPredicted = predicted[(i - period + 1)..(i + 1)];
// RMSE = sqrt(MSE)
double expected = Math.Sqrt(Distance.MSE(windowActual, windowPredicted));
Assert.Equal(expected, output[i], 1e-9);
}
}
}
}
}
+70
View File
@@ -0,0 +1,70 @@
using System.Runtime.CompilerServices;
namespace QuanTAlib;
/// <summary>
/// RMSE: Root Mean Squared Error
/// </summary>
/// <remarks>
/// RMSE is the square root of MSE, bringing the error metric back to the
/// original units of the data while retaining the outlier sensitivity
/// of squared errors.
///
/// Formula:
/// RMSE = √((1/n) * Σ(actual - predicted)²) = √MSE
///
/// Uses a RingBuffer for O(1) streaming updates with running sum.
///
/// Key properties:
/// - Always non-negative (RMSE ≥ 0)
/// - Same units as the original data
/// - Heavily penalizes outliers due to squaring before averaging
/// - RMSE = 0 indicates perfect prediction
/// </remarks>
[SkipLocalsInit]
public sealed class Rmse : BiInputIndicatorBase
{
/// <summary>
/// Creates RMSE with specified period.
/// </summary>
/// <param name="period">Number of values to average (must be > 0)</param>
public Rmse(int period) : base(period, $"Rmse({period})") { }
/// <inheritdoc/>
[MethodImpl(MethodImplOptions.AggressiveInlining)]
protected override double ComputeError(double actual, double predicted)
{
double diff = actual - predicted;
return diff * diff;
}
/// <inheritdoc/>
[MethodImpl(MethodImplOptions.AggressiveInlining)]
protected override double PostProcess(double mean) => Math.Sqrt(mean);
/// <summary>
/// Calculates RMSE for entire series.
/// </summary>
public static TSeries Calculate(TSeries actual, TSeries predicted, int period)
=> CalculateImpl(actual, predicted, period, Batch);
/// <summary>
/// Batch calculation using SIMD-accelerated squared error computation with sqrt of rolling mean.
/// </summary>
[MethodImpl(MethodImplOptions.AggressiveInlining)]
public static void Batch(ReadOnlySpan<double> actual, ReadOnlySpan<double> predicted, Span<double> output, int period)
{
ValidateBatchInputs(actual, predicted, output, period);
int len = actual.Length;
if (len == 0) return;
const int StackAllocThreshold = 256;
Span<double> sqErrors = len <= StackAllocThreshold
? stackalloc double[len]
: new double[len];
ErrorHelpers.ComputeSquaredErrors(actual, predicted, sqErrors);
ErrorHelpers.ApplyRollingMeanSqrt(sqErrors, output, period);
}
}
+41
View File
@@ -0,0 +1,41 @@
# RMSE: Root Mean Squared Error
> "MSE's more interpretable sibling that speaks the language of your data."
Root Mean Squared Error (RMSE) is the square root of MSE, providing an error metric in the same units as the original data while retaining sensitivity to large errors.
## Mathematical Foundation
### Formula
$$RMSE = \sqrt{\frac{1}{n} \sum_{i=1}^{n} (y_i - \hat{y}_i)^2} = \sqrt{MSE}$$
## Properties
* **Non-negative**: RMSE ≥ 0
* **Same units**: Unlike MSE, RMSE is in original data units
* **Outlier sensitive**: Inherits MSE's penalty for large errors
* **Always ≥ MAE**: RMSE ≥ MAE due to Jensen's inequality
## Usage
```csharp
var rmse = new Rmse(period: 20);
var result = rmse.Update(actualValue, predictedValue);
// Batch calculation
var results = Rmse.Calculate(actualSeries, predictedSeries, period: 20);
```
## Performance Profile
| Metric | Score | Notes |
| :--- | :--- | :--- |
| **Throughput** | ~15 ns/bar | O(1) with sqrt operation |
| **Allocations** | 0 | Pre-allocated ring buffer |
| **Complexity** | O(1) | Constant time per update |
## Related Indicators
* [MSE](../mse/Mse.md) - Mean Squared Error
* [MAE](../mae/Mae.md) - Mean Absolute Error