MCPcopy Create free account
hub / github.com/pytorch/executorch / EncodeWrapper

Class EncodeWrapper

examples/models/sortformer/export_sortformer.py:90–128  ·  view source on GitHub ↗

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

Source from the content-addressed store, hash-verified

88
89
90class 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
131def load_model(nemo_path: Optional[str] = None, hf_model: Optional[str] = None):

Callers 1

export_allFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected