Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions parakeet_mlx/__init__.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,9 @@
from parakeet_mlx.alignment import AlignedResult, AlignedSentence, AlignedToken
from parakeet_mlx.parakeet import (
BaseParakeet,
DecodingConfig,
ParakeetCTC,
ParakeetCTCArgs,
ParakeetDecodingConfig,
ParakeetRNNT,
ParakeetRNNTArgs,
ParakeetTDT,
Expand All @@ -15,7 +15,7 @@
from parakeet_mlx.utils import from_pretrained

__all__ = [
"DecodingConfig",
"ParakeetDecodingConfig",
"ParakeetTDTArgs",
"ParakeetTDT",
"ParakeetRNNT",
Expand Down
69 changes: 68 additions & 1 deletion parakeet_mlx/attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,7 +44,7 @@ def __call__(
k, v = cache.update_and_fetch_kv(k, v)

o = mx.fast.scaled_dot_product_attention(q, k, v, scale=self.scale, mask=mask)
o = o.transpose(0, 2, 1, 3).reshape(batch, q_seq, self.n_feat)
o = o.transpose(0, 2, 1, 3).reshape(batch, q_seq, self.head_dim * self.n_head)

return self.linear_out(o)

Expand Down Expand Up @@ -534,6 +534,52 @@ def matmul_pv(self, prob: mx.array, v: mx.array, w: int) -> mx.array:
return outputs[0]


# note that this has slight different scaling method than other encodings
class FixedPositionalEncoding(nn.Module):
def __init__(
self,
d_model: int,
max_len: int = 5000,
):
assert d_model % 2 == 0 and max_len > 0
super().__init__()

self.d_model = d_model
self.max_len = max_len
self.scale = math.sqrt(self.d_model)
self.calculate_pe()

def calculate_pe(self):
positions = mx.arange(self.max_len, dtype=mx.float32)
positions = mx.expand_dims(positions, axis=1)

div_term = mx.exp(
mx.arange(0, self.d_model, 2, dtype=mx.float32)
* -(math.log(10000.0) / self.d_model)
)
pe = mx.zeros((self.max_len, self.d_model), dtype=mx.float32)

pe[:, 0::2] = mx.sin(positions * div_term)
pe[:, 1::2] = mx.cos(positions * div_term)

self._pe = (
mx.expand_dims(pe, axis=0).astype(mx.float32) / self.scale
) # we scale here!

mx.eval(self._pe)

def __call__(self, x: mx.array, offset: int = 0) -> mx.array:
input_len = x.shape[1]

if offset + input_len > self.max_len:
self.max_len = offset + input_len
self.calculate_pe()

pos_emb = self._pe[:, offset : offset + input_len, :].astype(x.dtype)

return pos_emb


class RelPositionalEncoding(nn.Module):
def __init__(
self,
Expand Down Expand Up @@ -624,3 +670,24 @@ def __call__(self, x: mx.array, offset: int = 0) -> tuple[mx.array, mx.array]:
pos_emb = self._pe[:, :end_idx].astype(x.dtype)

return x, pos_emb


# utility
# thanks to mlx_lm
def create_causal_mask(
N: int,
offset: int = 0,
window_size: int | None = None,
lengths: mx.array | None = None,
):
rinds = mx.arange(offset + N)
linds = mx.arange(offset, offset + N) if offset else rinds
linds = linds[:, None]
rinds = rinds[None]
mask = linds >= rinds
if window_size is not None:
mask = mask & (linds <= rinds + window_size)
if lengths is not None:
lengths = lengths[:, None, None, None]
mask = mask & (rinds < lengths)
return mask
46 changes: 46 additions & 0 deletions parakeet_mlx/cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -154,3 +154,49 @@ def update_and_fetch_conv(self, x: mx.array, padding: int = 0) -> mx.array:
result = mx.pad(result, ((0, 0), (0, padding), (0, 0)))

return result


class TransformerDecoderCache:
keys: mx.array | None
values: mx.array | None

offset: int
step = 256

def __init__(self):
self.keys = None
self.values = None
self.conv = None
self.offset = 0

def update_and_fetch_kv(
self, keys: mx.array, values: mx.array
) -> tuple[mx.array, mx.array]:
# k, v is [batch, head, seq, dim]
prev = self.offset
if (
self.keys is None
or self.values is None
or (prev + keys.shape[2]) > self.keys.shape[2]
):
B, H, S, D_KEYS = keys.shape
_, _, _, D_VALUES = values.shape
S_CACHE = ((self.step + S - 1) // self.step) * self.step

new_k = mx.zeros((B, H, S_CACHE, D_KEYS), keys.dtype)
new_v = mx.zeros((B, H, S_CACHE, D_VALUES), keys.dtype)

if self.keys is None or self.values is None: # type safety!
self.keys, self.values = new_k, new_v
else:
if prev % self.step != 0:
self.keys = self.keys[..., :prev, :]
self.values = self.values[..., :prev, :]
self.keys = mx.concatenate([self.keys, new_k], axis=2)
self.values = mx.concatenate([self.values, new_v], axis=2)

self.offset += keys.shape[2]
self.keys[..., prev : self.offset, :] = keys
self.values[..., prev : self.offset, :] = values

return self.keys[..., : self.offset, :], self.values[..., : self.offset, :]
Loading