(
&self,
x: &Tensor,
_index_pos: usize,
_block_idx: usize,
_ctx: &mut Context,
)
| 98 | } |
| 99 | |
| 100 | async fn forward( |
| 101 | &self, |
| 102 | x: &Tensor, |
| 103 | _index_pos: usize, |
| 104 | _block_idx: usize, |
| 105 | _ctx: &mut Context, |
| 106 | ) -> Result<Tensor> { |
| 107 | // The input tensor packs time_emb as the first frame: |
| 108 | // [batch, 1 + seq_len, dim] where x[:, 0:1, :] is time_emb |
| 109 | let (batch, total_seq, _) = x.dims3()?; |
| 110 | let seq_len = total_seq - 1; |
| 111 | |
| 112 | // Extract time_emb (first frame) and data (rest) |
| 113 | let time_emb = x.narrow(1, 0, 1)?; // [batch, 1, dim] |
| 114 | let data = x.narrow(1, 1, seq_len)?; // [batch, seq_len, dim] |
| 115 | |
| 116 | // Generate relative position embeddings |
| 117 | let pos_emb = self.make_pos_emb(seq_len, x.device(), x.dtype())?; |
| 118 | |
| 119 | // Apply the Zipformer layer with time embedding |
| 120 | let out = self.layer.forward(&data, &pos_emb, Some(&time_emb))?; |
| 121 | |
| 122 | // Re-pack time_emb as first frame for next layer |
| 123 | let result = Tensor::cat(&[&time_emb, &out], 1)?; // [batch, 1 + seq_len, dim] |
| 124 | let _ = batch; // suppress unused warning |
| 125 | Ok(result) |
| 126 | } |
| 127 | |
| 128 | async fn forward_mut( |
| 129 | &mut self, |
no test coverage detected