MCPcopy Create free account
hub / github.com/evilsocket/cake / simple_downsample

Function simple_downsample

cake-core/src/models/luxtts/model.rs:54–86  ·  view source on GitHub ↗

SimpleDownsample: weighted average over groups of `ds` frames using learned softmax weights.

(src: &Tensor, ds: usize, bias: &Tensor)

Source from the content-addressed store, hash-verified

52
53/// SimpleDownsample: weighted average over groups of `ds` frames using learned softmax weights.
54fn simple_downsample(src: &Tensor, ds: usize, bias: &Tensor) -> Result<Tensor> {
55 // src: [batch, seq_len, dim], bias: [ds]
56 let (batch, seq_len, dim) = src.dims3()?;
57 let d_seq_len = seq_len.div_ceil(ds);
58 let padded_len = d_seq_len * ds;
59
60 // Pad by repeating last frame if needed
61 let src = if padded_len > seq_len {
62 let last_frame = src.narrow(1, seq_len - 1, 1)?; // [batch, 1, dim]
63 let pad_count = padded_len - seq_len;
64 let padding = last_frame.expand((batch, pad_count, dim))?;
65 Tensor::cat(&[src, &padding], 1)?
66 } else {
67 src.clone()
68 };
69
70 // Reshape to [batch, d_seq_len, ds, dim]
71 let src = src.reshape((batch, d_seq_len, ds, dim))?;
72
73 // Softmax weights: bias [ds] -> softmax -> [1, 1, ds, 1]
74 let b = bias.unsqueeze(0)?;
75 let max = b.max_keepdim(1)?;
76 let exp = b.broadcast_sub(&max)?.exp()?;
77 let sum = exp.sum_keepdim(1)?;
78 let weights = exp.broadcast_div(&sum)?; // [1, ds]
79 let weights = weights.reshape((1, 1, ds, 1))?;
80
81 // Weighted sum over ds dimension
82 let weighted = src.broadcast_mul(&weights)?; // [batch, d_seq_len, ds, dim]
83 let result = weighted.sum(2)?; // [batch, d_seq_len, dim]
84
85 Ok(result)
86}
87
88/// SimpleUpsample: repeat each frame `ds` times.
89fn simple_upsample(src: &Tensor, ds: usize) -> Result<Tensor> {

Callers 1

generate_speechMethod · 0.85

Calls 1

cloneMethod · 0.45

Tested by

no test coverage detected