Skip to content

Commit 151d8f3

Browse files
authored
refactor(tensor): simplify select_dim signature by replacing generic type constraints with impl AsIndex, update related usage in silero_vad module (#152)
1 parent 14de5a9 commit 151d8f3

2 files changed

Lines changed: 10 additions & 9 deletions

File tree

‎crates/bunsen/src/burner/tensor/tensor_op_ext.rs‎

Lines changed: 8 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -34,10 +34,10 @@ where
3434
fn extract(&mut self) -> Self;
3535

3636
/// Select (and Squeeze) a dimension.
37-
fn select_dim<const D2: usize, I: AsIndex>(
37+
fn select_dim<const D2: usize>(
3838
self,
39-
dim: usize,
40-
index: I,
39+
dim: impl AsIndex,
40+
index: impl AsIndex,
4141
) -> Tensor<B, D2, K>;
4242
}
4343

@@ -59,13 +59,14 @@ where
5959
z
6060
}
6161

62-
fn select_dim<const D2: usize, I: AsIndex>(
62+
fn select_dim<const D2: usize>(
6363
self,
64-
dim: usize,
65-
index: I,
64+
dim: impl AsIndex,
65+
index: impl AsIndex,
6666
) -> Tensor<B, D2, K> {
67+
let dim = dim.expect_dim_index(D);
6768
let index = index.as_index();
68-
self.slice_dim(dim, index..index + 1).squeeze_dim::<D2>(dim)
69+
self.slice_dim(dim, index).squeeze_dim::<D2>(dim)
6970
}
7071
}
7172

‎crates/bunsen/src/kits/speech/silero_vad/blocks/module.rs‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -648,7 +648,7 @@ impl<B: Backend> SileroVad<B> {
648648
(mut $acc:ident) => {{
649649
for step in 0..steps {
650650
// [batch, d_hidden]
651-
let features = seq_features.clone().select_dim::<2, _>(0, step);
651+
let features = seq_features.clone().select_dim::<2>(0, step);
652652

653653
// [batch, d_hidden]
654654
(hidden, cell) = self.lstm_step(features, hidden, cell);
@@ -773,7 +773,7 @@ impl<B: Backend> SileroVad<B> {
773773
let mag = (real_2 + imag_2).sqrt();
774774

775775
// Encode, then take the first (and, for a single chunk, only) frame.
776-
let x = self.encoder.forward(mag).select_dim::<2, _>(2, 0);
776+
let x = self.encoder.forward(mag).select_dim::<2>(2, 0);
777777

778778
#[cfg(any(test, debug_assertions))]
779779
crate::contracts::assert_shape_contract_periodically!(

0 commit comments

Comments
 (0)