Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 3 additions & 3 deletions src/TMG-Framework/Data/Categories.cs
Original file line number Diff line number Diff line change
Expand Up @@ -80,10 +80,10 @@ private Categories(List<int> elements)
}

/// <summary>
///
/// Gives the flat index of the specified sparse index, or less than 0 if the sparse index is not in the map.
/// </summary>
/// <param name="sparseIndex"></param>
/// <returns></returns>
/// <param name="sparseIndex">The sparse index to look up</param>
/// <returns>Gives the flat index of the specified sparse index, or less than 0 if the sparse index is not in the map.</returns>
public int GetFlatIndex(CategoryIndex sparseIndex)
{
return _elements.BinarySearch(sparseIndex);
Expand Down
108 changes: 91 additions & 17 deletions src/TMG-Framework/Loading/LoadMatrixFromMTX.cs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -33,6 +34,9 @@ public sealed class LoadMatrixFromMTX : BaseFunction<ReadStream, Matrix>
[SubModule(Required = true, Name = "Map", Description = "The sparse map this vector will be shaped in.", Index = 0)]
public IFunction<Categories> 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<bool> ConvertBetweenZoneSystems = null!;

private const uint MagicNumber = 0xC4D4F1B2;

private const int FloatType = 0x1;
Expand Down Expand Up @@ -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<float, byte>(data)[soFar..dataSize]);
if (amount == 0)
{
var amount = reader.Read(MemoryMarshal.Cast<float,byte>(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<int, byte>(rows.AsSpan()));
reader.ReadExactly(MemoryMarshal.Cast<int, byte>(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<float, byte>(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)
Expand Down
42 changes: 42 additions & 0 deletions tests/TMG-Framework.Test/Loading/TestLoadMatrix.cs
Original file line number Diff line number Diff line change
Expand Up @@ -131,6 +131,7 @@ public void TestLoadMatrixFromMTX()
var matrixLoader = new LoadMatrixFromMTX()
{
Categories = Helper.CreateParameter(map),
ConvertBetweenZoneSystems = Helper.CreateParameter(false)
};
using (var stream = (new OpenReadStreamFromFile()
{
Expand All @@ -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);
}
}
}
}
}
}
}
3 changes: 2 additions & 1 deletion tests/TMG-Framework.Test/Saving/TestSaveMatrixAsMTX.cs
Original file line number Diff line number Diff line change
Expand Up @@ -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);
Expand Down
Loading