bimm-contracts

Shape Contracts for Burn Image Models (BIMM).
- static/stack-evaluated runtime shape contracts for tensors.
Changelog
0.1.7
- Removed
assert_shape_every_n in favor of run_every_nth! macro.
- Improved isolation of
run_every_nth!.
Example Usage
use bimm_contracts::{ShapeContract, DimMatcher, DimExpr, run_every_nth};
pub fn window_partition<B: Backend, K>(
tensor: Tensor<B, 4, K>,
window_size: usize,
) -> Tensor<B, 4, K>
where
K: BasicOps<B>,
{
static INPUT_CONTRACT: ShapeContract = ShapeContract::new(&[
DimMatcher::Expr(DimExpr::Param("batch")),
DimMatcher::Expr(DimExpr::Prod(&[
DimExpr::Param("h_wins"),
DimExpr::Param("window_size"),
])),
DimMatcher::Expr(DimExpr::Prod(&[
DimExpr::Param("w_wins"),
DimExpr::Param("window_size"),
])),
DimMatcher::Expr(DimExpr::Param("channels")),
]);
let [b, h_wins, w_wins, c] = INPUT_CONTRACT.unpack_shape(
&tensor,
&["batch", "h_wins", "w_wins", "channels"],
&[("window_size", window_size)],
);
let tensor = tensor
.reshape([b, h_wins, window_size, w_wins, window_size, c])
.swap_dims(2, 3)
.reshape([b * h_wins * w_wins, window_size, window_size, c]);
static OUTPUT_CONTRACT: ShapeContract = ShapeContract::new(&[
DimMatcher::Expr(DimExpr::Prod(&[
DimExpr::Param("batch"),
DimExpr::Param("h_wins"),
DimExpr::Param("w_wins"),
])),
DimMatcher::Expr(DimExpr::Param("window_size")),
DimMatcher::Expr(DimExpr::Param("window_size")),
DimMatcher::Expr(DimExpr::Param("channels")),
]);
run_every_nth!(OUTPUT_CONTRACT.assert_shape(
&tensor,
&[
("batch", b),
("h_wins", h_wins),
("w_wins", w_wins),
("window_size", window_size),
("channels", c),
]
));
tensor
}
Performance
Benchmark: 230.51 ns/iter (+/- 5.22)
#[bench]
fn bench_shape_contract(b: &mut Bencher) {
static PATTERN: ShapeContract = ShapeContract::new(&[
DimMatcher::Any,
DimMatcher::Expr(DimExpr::Param("b")),
DimMatcher::Ellipsis,
DimMatcher::Expr(DimExpr::Prod(&[DimExpr::Param("h"), DimExpr::Param("p")])),
DimMatcher::Expr(DimExpr::Prod(&[DimExpr::Param("w"), DimExpr::Param("p")])),
DimMatcher::Expr(DimExpr::Pow(&DimExpr::Param("z"), 3)),
DimMatcher::Expr(DimExpr::Param("c")),
]);
let batch = 2;
let height = 3;
let width = 2;
let padding = 4;
let channels = 5;
let z = 4;
let shape = [12, batch, 1, 2, 3, height * padding, width * padding, z * z * z, channels];
let env = [("p", padding), ("c", channels)];
let keys = ["b", "h", "w", "z"];
b.iter(|| {
let _ = PATTERN.unpack_shape(&shape, &keys, &env);
});
}
run_every_nth!(CONTRACT.assert_shape(&tensor, &env))
Benchmark: 4.38 ns/iter (+/- 0.03)
#[bench]
fn bench_run_every_nth_assert_shape(b: &mut Bencher) {
static PATTERN: ShapeContract = ShapeContract::new(&[
DimMatcher::Any,
DimMatcher::Expr(DimExpr::Param("b")),
DimMatcher::Ellipsis,
DimMatcher::Expr(DimExpr::Prod(&[DimExpr::Param("h"), DimExpr::Param("p")])),
DimMatcher::Expr(DimExpr::Prod(&[DimExpr::Param("w"), DimExpr::Param("p")])),
DimMatcher::Expr(DimExpr::Pow(&DimExpr::Param("z"), 3)),
DimMatcher::Expr(DimExpr::Param("c")),
]);
let batch = 2;
let height = 3;
let width = 2;
let padding = 4;
let channels = 5;
let z = 4;
let shape = [
12,
batch,
1,
2,
3,
height * padding,
width * padding,
z * z * z,
channels,
];
let env = [("p", padding), ("c", channels)];
b.iter(|| {
run_every_nth!(PATTERN.assert_shape(&shape, &env));
});
}