Wraps conformer + encoder_proj + transformer + speaker head. Input: embs (1, T_total, 512) float, emb_len (1,) int64. Output: preds (1, T_total, 4) float — per-frame speaker probabilities. The C++ runner concatenates [spkcache, fifo, chunk_embs] into a single contiguous tensor and
| 88 | |
| 89 | |
| 90 | class EncodeWrapper(torch.nn.Module): |
| 91 | """Wraps conformer + encoder_proj + transformer + speaker head. |
| 92 | |
| 93 | Input: embs (1, T_total, 512) float, emb_len (1,) int64. |
| 94 | Output: preds (1, T_total, 4) float — per-frame speaker probabilities. |
| 95 | |
| 96 | The C++ runner concatenates [spkcache, fifo, chunk_embs] into a single |
| 97 | contiguous tensor and passes it here. Calls encoder with bypass_pre_encode=True |
| 98 | to skip ConvSubsampling (already run separately). |
| 99 | """ |
| 100 | |
| 101 | def __init__(self, encoder, encoder_proj, transformer_encoder, sortformer_modules): |
| 102 | super().__init__() |
| 103 | self.encoder = encoder |
| 104 | self.encoder_proj = encoder_proj if encoder_proj is not None else nn.Identity() |
| 105 | self.transformer_encoder = transformer_encoder |
| 106 | self.sortformer_modules = sortformer_modules |
| 107 | |
| 108 | def forward(self, embs: torch.Tensor, emb_len: torch.Tensor) -> torch.Tensor: |
| 109 | # Conformer layers (skip pre_encode since input is already subsampled) |
| 110 | encoded, enc_len = self.encoder( |
| 111 | audio_signal=embs, length=emb_len, bypass_pre_encode=True |
| 112 | ) |
| 113 | # encoded shape: (B, d_model, T) — transpose to (B, T, d_model) |
| 114 | encoded = encoded.transpose(1, 2) |
| 115 | |
| 116 | # Project from conformer dim to transformer dim: (B, T, 512) -> (B, T, 192) |
| 117 | encoded = self.encoder_proj(encoded) |
| 118 | |
| 119 | # Transformer encoder |
| 120 | encoder_mask = self.sortformer_modules.length_to_mask(enc_len, encoded.shape[1]) |
| 121 | trans_out = self.transformer_encoder( |
| 122 | encoder_states=encoded, encoder_mask=encoder_mask |
| 123 | ) |
| 124 | |
| 125 | # Speaker head: Linear(192->192) -> ReLU -> Linear(192->4) -> Sigmoid |
| 126 | preds = self.sortformer_modules.forward_speaker_sigmoids(trans_out) |
| 127 | preds = preds * encoder_mask.unsqueeze(-1) |
| 128 | return preds |
| 129 | |
| 130 | |
| 131 | def load_model(nemo_path: Optional[str] = None, hf_model: Optional[str] = None): |