From 148f30a4672fee465dabd029f65576f12b420122 Mon Sep 17 00:00:00 2001 From: Jim Steele Date: Mon, 8 Sep 2025 09:20:18 -0700 Subject: [PATCH] MPS support workaround for complex tensors Avoid repeat() on complex tensors in MPS by using cat fallback using torch.cat on MPS. Resolves the runtime error 'repeat(): Not supported for complex yet!' --- sam2/modeling/position_encoding.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/sam2/modeling/position_encoding.py b/sam2/modeling/position_encoding.py index f4b57ae..fafd042 100644 --- a/sam2/modeling/position_encoding.py +++ b/sam2/modeling/position_encoding.py @@ -211,6 +211,10 @@ def apply_rotary_enc( # repeat freqs along seq_len dim to match k seq_len if repeat_freqs_k: r = xk_.shape[-2] // xq_.shape[-2] - freqs_cis = freqs_cis.repeat(*([1] * (freqs_cis.ndim - 2)), r, 1) + if freqs_cis.is_complex() and freqs_cis.device.type == "mps": + # MPS doesn't support repeat on complex; cat works fine. + freqs_cis = torch.cat([freqs_cis] * r, dim=-2) + else: + freqs_cis = freqs_cis.repeat(*([1] * (freqs_cis.ndim - 2)), r, 1) xk_out = torch.view_as_real(xk_ * freqs_cis).flatten(3) return xq_out.type_as(xq).to(xq.device), xk_out.type_as(xk).to(xk.device)