Skip to content
7 changes: 2 additions & 5 deletions cutie/inference/memory_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -244,11 +244,8 @@ def add_memory(
for obj_id, obj in enumerate(objects):
if obj in self.obj_v:
# Keep embedding sums and counts for the object transformer's streaming average.
last_acc = self.obj_v[obj][:, :, -1]
new_acc = last_acc + obj_value[:, obj_id, :, -1]

self.obj_v[obj][:, :, :-1] = self.obj_v[obj][:, :, :-1] + obj_value[:, obj_id, :, :-1]
self.obj_v[obj][:, :, -1] = new_acc
self.obj_v[obj][:, :, :-1].add_(obj_value[:, obj_id, :, :-1])
self.obj_v[obj][:, :, -1].add_(obj_value[:, obj_id, :, -1])
else:
self.obj_v[obj] = obj_value[:, obj_id]

Expand Down
2 changes: 1 addition & 1 deletion cutie/model/big_modules.py
Original file line number Diff line number Diff line change
Expand Up @@ -80,7 +80,7 @@ def __init__(self, model_cfg: DictConfig):

def forward(self, x: torch.Tensor, *, need_s: bool, need_e: bool) -> (torch.Tensor, torch.Tensor, torch.Tensor):
x = self.pix_feat_proj(x)
shrinkage = self.d_proj(x) ** 2 + 1 if (need_s) else None
shrinkage = self.d_proj(x).pow(2).add_(1) if (need_s) else None
selection = torch.sigmoid(self.e_proj(x)) if (need_e) else None

return self.key_proj(x), shrinkage, selection
Expand Down
2 changes: 1 addition & 1 deletion cutie/model/channel_attn.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,4 +31,4 @@ def forward(self, x: torch.Tensor) -> torch.Tensor:
w = self.pool(x).view(b, 1, c)
w = self.conv(w).transpose(-1, -2).unsqueeze(-1).sigmoid() # B*C*1*1

return x * w + self.downsample(r) if self.residual else x * w
return torch.addcmul(self.downsample(r), x, w) if self.residual else x * w
5 changes: 3 additions & 2 deletions cutie/model/cutie.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,7 +54,8 @@ def _get_others(self, masks: torch.Tensor) -> torch.Tensor:
return (masks.sum(dim=1, keepdim=True) - masks).clamp(0, 1) if num_objects >= 1 else torch.zeros_like(masks)

def encode_image(self, image: torch.Tensor) -> (Iterable[torch.Tensor], torch.Tensor):
image = (image - self.pixel_mean) / self.pixel_std
# sub() copies (never mutates the caller's frame); only the fresh copy is div_'d in place
image = image.sub(self.pixel_mean).div_(self.pixel_std)
ms_image_feat = self.pixel_encoder(image)
return ms_image_feat, self.pix_feat_proj(ms_image_feat[0])

Expand All @@ -69,7 +70,7 @@ def encode_mask(
chunk_size: int = -1,
need_weights: bool = False,
) -> (torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor):
image = (image - self.pixel_mean) / self.pixel_std
image = image.sub(self.pixel_mean).div_(self.pixel_std)
others = self._get_others(masks)
mask_value, new_sensory = self.mask_encoder(
image,
Expand Down
2 changes: 1 addition & 1 deletion cutie/model/group_modules.py
Original file line number Diff line number Diff line change
Expand Up @@ -89,7 +89,7 @@ def forward(self, x: torch.Tensor, g: torch.Tensor, skip_expand: bool = False) -
elif self.method == 'mulcat':
g = torch.cat([x * g, g], dim=2)
elif self.method == 'muladd':
g = x * g + g
g = torch.addcmul(g, x, g)
else:
raise NotImplementedError

Expand Down
17 changes: 11 additions & 6 deletions cutie/model/utils/memory_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,16 +30,18 @@ def get_similarity(
# See XMem's appendix for derivation
mk = mk.transpose(1, 2)
a_sq = mk.pow(2) @ qe
two_ab = 2 * (mk @ (qk * qe))
two_ab = (mk @ (qk * qe)).mul_(2)
b_sq = (qe * qk.pow(2)).sum(1, keepdim=True)
similarity = -a_sq + two_ab - b_sq
similarity = two_ab.sub_(a_sq).sub_(b_sq)
else:
# similar to STCN if we don't have the selection term
a_sq = mk.pow(2).sum(1).unsqueeze(2)
two_ab = 2 * (mk.transpose(1, 2) @ qk)
similarity = -a_sq + two_ab
two_ab = (mk.transpose(1, 2) @ qk).mul_(2)
similarity = two_ab.sub_(a_sq)

return similarity * ms / math.sqrt(CK) if ms is not None else similarity / math.sqrt(CK) # B*N*HW
if ms is not None:
return similarity.mul_(ms).div_(math.sqrt(CK))
return similarity.div_(math.sqrt(CK)) # B*N*HW


def do_softmax(
Expand All @@ -63,7 +65,10 @@ def do_softmax(
affinity = torch.zeros_like(similarity).scatter_(1, indices, x_exp) # B*N*HW
else:
maxes = torch.max(similarity, dim=1, keepdim=True)[0]
x_exp = torch.exp(similarity - maxes)
# exp_() is safe (sub() output is a fresh throwaway tensor); div_() is NOT —
# Exp's backward needs its own output value preserved, so the final
# normalization must stay out-of-place (confirmed by backward-parity test).
x_exp = similarity.sub(maxes).exp_()
x_exp_sum = torch.sum(x_exp, dim=1, keepdim=True)
affinity = x_exp / x_exp_sum
indices = None
Expand Down
83 changes: 83 additions & 0 deletions scripts/bench_vram_mps.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,83 @@
"""Standalone MPS VRAM bench for the memory-read hot path.

Not a pytest assertion — MPS allocator memory numbers are noisy run-to-run, so
this is evidence to eyeball (before/after the in-place-op refactor), not a CI
gate. Correctness is covered by tests/test_mps_parity.py; this script only
measures peak Apple-GPU memory and wall-clock for cutie.model.utils.memory_utils
get_similarity + do_softmax at a realistic memory-bank size.

Usage:
python scripts/bench_vram_mps.py [--frames N] [--hw N] [--ck N] [--iters N]
"""

import argparse
import time

import torch
from torch import mps

from cutie.model.utils.memory_utils import do_softmax, get_similarity


def _build_inputs(
*, batch: int, ck: int, num_memory_frames: int, hw: int, device: str
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
n = num_memory_frames * hw
mk = torch.randn(batch, ck, n, device=device)
ms = torch.rand(batch, 1, n, device=device)
qk = torch.randn(batch, ck, hw, device=device)
qe = torch.rand(batch, ck, hw, device=device)
return mk, ms, qk, qe


def _run_once(mk: torch.Tensor, ms: torch.Tensor, qk: torch.Tensor, qe: torch.Tensor) -> torch.Tensor:
similarity = get_similarity(mk, ms, qk, qe)
return do_softmax(similarity)


def bench(*, batch: int, ck: int, num_memory_frames: int, hw: int, iters: int) -> None:
if not torch.backends.mps.is_available():
print('MPS not available on this machine — nothing to bench.')
return
Comment thread
Copilot marked this conversation as resolved.

device = 'mps'
mk, ms, qk, qe = _build_inputs(batch=batch, ck=ck, num_memory_frames=num_memory_frames, hw=hw, device=device)

# warm up (first MPS dispatch pays kernel-compile cost, not representative)
_run_once(mk, ms, qk, qe)
torch.mps.synchronize()

mps.empty_cache()
baseline_allocated = mps.current_allocated_memory()

start = time.perf_counter()
for _ in range(iters):
affinity = _run_once(mk, ms, qk, qe)
torch.mps.synchronize()
elapsed = time.perf_counter() - start

peak_allocated = mps.driver_allocated_memory()

print('MPS VRAM bench — cutie.model.utils.memory_utils (get_similarity + do_softmax)')
print(f' shapes: batch={batch} ck={ck} memory_frames={num_memory_frames} hw={hw} -> N={num_memory_frames * hw}')
print(f' iters: {iters}')
print(f' baseline allocated (post-warmup, pre-loop): {baseline_allocated / 2**20:.2f} MiB')
print(f' driver allocated (peak, post-loop): {peak_allocated / 2**20:.2f} MiB')
print(f' wall-clock: {elapsed:.4f}s total, {elapsed / iters * 1000:.3f}ms/iter')
print(f' output shape: {tuple(affinity.shape)}')


def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument('--batch', type=int, default=1)
parser.add_argument('--ck', type=int, default=64, help='key channel dim')
parser.add_argument('--frames', type=int, default=20, dest='num_memory_frames', help='accumulated memory frames')
parser.add_argument('--hw', type=int, default=30 * 54, help='flattened spatial size (H*W/patch)')
parser.add_argument('--iters', type=int, default=50)
args = parser.parse_args()

bench(batch=args.batch, ck=args.ck, num_memory_frames=args.num_memory_frames, hw=args.hw, iters=args.iters)


if __name__ == '__main__':
main()
Loading
Loading