|
12 | 12 | from torchjd.sparse._aten_function_overrides.shape import unsquash_pdim |
13 | 13 | from torchjd.sparse._structured_sparse_tensor import ( |
14 | 14 | StructuredSparseTensor, |
| 15 | + clear_null_stride_columns, |
15 | 16 | encode_by_order, |
16 | 17 | fix_ungrouped_dims, |
17 | 18 | get_groupings, |
@@ -277,3 +278,31 @@ def test_concatenate( |
277 | 278 |
|
278 | 279 | assert isinstance(res, StructuredSparseTensor) |
279 | 280 | assert torch.all(torch.eq(res.to_dense(), expected)) |
| 281 | + |
| 282 | + |
| 283 | +@mark.parametrize( |
| 284 | + ["physical", "strides", "expected_physical", "expected_strides"], |
| 285 | + [ |
| 286 | + ([[1, 2, 3], [4, 5, 6]], [[1, 0], [1, 0], [2, 0]], [6, 15], [[1], [1], [2]]), |
| 287 | + ( |
| 288 | + [[1, 2, 3], [4, 5, 6]], |
| 289 | + [[1, 1], [1, 0], [2, 0]], |
| 290 | + [[1, 2, 3], [4, 5, 6]], |
| 291 | + [[1, 1], [1, 0], [2, 0]], |
| 292 | + ), |
| 293 | + ], |
| 294 | +) |
| 295 | +def test_clear_null_stride_columns( |
| 296 | + physical: list, |
| 297 | + strides: list, |
| 298 | + expected_physical: list, |
| 299 | + expected_strides: list, |
| 300 | +): |
| 301 | + physical, strides = torch.tensor(physical), torch.tensor(strides) |
| 302 | + expected_physical, expected_strides = torch.tensor(expected_physical), torch.tensor( |
| 303 | + expected_strides |
| 304 | + ) |
| 305 | + |
| 306 | + physical, strides = clear_null_stride_columns(physical, strides) |
| 307 | + assert_close(physical, expected_physical) |
| 308 | + assert_close(strides, expected_strides) |
0 commit comments