Skip to content

Repository files navigation

sam2-rs

Native Candle-backed SAM2.1 inference with static, video/z-slice, automatic-mask, and direct PyTorch/safetensors checkpoint-loading APIs.

Translation coverage and the bottom-up remaining-feature checklist are tracked in TRANSLATION.md. The original architecture handover remains in handover.md.

Inference APIs

Published .pt checkpoints load directly; the same APIs also accept converted .safetensors files. Cached image prediction reuses the image encoder:

use sam2_rs::{ImagePromptBatch, Sam2ImagePredictor, Sam2Variant};

let mut predictor = Sam2ImagePredictor::load_sam2_1(
    Sam2Variant::HieraTiny,
    "checkpoints/sam2.1_hiera_tiny.pt",
)?;
predictor.set_image(&rgb_bytes, width, height)?;
let prediction = predictor.predict_cached(
    &ImagePromptBatch::default(),
    false,
    false,
)?;

Automatic masks and editable multi-object video use typed configurations:

use sam2_rs::{
    AutomaticMaskConfig, PointPrompt, Sam2AutomaticMaskGenerator,
    Sam2InteractiveVideoPredictor, Sam2Variant, VideoPredictorConfig,
};

let generator = Sam2AutomaticMaskGenerator::load_sam2_1(
    Sam2Variant::HieraTiny,
    "checkpoints/sam2.1_hiera_tiny.pt",
    AutomaticMaskConfig::default(),
)?;
let masks = generator.generate_rgb8(&rgb_bytes, width, height)?;

let video = Sam2InteractiveVideoPredictor::load_sam2_1(
    Sam2Variant::HieraTiny,
    "checkpoints/sam2.1_hiera_tiny.pt",
    VideoPredictorConfig::default(),
)?;
let mut session = video.init_state(&frames)?;
video.add_new_points(
    &mut session,
    0,
    101,
    &[PointPrompt { x: 120.0, y: 80.0, positive: true }],
    true,
)?;
let tracked = video.propagate_in_video(&mut session, None, None, false)?;

Here rgb_bytes is interleaved RGB8 and frames is a slice of (&[u8], width, height) tuples. See PARITY.md for reproducible Python/Rust conformance commands.

Optional downloading and media decoding

The default build remains network- and codec-free. Enable only the integrations an application needs:

sam2-rs = { version = "0.1", features = ["huggingface", "video-decoding"] }

huggingface adds cached checkpoint retrieval and from_pretrained constructors. It honors the standard HF_TOKEN, HF_HOME, HF_HUB_CACHE, and HF_ENDPOINT environment variables:

use sam2_rs::{Sam2ImagePredictor, Sam2Variant};

let mut predictor = Sam2ImagePredictor::from_pretrained(Sam2Variant::HieraTiny)?;
predictor.set_image_path("input.jpg")?;

Use HuggingFaceCheckpoint::for_variant(variant) when a pinned revision, custom cache directory, offline-only lookup, or forced refresh is needed.

image-decoding accepts JPEG/PNG images and upstream-style directories whose JPEG filenames have numeric stems. video-decoding includes image decoding and adds encoded-video support through the system FFmpeg libraries:

use sam2_rs::{Sam2InteractiveVideoPredictor, Sam2Variant, VideoPredictorConfig};

let predictor = Sam2InteractiveVideoPredictor::from_pretrained(
    Sam2Variant::HieraTiny,
    VideoPredictorConfig::default(),
)?;
let jpeg_session = predictor.init_state_from_jpeg_directory("frames")?;
let mp4_session = predictor.init_state_from_video("clip.mp4")?;
# let _ = (jpeg_session, mp4_session);

Building video-decoding requires FFmpeg development packages discoverable by pkg-config (on Debian/Ubuntu: libavcodec-dev, libavformat-dev, libavutil-dev, and libswscale-dev). decode_video_prefix is available for bounded previews without retaining an entire video in memory.

Reproducible performance baseline

CPU latency, host RSS, and output parity are the active performance scope. Ratios are Rust divided by pinned PyTorch; lower than 1.0× favors Rust. The full P0/P1 scorecard is not complete yet, so the project does not claim overall CPU parity.

