Commit 7738c057 authored by brkirch's avatar brkirch
Browse files

MPS fix is still needed :(

Apparently I did not test with large enough images to trigger the bug with torch.narrow on MPS
parent 2c1bb46c
Loading
Loading
Loading
Loading
+3 −0
Original line number Diff line number Diff line
@@ -207,3 +207,6 @@ if has_mps():
        cumsum_needs_bool_fix = not torch.BoolTensor([True,True]).to(device=torch.device("mps"), dtype=torch.int64).equal(torch.BoolTensor([True,False]).to(torch.device("mps")).cumsum(0))
        torch.cumsum = lambda input, *args, **kwargs: ( cumsum_fix(input, orig_cumsum, *args, **kwargs) )
        torch.Tensor.cumsum = lambda self, *args, **kwargs: ( cumsum_fix(self, orig_Tensor_cumsum, *args, **kwargs) )
        orig_narrow = torch.narrow
        torch.narrow = lambda *args, **kwargs: ( orig_narrow(*args, **kwargs).clone() )