diff --git a/src/Numerics/LinearAlgebra/Storage/SparseCompressedRowMatrixStorage.cs b/src/Numerics/LinearAlgebra/Storage/SparseCompressedRowMatrixStorage.cs index 7d6d63da..84d29680 100644 --- a/src/Numerics/LinearAlgebra/Storage/SparseCompressedRowMatrixStorage.cs +++ b/src/Numerics/LinearAlgebra/Storage/SparseCompressedRowMatrixStorage.cs @@ -36,7 +36,7 @@ namespace MathNet.Numerics.LinearAlgebra.Storage /// The number of non zero elements. public int ValueCount; - internal SparseCompressedRowMatrixStorage(int rows, int columns, T zero) + internal SparseCompressedRowMatrixStorage(int rows, int columns, T zero = default(T)) : base(rows, columns) { _zero = zero; @@ -81,7 +81,7 @@ namespace MathNet.Numerics.LinearAlgebra.Storage if (_zero.Equals(value)) { // Delete existing item - DeleteItemByIndex(index, row); + RemoveAtIndexUnchecked(index, row); } else { @@ -135,46 +135,13 @@ namespace MathNet.Numerics.LinearAlgebra.Storage } } - public override void Clear() - { - ValueCount = 0; - Array.Clear(RowPointers, 0, RowPointers.Length); - } - - public override void Clear(int rowIndex, int rowCount, int columnIndex, int columnCount) - { - if (rowIndex == 0 && columnIndex == 0 && rowCount == RowCount && columnCount == ColumnCount) - { - Clear(); - return; - } - - for (int i = rowIndex + rowCount - 1, row = rowCount - 1; i >= rowIndex; i--, row--) - { - var startIndex = RowPointers[i]; - var endIndex = i < RowPointers.Length - 1 ? RowPointers[i + 1] : ValueCount; - - for (int j = endIndex - 1; j >= startIndex; j--) - { - // check if the column index is in the range - if ((ColumnIndices[j] >= columnIndex) && (ColumnIndices[j] < columnIndex + columnCount)) - { - var column = ColumnIndices[j]; - - // NOTE: potential for more efficient implementation - At(row, column, _zero); - } - } - } - } - /// /// Delete value from internal storage /// /// Index of value in nonZeroValues array /// Row number of matrix /// WARNING: This method is not thread safe. Use "lock" with it and be sure to avoid deadlocks - void DeleteItemByIndex(int itemIndex, int row) + void RemoveAtIndexUnchecked(int itemIndex, int row) { // Move all values (with an position larger than index) in the value array to the previous position // move all values (with an position larger than index) in the columIndices array to the previous position @@ -240,6 +207,64 @@ namespace MathNet.Numerics.LinearAlgebra.Storage return delta; } + public override void Clear() + { + ValueCount = 0; + Array.Clear(RowPointers, 0, RowPointers.Length); + } + + public override void Clear(int rowIndex, int rowCount, int columnIndex, int columnCount) + { + if (rowIndex == 0 && columnIndex == 0 && rowCount == RowCount && columnCount == ColumnCount) + { + Clear(); + return; + } + + for (int row = rowIndex + rowCount - 1; row >= rowIndex; row--) + { + var startIndex = RowPointers[row]; + var endIndex = row < RowPointers.Length - 1 ? RowPointers[row + 1] : ValueCount; + + // empty row + if (startIndex == endIndex) + { + continue; + } + + // multiple entries in row + var first = Array.BinarySearch(ColumnIndices, startIndex, endIndex - startIndex, columnIndex); + var last = Array.BinarySearch(ColumnIndices, startIndex, endIndex - startIndex, columnIndex + columnCount - 1); + if (first < 0) first = ~first; + if (last < 0) last = ~last - 1; + int count = last - first + 1; + + if (count > 0) + { + // Move all values (with an position larger than index) in the value array to the previous position + // move all values (with an position larger than index) in the columIndices array to the previous position + Array.Copy(Values, first + count, Values, first, ValueCount - first - count); + Array.Copy(ColumnIndices, first + count, ColumnIndices, first, ValueCount - first - count); + + // Decrease value in Row + for (var k = row + 1; k < RowPointers.Length; k++) + { + RowPointers[k] -= count; + } + + ValueCount -= count; + } + } + + // Check if the storage needs to be shrink. This is reasonable to do if + // there are a lot of non-zero elements and storage is two times bigger + if ((ValueCount > 1024) && (ValueCount < Values.Length / 2)) + { + Array.Resize(ref Values, ValueCount); + Array.Resize(ref ColumnIndices, ValueCount); + } + } + /// /// Indicates whether the current object is equal to another object of the same type. /// @@ -513,13 +538,12 @@ namespace MathNet.Numerics.LinearAlgebra.Storage return; } - // NOTE: potential for more efficient implementation - if (!skipClearing) { target.Clear(targetRowIndex, rowCount, targetColumnIndex, columnCount); } + // NOTE: potential for more efficient implementation for (int i = sourceRowIndex, row = 0; i < sourceRowIndex + rowCount; i++, row++) { var startIndex = RowPointers[i]; diff --git a/src/UnitTests/LinearAlgebraTests/Double/SparseMatrixTests.cs b/src/UnitTests/LinearAlgebraTests/Double/SparseMatrixTests.cs index 6396f95d..c5dc5900 100644 --- a/src/UnitTests/LinearAlgebraTests/Double/SparseMatrixTests.cs +++ b/src/UnitTests/LinearAlgebraTests/Double/SparseMatrixTests.cs @@ -29,6 +29,7 @@ namespace MathNet.Numerics.UnitTests.LinearAlgebraTests.Double using System; using System.Collections.Generic; using LinearAlgebra.Double; + using LinearAlgebra.Double.IO; using NUnit.Framework; /// @@ -315,5 +316,29 @@ namespace MathNet.Numerics.UnitTests.LinearAlgebraTests.Double Assert.AreEqual(Order, matrix.ColumnCount); Assert.DoesNotThrow(() => matrix[0, 0] = 1); } + + [Test] + public void CanClearSubMatrixEx() + { + var dmr = new MatlabMatrixReader("./data/Matlab/sparse-small.mat"); + var matrix = dmr.ReadMatrix("S"); + var matrix2 = matrix.Clone(); + + // Zero the 4th column + for (int i = 0; i < matrix.RowCount; i++) + { + matrix.At(i, 3, 0.0); + } + matrix2.ClearColumn(3); + + // Zero the 40th row + for (int i = 0; i < matrix.ColumnCount; i++) + { + matrix.At(39, i, 0.0); + } + matrix2.ClearRow(39); + + Assert.That(matrix2.Equals(matrix), Is.True); + } } }