@@ -1496,10 +1496,46 @@ def test_ndim(self, xp: ModuleType):
14961496 padded = pad (a , 2 )
14971497 assert padded .shape == (6 , 7 , 8 )
14981498
1499+ def test_edge (self , xp : ModuleType ):
1500+ a = xp .asarray ([1 , 2 , 3 ])
1501+ padded = pad (a , (2 , 1 ), mode = "edge" )
1502+ assert_equal (padded , xp .asarray ([1 , 1 , 1 , 2 , 3 , 3 ]))
1503+
1504+ def test_edge_ndim (self , xp : ModuleType ):
1505+ a = xp .asarray ([[1 , 2 ], [3 , 4 ]])
1506+ padded = pad (a , ((1 , 2 ), (2 , 1 )), mode = "edge" )
1507+ expected = xp .asarray (
1508+ [
1509+ [1 , 1 , 1 , 2 , 2 ],
1510+ [1 , 1 , 1 , 2 , 2 ],
1511+ [3 , 3 , 3 , 4 , 4 ],
1512+ [3 , 3 , 3 , 4 , 4 ],
1513+ [3 , 3 , 3 , 4 , 4 ],
1514+ ]
1515+ )
1516+ assert_equal (padded , expected )
1517+
1518+ def test_wrap (self , xp : ModuleType ):
1519+ a = xp .asarray ([1 , 2 , 3 ])
1520+ padded = pad (a , (5 , 4 ), mode = "wrap" )
1521+ assert_equal (padded , xp .asarray ([2 , 3 , 1 , 2 , 3 , 1 , 2 , 3 , 1 , 2 , 3 , 1 ]))
1522+
1523+ def test_wrap_ndim (self , xp : ModuleType ):
1524+ a = xp .asarray ([[1 , 2 ], [3 , 4 ]])
1525+ padded = pad (a , ((1 , 1 ), (1 , 1 )), mode = "wrap" )
1526+ expected = xp .asarray ([[4 , 3 , 4 , 3 ], [2 , 1 , 2 , 1 ], [4 , 3 , 4 , 3 ], [2 , 1 , 2 , 1 ]])
1527+ assert_equal (padded , expected )
1528+
1529+ @pytest .mark .parametrize ("mode" , ["edge" , "wrap" ])
1530+ def test_empty_axis (self , xp : ModuleType , mode : str ):
1531+ a = xp .asarray ([])
1532+ with pytest .raises (ValueError , match = "can't extend empty axis" ):
1533+ _ = pad (a , 1 , mode = mode ) # type: ignore[arg-type] # pyright: ignore[reportArgumentType]
1534+
14991535 def test_mode_not_implemented (self , xp : ModuleType ):
15001536 a = xp .asarray ([1 , 2 , 3 ])
1501- with pytest .raises (NotImplementedError , match = "Only `'constant'` " ):
1502- _ = pad (a , 2 , mode = "edge " ) # type: ignore[arg-type] # pyright: ignore[reportArgumentType]
1537+ with pytest .raises (NotImplementedError , match = "Unsupported padding mode " ):
1538+ _ = pad (a , 2 , mode = "reflect " ) # type: ignore[arg-type] # pyright: ignore[reportArgumentType]
15031539
15041540 def test_device (self , xp : ModuleType , device : Device ):
15051541 a = xp .asarray (0.0 , device = device )
0 commit comments