Translation
MLX
Core ML
ONNX
Safetensors
Japanese
Chinese
jmangatranslator-fast
manga
japanese
chinese
Instructions to use muscgab/JMangaTranslator-Fast with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- MLX
How to use muscgab/JMangaTranslator-Fast with MLX:
# Download the model from the Hub pip install huggingface_hub[hf_xet] hf download muscgab/JMangaTranslator-Fast --local-dir JMangaTranslator-Fast
- Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- LM Studio
- Atomic Chat
Download src/model.py from muscgab/JMangaTranslator-Fast: direct link, hf CLI and curl.
- Browser
- Download file 20.6 kB
-
https://huggingface.co/muscgab/JMangaTranslator-Fast/resolve/main/src/model.py
- Command line
-
hf download hf://muscgab/JMangaTranslator-Fast/src/model.py
-
curl -L -o model.py https://huggingface.co/muscgab/JMangaTranslator-Fast/resolve/main/src/model.py
20.6 kB
| """ModernBERT-ja encoder + shallow autoregressive decoder for ja->zh block translation. | |
| encoder HF ModernBERT (bidirectional), all hidden states returned. | |
| memory base = Bridge(last_hidden_state); decoder layer j reads | |
| memory_j = [null x2 ; base + gamma_j * Wo(sum_l softmax(w_j)_l RMSNorm_l(h_l))] over h_0..h_{L-1} | |
| (static depth fusion of mangaOCR-NAR; gamma starts at 0 so memory_j == base at init). | |
| bridge "swiglu": RMSNorm -> SwiGLU(E -> hidden -> d) -> RMSNorm | |
| "mlp": RMSNorm -> Linear(E -> hidden) -> GELU -> Linear(hidden -> d) -> RMSNorm (OCR bridge) | |
| "linear": RMSNorm -> Linear(E -> d) -> RMSNorm | |
| decoder pre-RMSNorm blocks: causal self-attention, cross-attention to memory_j, SwiGLU; learned positions; | |
| output tied to the (compact) target embedding table. Inputs are scaled by sqrt(d) as in the DAT, | |
| so DAT embedding rows can initialise the table. | |
| context (cfg ctx_max_dist > 0) every block is encoded once (encode_blocks); a sample's memory is | |
| [null x2 ; current block ; earlier blocks] gathered from the encoded blocks (assemble), each position plus | |
| dist[k] (k = 0 current block, k = 1.. blocks back, zero-initialised). The decoder input is | |
| <ctx_zh> zh(earlier blocks, <sep>-joined) <bos> target; the loss covers the target only. | |
| """ | |
| from __future__ import annotations | |
| import math | |
| import torch | |
| from torch import nn | |
| from torch.nn import functional as F | |
| class RMSNorm(nn.Module): | |
| def __init__(self, d: int, eps: float = 1e-6): | |
| super().__init__() | |
| self.weight = nn.Parameter(torch.ones(d)) | |
| self.eps = eps | |
| def forward(self, x): | |
| return (x.float() * torch.rsqrt(x.float().pow(2).mean(-1, keepdim=True) + self.eps)).type_as(x) * self.weight | |
| class SwiGLU(nn.Module): | |
| def __init__(self, d_in: int, hidden: int, d_out: int): | |
| super().__init__() | |
| self.gate_up = nn.Linear(d_in, 2 * hidden, bias=False) | |
| self.down = nn.Linear(hidden, d_out, bias=False) | |
| def forward(self, x): | |
| g, u = self.gate_up(x).chunk(2, -1) | |
| return self.down(F.silu(g) * u) | |
| def make_bridge(kind: str, e: int, d: int, hidden: int) -> nn.Module: | |
| if kind == "swiglu": | |
| return nn.Sequential(RMSNorm(e), SwiGLU(e, hidden, d), RMSNorm(d)) | |
| if kind == "mlp": | |
| return nn.Sequential(RMSNorm(e), nn.Linear(e, hidden, bias=False), nn.GELU(), nn.Linear(hidden, d, bias=False), | |
| RMSNorm(d)) | |
| if kind == "linear": | |
| return nn.Sequential(RMSNorm(e), nn.Linear(e, d, bias=False), RMSNorm(d)) | |
| raise ValueError(kind) | |
| class Fusion(nn.Module): | |
| def __init__(self, e: int, d: int, dec_layers: int, depths: int): | |
| super().__init__() | |
| self.norms = nn.ModuleList(RMSNorm(e) for _ in range(depths)) | |
| self.wo = nn.Linear(e, d, bias=False) # random init so gamma receives gradient | |
| self.gamma = nn.Parameter(torch.zeros(dec_layers, d)) | |
| self.logits = nn.Parameter(torch.zeros(dec_layers, depths)) | |
| def forward(self, states): | |
| h = torch.stack([n(s) for n, s in zip(self.norms, states)], 0) # [D,B,N,E] | |
| f = torch.einsum("jl,lbne->jbne", self.logits.softmax(-1).to(h.dtype), h) | |
| return self.gamma[:, None, None].to(h.dtype) * self.wo(f) # [J,B,N,d] | |
| class DecoderLayer(nn.Module): | |
| def __init__(self, d: int, heads: int, ffn: int, dropout: float): | |
| super().__init__() | |
| self.h, self.hd, self.dropout = heads, d // heads, dropout | |
| self.n1, self.n2, self.n3 = RMSNorm(d), RMSNorm(d), RMSNorm(d) | |
| self.qkv = nn.Linear(d, 3 * d, bias=False) | |
| self.o1 = nn.Linear(d, d, bias=False) | |
| self.q2 = nn.Linear(d, d, bias=False) | |
| self.kv2 = nn.Linear(d, 2 * d, bias=False) | |
| self.o2 = nn.Linear(d, d, bias=False) | |
| self.ffn = SwiGLU(d, ffn, d) | |
| def heads(self, x): | |
| return x.view(x.shape[0], x.shape[1], self.h, self.hd).transpose(1, 2) | |
| def merge(self, y): | |
| return y.transpose(1, 2).reshape(y.shape[0], y.shape[2], -1) | |
| def drop(self, y): | |
| return F.dropout(y, self.dropout, self.training) | |
| def cross_kv(self, mem): | |
| k, v = self.kv2(mem).chunk(2, -1) | |
| return self.heads(k), self.heads(v) | |
| def forward(self, x, ck, cv, cmask): | |
| q, k, v = self.qkv(self.n1(x)).chunk(3, -1) | |
| x = x + self.drop(self.o1(self.merge(F.scaled_dot_product_attention( | |
| self.heads(q), self.heads(k), self.heads(v), is_causal=True)))) | |
| x = x + self.drop(self.o2(self.merge(F.scaled_dot_product_attention( | |
| self.heads(self.q2(self.n2(x))), ck, cv, attn_mask=cmask)))) | |
| return x + self.drop(self.ffn(self.n3(x))) | |
| def step(self, x, t, cache, ck, cv, cmask): | |
| """One position for every row; cache [2,B,H,T,hd] holds self-attention K/V of earlier positions.""" | |
| q, k, v = self.qkv(self.n1(x)).chunk(3, -1) | |
| cache[0, :, :, t] = self.heads(k)[:, :, 0] | |
| cache[1, :, :, t] = self.heads(v)[:, :, 0] | |
| x = x + self.o1(self.merge(F.scaled_dot_product_attention( | |
| self.heads(q), cache[0, :, :, : t + 1], cache[1, :, :, : t + 1]))) | |
| x = x + self.o2(self.merge(F.scaled_dot_product_attention( | |
| self.heads(self.q2(self.n2(x))), ck, cv, attn_mask=cmask))) | |
| return x + self.ffn(self.n3(x)) | |
| class ARMT(nn.Module): | |
| def __init__(self, cfg: dict): | |
| super().__init__() | |
| from transformers import AutoConfig, AutoModel | |
| self.cfg = cfg | |
| enc_cfg = AutoConfig.from_pretrained(cfg["encoder"]) | |
| self.encoder = AutoModel.from_pretrained(cfg["encoder"], attn_implementation="sdpa", dtype=torch.float32) \ | |
| if cfg.get("load_encoder_weights", True) else AutoModel.from_config(enc_cfg, attn_implementation="sdpa") | |
| e = enc_cfg.hidden_size | |
| d = cfg.get("width") or e | |
| self.d, L = d, cfg["dec_layers"] | |
| self.bridge = make_bridge(cfg["bridge"], e, d, cfg["bridge_hidden"]) | |
| self.fusion = Fusion(e, d, L, enc_cfg.num_hidden_layers) if cfg["fusion"] else None | |
| self.null = nn.Parameter(torch.randn(cfg["null_tokens"], d) * 0.02) if cfg["null_tokens"] else None | |
| self.emb = nn.Embedding(cfg["vocab"], d) | |
| self.pos = nn.Embedding(cfg["max_tgt"], d) | |
| self.layers = nn.ModuleList(DecoderLayer(d, d // 64, cfg["ffn"], cfg["dropout"]) for _ in range(L)) | |
| self.norm = RMSNorm(d) | |
| self.dist = nn.Embedding(cfg["ctx_max_dist"] + 1, d) if cfg.get("ctx_max_dist") else None | |
| nn.init.normal_(self.emb.weight, std=d ** -0.5) | |
| nn.init.normal_(self.pos.weight, std=0.02) | |
| if self.dist is not None: | |
| nn.init.zeros_(self.dist.weight) | |
| for lin in [m for m in self.modules() if isinstance(m, nn.Linear) and not self._in_encoder(m)]: | |
| nn.init.normal_(lin.weight, std=0.02) | |
| for layer in self.layers: # depth-scaled residual outputs | |
| for lin in (layer.o1, layer.o2, layer.ffn.down): | |
| nn.init.normal_(lin.weight, std=0.02 / math.sqrt(2 * L)) | |
| def _in_encoder(self, module) -> bool: | |
| return any(module is m for m in self.encoder.modules()) | |
| # ------------------------------------------------------------------ init helpers | |
| def init_embeddings(self, table: torch.Tensor, rows: list[int]) -> int: | |
| """Copy DAT embedding rows (compact index i <- DAT id rows[i]; rows[i] < 0 keeps the random init).""" | |
| n = 0 | |
| for i, r in enumerate(rows): | |
| if r >= 0: | |
| self.emb.weight[i] = table[r].to(self.emb.weight.dtype) | |
| n += 1 | |
| return n | |
| def decoder_parameters(self): | |
| return [p for n, p in self.named_parameters() if not n.startswith("encoder.")] | |
| # ------------------------------------------------------------------ forward pieces | |
| def memories(self, src, src_mask, encoder_grad: bool = True): | |
| with torch.set_grad_enabled(encoder_grad and torch.is_grad_enabled()): | |
| out = self.encoder(input_ids=src, attention_mask=src_mask.long(), output_hidden_states=self.fusion is not None) | |
| base = self.bridge(out.last_hidden_state) | |
| mems = base[None].expand(len(self.layers), -1, -1, -1) | |
| if self.fusion is not None: | |
| mems = mems + self.fusion(list(out.hidden_states[:-1])) | |
| if self.dist is not None: # single-block path = current block (distance 0) | |
| mems = mems + self.dist.weight[0].to(mems.dtype) | |
| mask = src_mask | |
| if self.null is not None: | |
| b = src.shape[0] | |
| mems = torch.cat((self.null[None, None].expand(len(self.layers), b, -1, -1).to(mems.dtype), mems), 2) | |
| mask = torch.cat((torch.ones(b, self.null.shape[0], dtype=torch.bool, device=src.device), src_mask), 1) | |
| return mems, mask[:, None, None, :] | |
| def _encode(self, src, src_mask, encoder_grad: bool): | |
| with torch.set_grad_enabled(encoder_grad and torch.is_grad_enabled()): | |
| out = self.encoder(input_ids=src, attention_mask=src_mask.long(), output_hidden_states=self.fusion is not None) | |
| base = self.bridge(out.last_hidden_state) | |
| mems = base[None].expand(len(self.layers), -1, -1, -1) | |
| if self.fusion is not None: | |
| mems = mems + self.fusion(list(out.hidden_states[:-1])) | |
| return mems | |
| def encode_blocks(self, src, src_mask, encoder_grad: bool = True, groups: int = 1): | |
| """Encode every block once and keep only real tokens: returns (packed [J, N, d], start [U], L) where block u | |
| occupies packed[:, start[u] : start[u] + len_u]. groups > 1 encodes length-sorted groups, each padded only to | |
| its own longest block.""" | |
| U, L = src.shape | |
| lens = src_mask.sum(1) | |
| order = torch.argsort(lens) if groups > 1 and U >= 2 * groups else torch.arange(U, device=src.device) | |
| parts = torch.tensor_split(order, groups) if groups > 1 and U >= 2 * groups else [order] | |
| chunks, start = [], torch.zeros(U, dtype=torch.long, device=src.device) | |
| off = 0 | |
| for part in parts: | |
| if part.numel() == 0: | |
| continue | |
| lg = int(lens[part].max()) | |
| m = self._encode(src[part, :lg], src_mask[part, :lg], encoder_grad) # [J, g, lg, d] | |
| pm = src_mask[part, :lg] | |
| chunks.append(m[:, pm]) # row-major: block by block | |
| pl = lens[part] | |
| start[part] = off + torch.cumsum(pl, 0) - pl | |
| off += int(pl.sum()) | |
| return torch.cat(chunks, 1), start, L | |
| def assemble(self, blocks, idx, dist, valid): | |
| """Gather sample memories from encoded blocks (encode_blocks output). | |
| idx [B, M] = u * L + p (block u, position p; pad entries arbitrary, masked by valid), dist [B, M] distance ids, | |
| valid [B, M] bool. Returns mems [J, B, n_null + M, d] and mask [B, 1, 1, n_null + M].""" | |
| packed, start, L = blocks | |
| J, d = packed.shape[0], packed.shape[-1] | |
| flat = torch.where(valid, start[idx // L] + idx % L, torch.zeros_like(idx)) | |
| mems = packed[:, flat] # [J, B, M, d] | |
| if self.dist is not None: | |
| mems = mems + self.dist(dist)[None].to(mems.dtype) | |
| mask = valid | |
| if self.null is not None: | |
| b = idx.shape[0] | |
| mems = torch.cat((self.null[None, None].expand(J, b, -1, -1).to(mems.dtype), mems), 2) # noqa: E501 | |
| mask = torch.cat((torch.ones(b, self.null.shape[0], dtype=torch.bool, device=idx.device), valid), 1) | |
| return mems, mask[:, None, None, :] | |
| def decode_logits(self, mems, cmask, tgt_in): | |
| x = F.dropout(self.embed(tgt_in), self.cfg["dropout"], self.training) | |
| for layer, mem in zip(self.layers, mems): | |
| ck, cv = layer.cross_kv(mem) | |
| x = layer(x, ck, cv, cmask) | |
| return self.norm(x) @ self.emb.weight.T | |
| def loss_ctx(self, src, src_mask, idx, dist, valid, tgt_in, gold, pad: int, smoothing: float, | |
| encoder_grad: bool = True, groups: int = 1): | |
| """gold = next-token targets with prefix positions set to pad (ignored). Logits are computed only at target | |
| positions. groups > 1: blocks are encoded in length groups and samples are decoded in length groups, each | |
| padded to its own longest decoder input / memory.""" | |
| blocks = self.encode_blocks(src, src_mask, encoder_grad, groups) | |
| tlen = tgt_in.ne(pad).sum(1) | |
| order = torch.argsort(tlen) if groups > 1 else torch.arange(tgt_in.shape[0], device=tgt_in.device) | |
| tot_s = tot_p = 0.0 | |
| n_all = gold.ne(pad).sum() | |
| for part in (torch.tensor_split(order, groups) if groups > 1 else [order]): | |
| if part.numel() == 0: | |
| continue | |
| T = int(tlen[part].max()) | |
| M = int(valid[part].sum(1).max()) | |
| mems, cmask = self.assemble(blocks, idx[part, :M], dist[part, :M], valid[part, :M]) | |
| x = F.dropout(self.embed(tgt_in[part, :T]), self.cfg["dropout"], self.training) | |
| for layer, mem in zip(self.layers, mems): | |
| ck, cv = layer.cross_kv(mem) | |
| x = layer(x, ck, cv, cmask) | |
| g = gold[part, :T] | |
| ok = g.ne(pad) | |
| logits = (self.norm(x[ok]) @ self.emb.weight.T).float() | |
| tot_s = tot_s + F.cross_entropy(logits, g[ok], reduction="sum", label_smoothing=smoothing) | |
| tot_p = tot_p + F.cross_entropy(logits.detach(), g[ok], reduction="sum") | |
| return tot_s / n_all, tot_p / n_all, n_all | |
| def generate_ctx(self, mems, cmask, prefixes, bos: int, eos: int, pad: int, max_len: int, byte_compact=None): | |
| """Greedy decoding after a per-row forced prefix (list of id lists ending before <bos>; [] = no context). | |
| Rows keep absolute positions from 0, so each row reads prefix, <bos>, then its own output.""" | |
| b = len(prefixes) | |
| dev = cmask.device | |
| rules = self.kana_byte_rules(byte_compact, self.emb.weight.shape[0], dev) if byte_compact else [] | |
| lp = torch.tensor([len(p) for p in prefixes], device=dev) | |
| P = int(lp.max()) if b else 0 | |
| forced = torch.full((b, P + 1), bos, dtype=torch.long, device=dev) | |
| for r, p in enumerate(prefixes): | |
| if p: | |
| forced[r, :len(p)] = torch.tensor(p, device=dev) | |
| cross = [layer.cross_kv(mem) for layer, mem in zip(self.layers, mems)] | |
| steps = min(P + max_len, self.cfg["max_tgt"]) # outputs start at t = len(prefix) | |
| caches = [torch.zeros(2, b, layer.h, steps, layer.hd, device=dev, dtype=mems.dtype) for layer in self.layers] | |
| prev1 = torch.full((b,), -1, dtype=torch.long, device=dev) | |
| prev2 = prev1.clone() | |
| done = torch.zeros(b, dtype=torch.bool, device=dev) | |
| tok = forced[:, :1] | |
| out = [] | |
| for t in range(steps): | |
| x = self.embed(tok, t) | |
| for layer, cache, (ck, cv) in zip(self.layers, caches, cross): | |
| x = layer.step(x, t, cache, ck, cv, cmask) | |
| gen = t >= lp # this step's output is a target token | |
| logits = (self.norm(x) @ self.emb.weight.T)[:, -1] | |
| for p2, p1, m in rules: | |
| hit = prev1.eq(p1) if p2 is None else prev1.eq(p1) & prev2.eq(p2) | |
| logits = logits.masked_fill((hit & gen)[:, None] & m[None], float("-inf")) | |
| nxt = logits.argmax(-1) | |
| nxt = torch.where(done, torch.full_like(nxt, pad), nxt) | |
| out.append(torch.where(gen, nxt, torch.full_like(nxt, -1))) | |
| done |= gen & nxt.eq(eos) | |
| if bool(done.all()) or t + 1 >= steps: | |
| break | |
| nf = forced[:, min(t + 1, P)] | |
| feed = torch.where(t + 1 <= lp, nf, nxt) # still inside prefix (incl. <bos>) -> forced | |
| prev2, prev1 = prev1, torch.where(gen, nxt, prev1) | |
| tok = feed[:, None] | |
| seqs = torch.stack(out, 1).tolist() if out else [[] for _ in range(b)] | |
| res = [] | |
| for s in seqs: | |
| r = [] | |
| for x in s: | |
| if x == -1: | |
| continue | |
| if x in (eos, pad): | |
| break | |
| r.append(x) | |
| res.append(r) | |
| return res | |
| def embed(self, ids, start: int = 0): | |
| pos = torch.arange(start, start + ids.shape[1], device=ids.device) | |
| return self.emb(ids) * math.sqrt(self.d) + self.pos(pos)[None] | |
| def forward(self, src, src_mask, tgt_in, encoder_grad: bool = True): | |
| mems, cmask = self.memories(src, src_mask, encoder_grad) | |
| x = F.dropout(self.embed(tgt_in), self.cfg["dropout"], self.training) | |
| for layer, mem in zip(self.layers, mems): | |
| ck, cv = layer.cross_kv(mem) | |
| x = layer(x, ck, cv, cmask) | |
| return self.norm(x) @ self.emb.weight.T | |
| def loss(self, src, src_mask, tgt, pad: int, smoothing: float, encoder_grad: bool = True): | |
| logits = self(src, src_mask, tgt[:, :-1], encoder_grad).float() | |
| gold = tgt[:, 1:] | |
| valid = gold.ne(pad) | |
| nll = F.cross_entropy(logits.flatten(0, 1), gold.flatten(), reduction="none", label_smoothing=smoothing) | |
| plain = F.cross_entropy(logits.flatten(0, 1).detach(), gold.flatten(), reduction="none") | |
| n = valid.sum() | |
| return (nll * valid.flatten()).sum() / n, (plain * valid.flatten()).sum() / n, n | |
| def kana_byte_rules(byte_compact: list[int], vocab: int, device): | |
| """Masks that stop byte-fallback pieces from spelling kana in UTF-8: | |
| E3 81|82|83 xx (U+3040-30FF), E3 87 B0-BF (U+31F0-31FF), EF BD A6-BF and EF BE 80-9D (half-width).""" | |
| B = torch.tensor(byte_compact, device=device) | |
| def mask(byte_values): | |
| m = torch.zeros(vocab, dtype=torch.bool, device=device) | |
| m[B[list(byte_values)]] = True | |
| return m | |
| return [(None, int(B[0xE3]), mask([0x81, 0x82, 0x83])), | |
| (int(B[0xE3]), int(B[0x87]), mask(range(0xB0, 0xC0))), | |
| (int(B[0xEF]), int(B[0xBD]), mask(range(0xA6, 0xC0))), | |
| (int(B[0xEF]), int(B[0xBE]), mask(range(0x80, 0x9E)))] | |
| def generate(self, src, src_mask, bos: int, eos: int, pad: int, max_len: int, byte_compact=None): | |
| """Batched greedy decoding with KV cache; returns compact id lists without BOS/EOS. | |
| byte_compact (compact ids of <0x00>..<0xFF>) enables the kana byte-sequence ban.""" | |
| b = src.shape[0] | |
| rules = self.kana_byte_rules(byte_compact, self.emb.weight.shape[0], src.device) if byte_compact else [] | |
| prev1 = torch.full((b,), -1, dtype=torch.long, device=src.device) | |
| prev2 = prev1.clone() | |
| mems, cmask = self.memories(src, src_mask, encoder_grad=False) | |
| cross = [layer.cross_kv(mem) for layer, mem in zip(self.layers, mems)] | |
| steps = min(max_len, self.cfg["max_tgt"] - 1) | |
| caches = [torch.zeros(2, b, layer.h, steps, layer.hd, device=src.device, dtype=mems.dtype) for layer in self.layers] | |
| tok = torch.full((b, 1), bos, dtype=torch.long, device=src.device) | |
| done = torch.zeros(b, dtype=torch.bool, device=src.device) | |
| out = [] | |
| for t in range(steps): | |
| x = self.embed(tok, t) | |
| for layer, cache, (ck, cv) in zip(self.layers, caches, cross): | |
| x = layer.step(x, t, cache, ck, cv, cmask) | |
| logits = (self.norm(x) @ self.emb.weight.T)[:, -1] | |
| for p2, p1, m in rules: | |
| hit = prev1.eq(p1) if p2 is None else prev1.eq(p1) & prev2.eq(p2) | |
| logits = logits.masked_fill(hit[:, None] & m[None], float("-inf")) | |
| nxt = logits.argmax(-1) | |
| nxt = torch.where(done, torch.full_like(nxt, pad), nxt) | |
| prev2, prev1 = prev1, nxt | |
| out.append(nxt) | |
| done |= nxt.eq(eos) | |
| if bool(done.all()): | |
| break | |
| tok = nxt[:, None] | |
| seqs = torch.stack(out, 1).tolist() if out else [[] for _ in range(b)] | |
| res = [] | |
| for s in seqs: | |
| r = [] | |
| for x in s: | |
| if x in (eos, pad): | |
| break | |
| r.append(x) | |
| res.append(r) | |
| return res | |