facebookresearch / facebookresearch/sam3

SAM3.1. size assert error during compile

Open
#534 1 comment 0 reactions 0 assignees View on GitHub
Dominant language
Python
Stars
11.7k
Forks
1.8k
PR merge metrics
No merged PRs in 30d

Description

Hi there,
trying to compile SAM3.1 as I used to with SAM3
```python
predictor = build_sam3_multiplex_video_predictor(
use_fa3=False,
compile=True,
)
```

However, upon calling
```python
for response in predictor.handle_stream_request(
request=dict(
type="propagate_in_video",
session_id=session_id,
start_frame_index=0,
max_frame_num_to_track=num_frames
)
):
pass
```

Some compilation steps are successful before I receive the following error:
```bash
File "/PROJ_HOME/test_sam3.py", line 543, in run_sam31
for response in predictor.handle_stream_request(
File "/ENV_HOME/lib/python3.10/site-packages/torch/utils/_contextlib.py", line 38, in generator_context
response = gen.send(None)
File "/PROJ_HOME/subtrees/SAM3/sam3/model/sam3_base_predictor.py", line 90, in handle_stream_request
yield from self.propagate_in_video(
File "/PROJ_HOME/subtrees/SAM3/sam3/model/sam3_base_predictor.py", line 279, in propagate_in_video
for frame_idx, outputs in self.model.propagate_in_video(
File "/ENV_HOME/lib/python3.10/site-packages/torch/utils/_contextlib.py", line 38, in generator_context
response = gen.send(None)
File "/PROJ_HOME/subtrees/SAM3/sam3/model/sam3_multiplex_tracking.py", line 2332, in propagate_in_video
yield from super().propagate_in_video(
File "/ENV_HOME/lib/python3.10/site-packages/torch/utils/_contextlib.py", line 38, in generator_context
response = gen.send(None)
File "/PROJ_HOME/subtrees/SAM3/sam3/model/sam3_multiplex_tracking.py", line 366, in propagate_in_video
out = self._run_single_frame_inference(
File "/PROJ_HOME/subtrees/SAM3/sam3/model/sam3_multiplex_tracking.py", line 635, in _run_single_frame_inference
) = self._det_track_one_frame(
File "/PROJ_HOME/subtrees/SAM3/sam3/model/sam3_multiplex_base.py", line 447, in _det_track_one_frame
return self._det_track_one_frame_impl(
File "/PROJ_HOME/subtrees/SAM3/sam3/model/sam3_multiplex_base.py", line 527, in _det_track_one_frame_impl
self.run_tracker_propagation(
File "/PROJ_HOME/subtrees/SAM3/sam3/model/sam3_multiplex_base.py", line 810, in run_tracker_propagation
self._propogate_tracker_one_frame_local_gpu(
File "/PROJ_HOME/subtrees/SAM3/sam3/model/sam3_multiplex_base.py", line 1708, in _propogate_tracker_one_frame_local_gpu
for out in self.tracker.propagate_in_video(
File "/ENV_HOME/lib/python3.10/site-packages/torch/utils/_contextlib.py", line 38, in generator_context
response = gen.send(None)
File "/PROJ_HOME/subtrees/SAM3/sam3/model/video_tracking_multiplex_demo.py", line 3447, in propagate_in_video
current_out, pred_masks = self._run_single_frame_inference(
File "/PROJ_HOME/subtrees/SAM3/sam3/model/video_tracking_multiplex_demo.py", line 2793, in _run_single_frame_inference
current_out = self.track_step(
File "/PROJ_HOME/subtrees/SAM3/sam3/model/video_tracking_multiplex.py", line 3382, in track_step
current_out, aux_out = self._track_step_aux(
File "/PROJ_HOME/subtrees/SAM3/sam3/model/video_tracking_multiplex.py", line 2057, in _track_step_aux
pix_feat_with_mem = self._prepare_memory_conditioned_features(
File "/PROJ_HOME/subtrees/SAM3/sam3/model/video_tracking_multiplex.py", line 1590, in _prepare_memory_conditioned_features
encoder_out = self.transformer.encoder(
File "/ENV_HOME/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1775, in _wrapped_call_impl
return self._call_impl(*args, **kwargs)
File "/ENV_HOME/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1786, in _call_impl
return forward_call(*args, **kwargs)
File "/PROJ_HOME/subtrees/SAM3/sam3/perflib/compile.py", line 113, in wrapper
return fn(*args, **kwargs)
File "/PROJ_HOME/subtrees/SAM3/sam3/perflib/compile.py", line 75, in compiled_fn_wrapper
result = compiled_fn(*args, **kwargs)
File "/ENV_HOME/lib/python3.10/site-packages/torch/_dynamo/eval_frame.py", line 832, in compile_wrapper
return fn(*args, **kwargs)
File "/PROJ_HOME/subtrees/SAM3/sam3/model/decoder.py", line 1282, in forward
def forward(
File "/ENV_HOME/lib/python3.10/site-packages/torch/_dynamo/eval_frame.py", line 1044, in _fn
return fn(*args, **kwargs)
File "/ENV_HOME/lib/python3.10/site-packages/torch/_functorch/aot_autograd.py", line 1130, in forward
return compiled_fn(full_args)
File "/ENV_HOME/lib/python3.10/site-packages/torch/_functorch/_aot_autograd/runtime_wrappers.py", line 353, in runtime_wrapper
all_outs = call_func_at_runtime_with_args(
File "/ENV_HOME/lib/python3.10/site-packages/torch/_functorch/_aot_autograd/utils.py", line 129, in call_func_at_runtime_with_args
out = normalize_as_list(f(args))
File "/ENV_HOME/lib/python3.10/site-packages/torch/_functorch/_aot_autograd/runtime_wrappers.py", line 724, in inner_fn
outs = compiled_fn(args)
File "/ENV_HOME/lib/python3.10/site-packages/torch/_functorch/_aot_autograd/runtime_wrappers.py", line 526, in wrapper
return compiled_fn(runtime_args)
File "/ENV_HOME/lib/python3.10/site-packages/torch/_inductor/output_code.py", line 613, in __call__
return self.current_callable(inputs)
File "/ENV_HOME/lib/python3.10/site-packages/torch/_inductor/utils.py", line 2962, in run
out = model(new_inputs)
File "/TORCHINDUCTOR_CACHE_DIR/wl/cwl2lkbm2n6yfn7ps7vfaw4wqz4n6m56rxcq7ga6qfv3d2dmalts.py", line 2884, in call
assert_size_stride(arg6_1, (s28, 1, 256), (256, 256, 1))
AssertionError: expected size 10368==5184, stride 256==256 at dim=0
This error most often comes from a incorrect fake (aka meta) kernel for a custom op.
Use torch.library.opcheck to test your custom op.
See https://pytorch.org/docs/stable/library.html#torch.library.opcheck
```

Any suggestions?

Contributor guide

Open the contributing guide

Research direction

Reproduce the failure with build_sam3_multiplex_video_predictor(compile=True) and the propagate_in_video request. Start with sam3/perflib/compile.py and the decoder call in sam3/model/decoder.py, then trace the tensor through video_tracking_multiplex.py; the relevant assertion reports a size mismatch during the compiled encoder path. Done means SAM3.1 video propagation completes with compilation enabled without the assertion.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
computer-vision, machine-learning
Issue type
Bug
Difficulty
4/5
Estimated time
3-5 days
Activity status
Quiet
Clarity
Mostly clear
Newbie friendliness
48/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.