Representative current CPU evidence Latency Peak RSS Status
Static image, one point, Tiny/20 workers 0.889× 0.446× faster on this fixture only
Cached prompt decode (I3) 1.543× clean; 1.360× latest diagnostic 0.322× latest layer-specific candidate fixes parity; speed and four-pair promotion remain open
Batched/prior-mask decode (I4) 1.600× 0.404× improved; speed fails and cars is narrowly noisy
Video initialization (V1) 0.919× 0.435× passes current cell gates
Video propagation (V2) 1.213× 0.583× clean row fails speed; latest 0.932× candidate is noisy
Practical automatic masks (A1) 1.049× 0.566× cells pass; weighted target open
Default-grid automatic masks (A2) 0.927× 0.854× speed passes; weighted RSS open
Model load (L1) 0.132× 0.276× faster, but Python run was noisy
Small-model first interaction (S1) 1.047× 0.338× passes current cell gates
Base-Plus/Large scaling (M1) 1.012× latest diagnostic 0.515× both variants pass provisionally; four-pair promotion pending

Run the model-free schema/timing/parity check first:

python3 tools/run_parity_suite.py --validate-only

Pinned performance runs use tools/run_parity_suite.py --device cpu with the cell, fixture, 8/20-worker count, and four alternating process pairs selected as described in CPU_PARITY_PLAN.md. Current row-level evidence and exact archive paths are in CPU_PARITY_RESULTS.md; accepted and rejected optimizations are in PERFORMANCE_INVESTIGATION.md.

The prior GPU suite is historical and outside the current completion gates. Its last complete eager result was 0.941× weighted latency and 0.312× host RSS; details remain archived in CPU_PARITY_RESULTS.md.

Historical early CPU baseline

The generated 17.351× row below predates the retained Candle and sam2-rs CPU optimizations. It is kept only to show the starting point and must not be read as current performance.

Generated on Linux 6.8.0-58-generic, CPU: Intel(R) Xeon(R) Gold 6138 CPU @ 2.00GHz; fixture SHA-256: 24271176ecdc. Peak RSS is Linux VmHWM.

Fixture / operation Rust median ms Python median ms Rust/Python latency Rust peak RSS MiB Python peak RSS MiB Rust/Python RSS Max abs mask-logit error
bedroom/00089.jpg, Tiny static point prompt (3 measured iterations) 12606.34 726.56 17.351× 1713.5 1434.2 1.195× 4.50611e-05

Historical CPU stage-attribution diagnostic

This one-iteration diagnostic is separate from the median baseline above. It uses the same warmed-up fixture and includes prompt encoding with the decoder.

Stage Rust ms Python ms Rust/Python
Image encoder (Hiera + FPN) 10813.10 738.39 14.644×
Prompt encoder + mask decoder 484.93 20.92 23.181×

Reproduce with target/release/bench-static checkpoints/sam2.1_hiera_tiny.safetensors checkpoints/bench_bedroom.safetensors 1 --profile-stages and the equivalent tools/bench_static_python.py command.

Local Candle-fork CPU diagnostic

The direct local Candle fork includes its CPU Conv2D fast paths. On the same warmed-up one-iteration diagnostic, it preserved the full mask-logit maximum error (4.50611e-05) and improved the default-Candle result as follows.

Metric Default Candle Local Candle fork Fork/default
Total static inference 10939.36 ms 9758.35 ms 0.892×
Image encoder 10398.43 ms 9349.32 ms 0.899×
Prompt encoder + mask decoder 529.60 ms 396.46 ms 0.749×
FPN 776.81 ms 576.28 ms 0.742×

SIMD CPU flash-attention diagnostic

The local Candle commit 480dbc1b adds a float32 AVX2/FMA online-softmax path for long attention. It is opt-in with --features cpu-flash-attention; the default remains available for comparison.

Metric Fork default SIMD flash feature Feature/default
Total static inference 9758.35 ms 7007.11 ms 0.718×
Image encoder 9349.32 ms 6569.83 ms 0.703×
Peak RSS 1688.9 MiB 661.4 MiB 0.392×
Max abs mask-logit error vs PyTorch 4.50611e-05 5.00679e-05 —

The feature is still 10.75× the one-iteration PyTorch CPU diagnostic, but its three global-attention blocks fall from about 1.5 s each to 0.53–0.55 s. Reproduce the full ratio table with python3 tools/run_benchmarks.py --rust-features cpu-flash-attention.

Historical experimental CPU baseline

The following focused rows were useful during optimization but are not the current multi-scenario CPU scorecard above.

