diff --git a/src/TMG-Framework/Data/Categories.cs b/src/TMG-Framework/Data/Categories.cs index a7775b1..ceb3d23 100644 --- a/src/TMG-Framework/Data/Categories.cs +++ b/src/TMG-Framework/Data/Categories.cs @@ -80,10 +80,10 @@ private Categories(List elements) } /// - /// + /// Gives the flat index of the specified sparse index, or less than 0 if the sparse index is not in the map. /// - /// - /// + /// The sparse index to look up + /// Gives the flat index of the specified sparse index, or less than 0 if the sparse index is not in the map. public int GetFlatIndex(CategoryIndex sparseIndex) { return _elements.BinarySearch(sparseIndex); diff --git a/src/TMG-Framework/Loading/LoadMatrixFromMTX.cs b/src/TMG-Framework/Loading/LoadMatrixFromMTX.cs index 243d62f..0967881 100644 --- a/src/TMG-Framework/Loading/LoadMatrixFromMTX.cs +++ b/src/TMG-Framework/Loading/LoadMatrixFromMTX.cs @@ -19,6 +19,7 @@ You should have received a copy of the GNU General Public License using System; using System.Collections.Generic; using System.IO; +using System.Runtime.CompilerServices; using System.Runtime.InteropServices; using System.Text; using TMG.Utilities; @@ -33,6 +34,9 @@ public sealed class LoadMatrixFromMTX : BaseFunction [SubModule(Required = true, Name = "Map", Description = "The sparse map this vector will be shaped in.", Index = 0)] public IFunction Categories = null!; + [Parameter(Name = "Convert Between Zone Systems", DefaultValue = "false", Description = "A function that converts between the zone system of the matrix and the zone system of the map.", Index = 1)] + public IFunction ConvertBetweenZoneSystems = null!; + private const uint MagicNumber = 0xC4D4F1B2; private const int FloatType = 0x1; @@ -60,32 +64,102 @@ public override Matrix Invoke(ReadStream context) { throw new XTMFRuntimeException(this, $"The matrix contained {numberOfIndexes} dimensions!"); } - int firstSize = reader.ReadInt32(); - int secondSize = reader.ReadInt32(); - if(categories.Count != firstSize) + + var convert = ConvertBetweenZoneSystems.Invoke(); + if(!convert) { - throw new XTMFRuntimeException(this, "The matrix had the wrong number of elements in the first dimension!"); + LoadWithoutConversion(categories, matrix, reader); } - if (categories.Count != secondSize) + else { - throw new XTMFRuntimeException(this, "The matrix had the wrong number of elements in the second dimension!"); + LoadWithConversion(categories, matrix, reader); } - ValidateIndexes(reader, categories); - ValidateIndexes(reader, categories); - var data = matrix.Data; - var dataSize = data.Length * sizeof(float); - var soFar = 0; - while (soFar < dataSize) + } + return matrix; + } + + private void LoadWithoutConversion(Categories categories, Matrix matrix, BinaryReader reader) + { + int firstSize = reader.ReadInt32(); + int secondSize = reader.ReadInt32(); + if (categories.Count != firstSize) + { + throw new XTMFRuntimeException(this, "The matrix had the wrong number of elements in the first dimension!"); + } + if (categories.Count != secondSize) + { + throw new XTMFRuntimeException(this, "The matrix had the wrong number of elements in the second dimension!"); + } + ValidateIndexes(reader, categories); + ValidateIndexes(reader, categories); + + var data = matrix.Data; + var dataSize = data.Length * sizeof(float); + var soFar = 0; + while (soFar < dataSize) + { + var amount = reader.Read(MemoryMarshal.Cast(data)[soFar..dataSize]); + if (amount == 0) { - var amount = reader.Read(MemoryMarshal.Cast(data)[soFar..dataSize]); - if(amount == 0) + throw new XTMFRuntimeException(this, $"The matrix expected {dataSize}bytes but we only could get {soFar}bytes!"); + } + soFar += amount; + } + } + + private void LoadWithConversion(Categories categories, Matrix matrix, BinaryReader reader) + { + int rowSize = reader.ReadInt32(); + int columnSize = reader.ReadInt32(); + + if (rowSize != columnSize) + { + throw new XTMFRuntimeException(this, "The matrix was not square!"); + } + + // Load in the column categories (sparse space) + var rows = new int[rowSize]; + var columns = new int[columnSize]; + + reader.ReadExactly(MemoryMarshal.Cast(rows.AsSpan())); + reader.ReadExactly(MemoryMarshal.Cast(columns.AsSpan())); + + // Load in the matrix data + var numberOfElements = rowSize * columnSize; + var dataSize = numberOfElements * sizeof(float); + var data = new float[numberOfElements]; + var dataSpan = data.AsSpan(); + var soFar = 0; + while (soFar < dataSize) + { + var amount = reader.Read(MemoryMarshal.Cast(dataSpan)[soFar..dataSize]); + if (amount == 0) + { + throw new XTMFRuntimeException(this, $"The matrix expected {dataSize}bytes but we only could get {soFar}bytes!"); + } + soFar += amount; + } + + ref var matrixData = ref MemoryMarshal.GetReference(matrix.Data); + ref var rData = ref MemoryMarshal.GetReference(dataSpan); + var matrixColumnSize = matrix.NumberOfColumns; + for (int i = 0; i < rowSize; i++) + { + var rowIndex = categories.GetFlatIndex(rows[i]); + if (rowIndex < 0) + { + continue; + } + for (int j = 0; j < columnSize; j++) + { + var columnIndex = categories.GetFlatIndex(columns[j]); + if (columnIndex >= 0) { - throw new XTMFRuntimeException(this, $"The matrix expected {dataSize}bytes but we only could get {soFar}bytes!"); + ref var writeTo = ref Unsafe.Add(ref matrixData, rowIndex * matrixColumnSize + columnIndex); + writeTo = Unsafe.Add(ref rData, i * columnSize + j); } - soFar += amount; } } - return matrix; } private void ValidateIndexes(BinaryReader reader, Categories categories) diff --git a/tests/TMG-Framework.Test/Loading/TestLoadMatrix.cs b/tests/TMG-Framework.Test/Loading/TestLoadMatrix.cs index 26e8fe7..d2dfcc3 100644 --- a/tests/TMG-Framework.Test/Loading/TestLoadMatrix.cs +++ b/tests/TMG-Framework.Test/Loading/TestLoadMatrix.cs @@ -131,6 +131,7 @@ public void TestLoadMatrixFromMTX() var matrixLoader = new LoadMatrixFromMTX() { Categories = Helper.CreateParameter(map), + ConvertBetweenZoneSystems = Helper.CreateParameter(false) }; using (var stream = (new OpenReadStreamFromFile() { @@ -152,5 +153,46 @@ public void TestLoadMatrixFromMTX() } } } + + [TestMethod] + public void TestLoadMatrixFromMTXDifferentZoneSystem() + { + var bigMap = MapHelper.LoadMap(MapHelper.WriteCSV(64)); + var smallMap = MapHelper.LoadMap(MapHelper.WriteCSV(32)); + float[][] data = new float[64][]; + for (int i = 0; i < data.Length; i++) + { + data[i] = new float[64]; + for (int j = 0; j < data[i].Length; j++) + { + data[i][j] = 2.0f + i * j; + } + } + var matrixFileName = MatrixHelper.WriteMatrixToMTX(bigMap, data); + var matrixLoader = new LoadMatrixFromMTX() + { + Categories = Helper.CreateParameter(smallMap), + ConvertBetweenZoneSystems = Helper.CreateParameter(true) + }; + using (var stream = (new OpenReadStreamFromFile() + { + FilePath = Helper.CreateParameter(matrixFileName) + }).Invoke()) + { + var matrix = matrixLoader.Invoke(stream); + Assert.AreSame(smallMap, matrix.RowCategories); + var vData = matrix.Data; + for (int i = 0; i < smallMap.Count; i++) + { + for (int j = 0; j < smallMap.Count; j++) + { + if (Math.Abs((2.0f + i * j) - vData[i * smallMap.Count + j]) < 0.00001f) + { + Assert.AreEqual(2.0f + i * j, vData[i * smallMap.Count + j], 0.00001f); + } + } + } + } + } } } diff --git a/tests/TMG-Framework.Test/Saving/TestSaveMatrixAsMTX.cs b/tests/TMG-Framework.Test/Saving/TestSaveMatrixAsMTX.cs index db30cb4..64bebba 100644 --- a/tests/TMG-Framework.Test/Saving/TestSaveMatrixAsMTX.cs +++ b/tests/TMG-Framework.Test/Saving/TestSaveMatrixAsMTX.cs @@ -70,7 +70,8 @@ public void SaveMatrixAsMTXAsMatrixAndLoad() { var readMatrix = new TMG.Loading.LoadMatrixFromMTX() { - Categories = Helper.CreateParameter(categories) + Categories = Helper.CreateParameter(categories), + ConvertBetweenZoneSystems = Helper.CreateParameter(false) }.Invoke(readStream); string? error = null; Assert.IsTrue(MatrixHelper.Compare(a, readMatrix, ref error), error);