Generated on Linux 6.8.0-58-generic, CPU: Intel(R) Xeon(R) Gold 6138 CPU @ 2.00GHz; fixture SHA-256: 24271176ecdc. CPU workers: 20, pinned to CPUs 0-19. Peak RSS is Linux VmHWM.

Fixture / operation Rust mean ms Python mean ms Rust/Python latency Rust mean peak RSS MiB Python mean peak RSS MiB Rust/Python RSS Max abs mask-logit error
bedroom/00089.jpg, Tiny static point prompt on cpu; Rust cpu-flash-attention,cpu-glibc-vector-math (8 independent runs × 5 measured iterations) 580.06 652.42 0.889× 629.6 1411.8 0.446× 2.86102e-05

This current-code row combines two four-run batches on the shared host. Their latency ratios were 0.891× and 0.887×; all eight individual Rust/PyTorch pairs favored Rust. The arithmetic mean is reproducible with the runner, but contention makes small cross-date changes inconclusive. CPU speed parity is reached for this fixture, machine, and 20-worker configuration, not established for every variant, prompt, or CPU. Use paired A/B measurements in PERFORMANCE_INVESTIGATION.md before attributing an optimization. Exact within-block decoder key/position reuse and image-time projection of high-resolution skips are included in this row. Their repeated-decode gains are measured, but paired full-model A/B runs have not established a clean gain under the shared host's varying load. The direct Hiera query-pool path is included in this row. Its paired full-model A/B averaged 0.967× of the prior Rust path with byte-identical masks. An earlier attempted cross-runtime batch had a 6.58 s PyTorch contention outlier and was excluded; see PERFORMANCE_INVESTIGATION.md. On Linux/glibc, loading a CPU model also sets glibc's mmap and trim thresholds to 256 MiB once per process. This keeps large temporary attention buffers on the heap and trades some peak RSS for speed. It has no effect on CUDA model loading. Set SAM2_CPU_KEEP_LARGE_ALLOCATIONS=0 to opt out; existing MALLOC_MMAP_THRESHOLD_ or MALLOC_TRIM_THRESHOLD_ settings also take precedence. This tuning is process-global, so embedding applications should consider its effect on other allocations. Run python3 tools/run_benchmarks.py --iterations 5 --repeats 4 --threads 20 --cpu-list 0-19 --rust-features cpu-flash-attention,cpu-glibc-vector-math --no-write twice to collect the same number of independent runs.

An earlier run before the tiny-query decoder change, with eight pinned CPU workers, measured 1190.83 ms Rust versus 828.61 ms PyTorch (1.437×) across four independent runs of five iterations. Mean peak RSS was 588.2 MiB versus 1582.1 MiB (0.372×), with maximum mask-logit error 2.90871e-05. Reproduce with python3 tools/run_benchmarks.py --iterations 5 --repeats 4 --threads 8 --cpu-list 0-7 --rust-features cpu-flash-attention,cpu-glibc-vector-math --no-write.

Historical GPU performance baseline

Generated on Linux 6.8.0-58-generic, CPU: Intel(R) Xeon(R) Gold 6138 CPU @ 2.00GHz. GPU: Quadro RTX 5000, 570.133.20. Fixture SHA-256: 24271176ecdc. CPU workers: implementation default, alternating process order. Peak RSS is Linux VmHWM.

Fixture / operation Rust mean ms Python mean ms Rust/Python latency Rust mean peak RSS MiB Python mean peak RSS MiB Rust/Python RSS Max abs mask-logit error
bedroom/00089.jpg, Tiny static point prompt on cuda; Rust cuda,cuda-graph-both (4 independent runs × 10 measured iterations) 67.97 67.16 1.012× 470.8 1192.6 0.395× 1.43051e-05

GPU timings synchronize device work around every sample. Peak RSS is host VmHWM, not device VRAM. The row uses CUDA graph capture and replay for both implementations because its image, checkpoint, shape, and prompt are static. Reproduce it with python3 tools/run_gpu_benchmarks.py --cuda-graph --alternate-order --repeats 4 --iterations 10. The retained eager path measured 75.90 ms Rust versus 67.20 ms PyTorch (1.129×) in an independent four-run batch. Graph mode removes Candle's remaining repeated allocation/event/launch bookkeeping; it does not change the model output.

About

No description, website, or topics provided.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages