初始化项目,由ModelHub XC社区提供模型

Model: KBlueLeaf/TIPOv2-1B-A200M
Source: Original Platform
This commit is contained in:
ModelHub XC
2026-09-22 07:23:17 +08:00
commit 47c526c365
11 changed files with 747 additions and 0 deletions

38
.gitattributes vendored Normal file
View File

@@ -0,0 +1,38 @@
*.7z filter=lfs diff=lfs merge=lfs -text
*.arrow filter=lfs diff=lfs merge=lfs -text
*.bin filter=lfs diff=lfs merge=lfs -text
*.bz2 filter=lfs diff=lfs merge=lfs -text
*.ckpt filter=lfs diff=lfs merge=lfs -text
*.ftz filter=lfs diff=lfs merge=lfs -text
*.gz filter=lfs diff=lfs merge=lfs -text
*.h5 filter=lfs diff=lfs merge=lfs -text
*.joblib filter=lfs diff=lfs merge=lfs -text
*.lfs.* filter=lfs diff=lfs merge=lfs -text
*.mlmodel filter=lfs diff=lfs merge=lfs -text
*.model filter=lfs diff=lfs merge=lfs -text
*.msgpack filter=lfs diff=lfs merge=lfs -text
*.npy filter=lfs diff=lfs merge=lfs -text
*.npz filter=lfs diff=lfs merge=lfs -text
*.onnx filter=lfs diff=lfs merge=lfs -text
*.ot filter=lfs diff=lfs merge=lfs -text
*.parquet filter=lfs diff=lfs merge=lfs -text
*.pb filter=lfs diff=lfs merge=lfs -text
*.pickle filter=lfs diff=lfs merge=lfs -text
*.pkl filter=lfs diff=lfs merge=lfs -text
*.pt filter=lfs diff=lfs merge=lfs -text
*.pth filter=lfs diff=lfs merge=lfs -text
*.rar filter=lfs diff=lfs merge=lfs -text
*.safetensors filter=lfs diff=lfs merge=lfs -text
saved_model/**/* filter=lfs diff=lfs merge=lfs -text
*.tar.* filter=lfs diff=lfs merge=lfs -text
*.tar filter=lfs diff=lfs merge=lfs -text
*.tflite filter=lfs diff=lfs merge=lfs -text
*.tgz filter=lfs diff=lfs merge=lfs -text
*.wasm filter=lfs diff=lfs merge=lfs -text
*.xz filter=lfs diff=lfs merge=lfs -text
*.zip filter=lfs diff=lfs merge=lfs -text
*.zst filter=lfs diff=lfs merge=lfs -text
*tfevents* filter=lfs diff=lfs merge=lfs -text
gguf/TIPOv2-1B-A200M-f16.gguf filter=lfs diff=lfs merge=lfs -text
gguf/TIPOv2-1B-A200M-Q8_0.gguf filter=lfs diff=lfs merge=lfs -text
TIPOv2-1B-A200M-f16.gguf filter=lfs diff=lfs merge=lfs -text

259
README.md Normal file
View File

@@ -0,0 +1,259 @@
---
tags:
- text-generation
- transformers
- safetensors
- gguf
- text-generation-inference
- mixture-of-experts
- tipo
datasets:
- KBlueLeaf/danbooru2023-metadata-database
- KBlueLeaf/danbooru-8M-qwen3.5-caption
- CaptionEmporium/coyo-hd-11m-llavanext
- CaptionEmporium/laion-coco-13m-molmo-d-7b
- pixparse/cc12m-wds
arxiv: 2411.08127
language:
- en
pipeline_tag: text-generation
library_name: transformers
---
# TIPOv2-1B-A200M: Next generation of T2I prompt optimization model.
**TIPOv2-1B-A200M** is a **1B-A200M** sparse model. 991M total parameters, ~200M active
per token, plus a ~50M embedding table. The second generation of TIPO, rebuilt
from the dataset up on the KohakUwU MoE architecture.
![image](https://cdn-uploads.huggingface.co/production/uploads/630593e2fca1d8d92b81d2a1/SDijcDRj115C0wYazinFY.png)
## Introduction
TIPO is a framework for improving Text-to-Image generation by *text presampling*:
a small language model expands a short user prompt into a detailed one before the
diffusion model ever sees it. Pre-sampling a smaller distribution by narrow the range indicate by a brief prompt to a more specifici description which match the original prompt,
allow diffusion model to have more information to work while persist overall diversity and fidelity.
Instead of asking the user to write 200 tokens of booru tags and natural language, TIPO samples that expansion from a distribution
learned over real caption data.
This is **v2**. It is not a fine-tune of TIPO-500M. The dataset, the captioner
and the architecture are all different.
## What makes it v2
### 1. Fully upgraded dataset
| | v1 (TIPO-500M) | **v2 (TIPOv2-1B-A200M)** |
|---|---|---|
| sources | GBC10M, Danbooru, CoyoHD-11M | **latest Danbooru, Nozomi, CC12M, CoyoHD-11M, LAION-COCO-13M** |
| breadth | 3 sources | 5 sources, both anime-domain and general-photography |
The v1 mix leaned heavily on one general-caption source. v2 adds **Nozomi** and
**LAION-COCO-13M** alongside a refreshed Danbooru, which broadens both the tag
vocabulary and the image domains the model has seen. Danbooru is weighted x3 and
the dedicated tagger view x2, so booru-style tag structure stays dominant while
the general sources supply natural-language variety.
### 2. Better natural-language captions
Every natural-language caption in v2 is regenerated with **Qwen3.5-2B**. In v1 the
caption quality varied by source, because each dataset shipped whatever captions
its authors produced. Regenerating them under a single captioner means caption
*style* is constant across sources, and the only thing that varies between
`coyo11m` and `laion_coco` is the **image distribution**, not the writing. That
makes the source weighting a choice about visual domain rather than an accidental
choice about prose quality.
### 3. Fully upgraded architecture: KohakUwU MoE
v1 was a 500M dense LLaMA-like arch. v2 uses the **KohakUwU MoE** architecture, a
DeepSeekMoE-style sparse decoder from
[KohakUwULLM](https://github.com/KohakuBlueleaf/KohakUwULLM).
**KohakUwU** is a series of projects for pretraining infrastructure.
**KohakUwULLM** is the general-purpose LLM training project within that series,
and it is where this architecture, the training framework and the kernels
described below come from. None of it is part of TIPO, and none of it was built
for TIPO. TIPOv2 is one model trained with it.
The configuration used here:
| Configuration | |
|---|---|
| total params | **990.8M** |
| **active params / token** | **193.1M** (excludes the embedding lookup) |
| input embedding | 50.3M (a gather, not a matmul, so not counted as active) |
| output head | 50.3M |
| routed experts | 854.1M total, **106.8M active** at top-8 |
| attention + shared expert + dense layer + router + norms | 36.1M, all active |
| layers | 16 (layer 0 dense, 15 MoE) |
| hidden size | 768 |
| attention | 12 heads, 2 KV heads (GQA), head dim 64, **QK-norm** |
| routed experts | **64**, top-8 per token |
| shared experts | 1 (always on) |
| expert hidden | 384 |
| dense MLP hidden | 2048 |
| router | sigmoid scoring, **aux-loss-free** bias balancing |
| position | RoPE, theta 100000, 4096 context |
| norm | RMSNorm, eps 1e-6 |
| vocab | 65536 |
Only **193M of 991M** parameters do work on any given token. The 854M of routed
experts contribute just 107M at top-8, and the 50M embedding is a lookup rather
than a matmul. So v2 carries roughly 2x v1's parameters while activating fewer
of them per token than v1's dense 500M. Context is **4096**, up from v1's 1024.
## Training recipe
Trained on 4x RTX 5090 (32 GB, sm_120) with KohakUwULLM.
| | |
|---|---|
| steps | **150,000** |
| tokens per step | 262,144 (16384 x 16 microbatches) |
| context | 2048 packed |
| parallelism | 4-stage pipeline, **1F1B** schedule |
| parameter dtype | **full fp16** (with dynamic loss scaling) |
| autocast | fp16 |
| **MXFP8** | **q/k/v/o projections and MLP up/down (`w_in`/`w_out`), including the shared expert. 111 modules.** |
| routed experts | fused MXFP8 expert path |
| optimizer | **Muon** on hidden matrices, AdamW on the rest |
| LR | 5e-4 (`muon_lr` 2e-3, `embed_lr` 2e-3) |
| schedule | inverse-sqrt power (s0 2500, b -0.5), then cosine to 1% |
| warmup | 2% of run (3000 steps) |
| grad clip | 1.0 |
| aux loss / router z-loss | **0.0 / 0.0**, since balancing is aux-loss-free |
Notes on the choices that are not obvious. All of these are KohakUwULLM
facilities, not TIPO-specific work:
- **Packed varlen, not padded.** Every sequence is concatenated onto one flat
token axis with `cu_seqlens` carrying document boundaries. For TIPO-shaped data
(50 to 600 tokens against a 2048 context) a padded batch would be ~80% padding.
- **fp16 parameters, not bf16.** fp16 carries 10 mantissa bits against bf16's 7.
It needs loss scaling to keep its narrower exponent range in bounds, which the
trainer supplies; the run reports zero overflows at scale 65536.
- **Aux-loss-free balancing.** Expert load is balanced by a selection-only bias
updated outside the gradient, not by an auxiliary loss term. A router z-loss was
measured at 1.59x end-to-end cost and left off.
- **MXFP8 on the dense projections.** Block-scaled fp8 (E4M3 with a shared
power-of-two scale per 32 elements) on q/k/v/o and up/down. The routed experts
use a fused MXFP8 path whose epilogues never materialize the
`(tokens x top_k, hidden)` intermediates.
## Tokenizer
The tokenizer is the **DeepSeek-V4 tokenizer, pruned to 64000 ordinary BPE
tokens**, plus a 1536-slot block reserved for special tokens. Total vocabulary is
**65536**.
| id range | count | contents |
|---|---|---|
| 0 to 63999 | 64000 | ordinary BPE tokens, kept in DeepSeek-V4 merge order |
| 64000 to 64016 | 17 | named specials: `<\|bos\|>`, `<\|eos\|>`, `<\|pad\|>`, `<\|unk\|>`, and the 13 TIPO control tokens |
| 64017 to 65535 | 1519 | `<\|reserved_N\|>` placeholders |
Two reasons the layout looks like this:
- **65536 is a power of two.** The output head is a GEMM whose N dimension is the
vocabulary, and a power-of-two N keeps that GEMM tile-aligned. An awkward vocab
size costs throughput on every token generated.
- **The reserved block is deliberate headroom.** Adding a control token later is
an id assignment inside the existing embedding table, not a resize and
re-embed. 1519 slots are still free in this release.
## Prompt format
```
quality: masterpiece
rating: general
target: <|long|> <|tag_to_long|>
tag: 1girl, cherry blossoms, outdoors
```
### Control tokens
**Length targets**, which set how long the generated result should be:
`<|empty|>` `<|very_short|>` `<|short|>` `<|long|>` `<|very_long|>`
**Task selectors**, which set what to generate from what:
| token | meaning |
|---|---|
| `<\|tag_to_long\|>` | tags to long natural-language caption |
| `<\|long_to_tag\|>` | long caption to tags |
| `<\|short_to_tag\|>` | short caption to tags |
| `<\|short_to_long\|>` | short caption to long caption |
| `<\|tag_to_short_to_long\|>` | tags, then short, then long |
| `<\|short_to_tag_to_long\|>` | short, then tags, then long |
| `<\|short_to_long_to_tag\|>` | short, then long, then tags |
| `<\|gen_meta\|>` | also predict the metadata fields |
Metadata lines the model understands, all optional: `quality`, `rating`, `artist`,
`characters`, `copyrights`, `meta`, `aspect ratio`.
## Usage
<!-- PLACEHOLDER: confirm extension support for v2 before publishing -->
```python
from transformers import AutoTokenizer, AutoModelForCausalLM
model = AutoModelForCausalLM.from_pretrained(
"KBlueLeaf/TIPOv2-1B-A200M", trust_remote_code=True, dtype="float16"
).cuda().eval()
tokenizer = AutoTokenizer.from_pretrained("KBlueLeaf/TIPOv2-1B-A200M")
prompt = (
"quality: masterpiece\n"
"rating: general\n"
"target: <|long|> <|tag_to_long|>\n"
"tag: 1girl, cherry blossoms, outdoors\n"
)
ids = tokenizer(prompt, return_tensors="pt").input_ids.cuda()
out = model.generate(ids, max_new_tokens=256, temperature=1.0, min_p=0.1, do_sample=True)
print(tokenizer.decode(out[0], skip_special_tokens=False))
```
`trust_remote_code=True` is required, because the KohakUwU MoE architecture
ships as `modeling_kohaku.py` beside the weights.
### Files
| file | size | use |
|---|---|---|
| `hf/model.safetensors` | 1.98 GB | transformers, fp16 |
| `gguf/TIPOv2-1B-A200M-f16.gguf` | 2.02 GB | llama.cpp, fp16 |
| `gguf/TIPOv2-1B-A200M-Q8_0.gguf` | 1.07 GB | llama.cpp, 8-bit |
## LICENSE
Released under **Kohaku License 1.0**.
## Citation
TIPO:
```bibtex
@misc{yeh2024tipotextimagetext,
title={TIPO: Text to Image with Text Presampling for Prompt Optimization},
author={Yeh, Shih-Ying and Park, Sang-Hyun and Oh, Giyeong and Song, Min and Yu, Youngjae},
year={2024},
eprint={2411.08127},
archivePrefix={arXiv}
}
```
The architecture, training framework and kernels:
```bibtex
@software{kohakuwullm,
title={KohakUwULLM: an extensible decoder-only LLM training framework},
author={Yeh, Shih-Ying},
url={https://github.com/KohakuBlueleaf/KohakUwULLM},
year={2026}
}
```

3
TIPOv2-1B-A200M-f16.gguf Normal file
View File

@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:940a33e1ffe30886a0452c580d0d08124e4ca38655aca5501127432257c29128
size 2015578528

View File

@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:0a10a20ab475f978ab3bfab9e8b688559a80152d79cb05e39030c051ba708fc0
size 1072689632

View File

@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:940a33e1ffe30886a0452c580d0d08124e4ca38655aca5501127432257c29128
size 2015578528

32
hf/config.json Normal file
View File

@@ -0,0 +1,32 @@
{
"architectures": [
"KohakuForCausalLM"
],
"model_type": "kohaku",
"auto_map": {
"AutoConfig": "configuration_kohaku.KohakuConfig",
"AutoModel": "modeling_kohaku.KohakuModel",
"AutoModelForCausalLM": "modeling_kohaku.KohakuForCausalLM"
},
"vocab_size": 65536,
"hidden_size": 768,
"num_hidden_layers": 16,
"num_attention_heads": 12,
"num_key_value_heads": 2,
"head_dim": 64,
"intermediate_size": 2048,
"moe_intermediate_size": 384,
"n_routed_experts": 64,
"n_shared_experts": 1,
"num_experts_per_tok": 8,
"first_k_dense": 1,
"norm_topk_prob": true,
"routed_scaling_factor": 1.0,
"scoring_func": "sigmoid",
"max_position_embeddings": 4096,
"rope_theta": 100000.0,
"rms_norm_eps": 1e-06,
"qk_norm": true,
"tie_word_embeddings": false,
"dtype": "float16"
}

View File

@@ -0,0 +1,75 @@
"""HF config for the Kohaku decoder. Ships inside an exported repository.
Standalone by construction: an exported repo is loaded with
``trust_remote_code=True`` on machines that do not have kohakuwullm installed, so
nothing here may import it. See docs/guides/hf-export.md.
"""
from transformers.configuration_utils import PretrainedConfig
class KohakuConfig(PretrainedConfig):
"""Kohaku: GQA + per-head QK-norm, SwiGLU, and DeepSeek-style sparse MLPs.
Layers below ``first_k_dense`` use a dense SwiGLU; the rest use one shared
expert plus ``num_experts_per_tok`` of ``n_routed_experts``, selected on
sigmoid scores offset by a selection-only bias and weighted by the unbiased
score.
"""
model_type = "kohaku"
keys_to_ignore_at_inference = ["past_key_values"]
def __init__(
self,
vocab_size: int = 65536,
hidden_size: int = 768,
num_hidden_layers: int = 16,
num_attention_heads: int = 12,
num_key_value_heads: int = 2,
head_dim: int = 64,
intermediate_size: int = 2048,
moe_intermediate_size: int = 384,
n_routed_experts: int = 64,
n_shared_experts: int = 1,
num_experts_per_tok: int = 8,
first_k_dense: int = 1,
norm_topk_prob: bool = True,
routed_scaling_factor: float = 1.0,
scoring_func: str = "sigmoid",
max_position_embeddings: int = 4096,
rope_theta: float = 100000.0,
rms_norm_eps: float = 1e-6,
qk_norm: bool = True,
tie_word_embeddings: bool = False,
bos_token_id: int | None = 64000,
eos_token_id: int | None = 64001,
pad_token_id: int | None = 64002,
**kwargs,
) -> None:
self.vocab_size = vocab_size
self.hidden_size = hidden_size
self.num_hidden_layers = num_hidden_layers
self.num_attention_heads = num_attention_heads
self.num_key_value_heads = num_key_value_heads
self.head_dim = head_dim
self.intermediate_size = intermediate_size
self.moe_intermediate_size = moe_intermediate_size
self.n_routed_experts = n_routed_experts
self.n_shared_experts = n_shared_experts
self.num_experts_per_tok = num_experts_per_tok
self.first_k_dense = first_k_dense
self.norm_topk_prob = norm_topk_prob
self.routed_scaling_factor = routed_scaling_factor
self.scoring_func = scoring_func
self.max_position_embeddings = max_position_embeddings
self.rope_theta = rope_theta
self.rms_norm_eps = rms_norm_eps
self.qk_norm = qk_norm
super().__init__(
bos_token_id=bos_token_id,
eos_token_id=eos_token_id,
pad_token_id=pad_token_id,
tie_word_embeddings=tie_word_embeddings,
**kwargs,
)

3
hf/model.safetensors Normal file
View File

@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:6670583539124105a989243e5aa111aff05728b3e0e7ecb7110267ac49602a77
size 1981598832

321
hf/modeling_kohaku.py Normal file
View File

@@ -0,0 +1,321 @@
"""HF modeling code for the Kohaku decoder. Ships inside an exported repository.
Standalone by construction: loaded with ``trust_remote_code=True`` on machines
without kohakuwullm, so nothing here may import it. Routed experts stay stacked
as ``(E, out, in)``, the layout training uses, so an export is a copy rather than
a reshape. See docs/guides/hf-export.md.
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers.cache_utils import Cache, DynamicCache
from transformers.generation import GenerationMixin
from transformers.modeling_outputs import (
BaseModelOutputWithPast,
CausalLMOutputWithPast,
)
from transformers.modeling_utils import PreTrainedModel
from .configuration_kohaku import KohakuConfig
class KohakuRMSNorm(nn.Module):
def __init__(self, dim: int, eps: float = 1e-6) -> None:
super().__init__()
self.weight = nn.Parameter(torch.ones(dim))
self.eps = eps
self.normalized_shape = (dim,)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return F.rms_norm(x, self.normalized_shape, self.weight, self.eps)
def rotate_half(x: torch.Tensor) -> torch.Tensor:
half = x.shape[-1] // 2
return torch.cat((-x[..., half:], x[..., :half]), dim=-1)
def apply_rope(q, k, cos, sin):
"""``cos``/``sin`` are ``(B, S, head_dim)``; q/k are ``(B, S, H, head_dim)``."""
cos, sin = cos.unsqueeze(2), sin.unsqueeze(2)
return q * cos + rotate_half(q) * sin, k * cos + rotate_half(k) * sin
class KohakuRotary(nn.Module):
def __init__(self, config: KohakuConfig) -> None:
super().__init__()
inv = 1.0 / (
config.rope_theta
** (
torch.arange(0, config.head_dim, 2, dtype=torch.int64).float()
/ config.head_dim
)
)
self.register_buffer("inv_freq", inv, persistent=False)
@torch.no_grad()
def forward(self, x: torch.Tensor, position_ids: torch.Tensor):
freqs = position_ids[:, :, None].float() * self.inv_freq[None, None, :]
angles = torch.cat((freqs, freqs), dim=-1)
return angles.cos().to(x.dtype), angles.sin().to(x.dtype)
class KohakuAttention(nn.Module):
"""GQA with per-head QK-norm applied before RoPE."""
def __init__(self, config: KohakuConfig, layer_idx: int) -> None:
super().__init__()
self.layer_idx = layer_idx
self.heads = config.num_attention_heads
self.kv_heads = config.num_key_value_heads
self.head_dim = config.head_dim
self.scale = self.head_dim**-0.5
q_out = self.heads * self.head_dim
kv_out = self.kv_heads * self.head_dim
self.q_proj = nn.Linear(config.hidden_size, q_out, bias=False)
self.k_proj = nn.Linear(config.hidden_size, kv_out, bias=False)
self.v_proj = nn.Linear(config.hidden_size, kv_out, bias=False)
self.o_proj = nn.Linear(q_out, config.hidden_size, bias=False)
if config.qk_norm:
self.q_norm = KohakuRMSNorm(self.head_dim, config.rms_norm_eps)
self.k_norm = KohakuRMSNorm(self.head_dim, config.rms_norm_eps)
else:
self.q_norm = nn.Identity()
self.k_norm = nn.Identity()
def forward(self, x, cos, sin, attention_mask=None, past_key_values=None, **kwargs):
b, s, _ = x.shape
q = self.q_norm(self.q_proj(x).view(b, s, self.heads, self.head_dim))
k = self.k_norm(self.k_proj(x).view(b, s, self.kv_heads, self.head_dim))
v = self.v_proj(x).view(b, s, self.kv_heads, self.head_dim)
q, k = apply_rope(q, k, cos, sin)
q, k, v = (t.transpose(1, 2) for t in (q, k, v))
if past_key_values is not None:
k, v = past_key_values.update(k, v, self.layer_idx)
out = F.scaled_dot_product_attention(
q,
k,
v,
attn_mask=attention_mask,
scale=self.scale,
is_causal=attention_mask is None and s > 1,
enable_gqa=self.kv_heads != self.heads,
)
return self.o_proj(out.transpose(1, 2).reshape(b, s, -1))
class KohakuMLP(nn.Module):
"""SwiGLU. ``gate_proj``/``up_proj`` are the two halves of the trained ``w_in``."""
def __init__(self, hidden_size: int, intermediate_size: int) -> None:
super().__init__()
self.gate_proj = nn.Linear(hidden_size, intermediate_size, bias=False)
self.up_proj = nn.Linear(hidden_size, intermediate_size, bias=False)
self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=False)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))
class KohakuMoE(nn.Module):
"""One shared expert plus top-k of ``n_routed_experts``, experts kept stacked.
Selection adds ``expert_bias`` to the sigmoid scores; the weights come from
the unbiased scores, which is what makes the balancer auxiliary-loss-free.
"""
def __init__(self, config: KohakuConfig) -> None:
super().__init__()
self.top_k = config.num_experts_per_tok
self.norm_topk_prob = config.norm_topk_prob
self.routed_scaling_factor = config.routed_scaling_factor
self.scoring_func = config.scoring_func
e, d, h = (
config.n_routed_experts,
config.hidden_size,
config.moe_intermediate_size,
)
self.gate = nn.Linear(d, e, bias=False)
self.register_buffer("expert_bias", torch.zeros(e), persistent=True)
self.gate_proj = nn.Parameter(torch.empty(e, h, d))
self.up_proj = nn.Parameter(torch.empty(e, h, d))
self.down_proj = nn.Parameter(torch.empty(e, d, h))
self.shared_expert = KohakuMLP(d, h * config.n_shared_experts)
def score(self, logits: torch.Tensor) -> torch.Tensor:
if self.scoring_func == "softmax":
return logits.softmax(-1)
if self.scoring_func == "sqrtsoftplus":
return F.softplus(logits).sqrt()
return logits.sigmoid()
def forward(self, x: torch.Tensor) -> torch.Tensor:
shape = x.shape
flat = x.reshape(-1, shape[-1])
# Router runs wholly in fp32, matching training.
scores = self.score(F.linear(flat.float(), self.gate.weight.float()))
index = (scores + self.expert_bias).topk(self.top_k, dim=-1).indices
weight = scores.gather(1, index)
if self.norm_topk_prob and self.top_k > 1:
weight = weight / weight.sum(-1, keepdim=True).clamp_min(1e-9)
weight = (weight * self.routed_scaling_factor).to(x.dtype)
out = torch.zeros_like(flat)
for expert in index.unique():
rows, slot = (index == expert).nonzero(as_tuple=True)
taken = flat[rows]
hidden = F.silu(taken @ self.gate_proj[expert].T) * (
taken @ self.up_proj[expert].T
)
out.index_add_(
0, rows, (hidden @ self.down_proj[expert].T) * weight[rows, slot, None]
)
return (out + self.shared_expert(flat)).view(shape)
class KohakuDecoderLayer(nn.Module):
def __init__(self, config: KohakuConfig, layer_idx: int) -> None:
super().__init__()
self.self_attn = KohakuAttention(config, layer_idx)
self.mlp = (
KohakuMLP(config.hidden_size, config.intermediate_size)
if layer_idx < config.first_k_dense
else KohakuMoE(config)
)
self.input_layernorm = KohakuRMSNorm(config.hidden_size, config.rms_norm_eps)
self.post_attention_layernorm = KohakuRMSNorm(
config.hidden_size, config.rms_norm_eps
)
def forward(self, x, cos, sin, attention_mask=None, past_key_values=None, **kwargs):
x = x + self.self_attn(
self.input_layernorm(x), cos, sin, attention_mask, past_key_values, **kwargs
)
return x + self.mlp(self.post_attention_layernorm(x))
class KohakuPreTrainedModel(PreTrainedModel):
config_class = KohakuConfig
base_model_prefix = "model"
supports_gradient_checkpointing = True
_no_split_modules = ["KohakuDecoderLayer"]
_supports_sdpa = True
_supports_cache_class = True
def _init_weights(self, module) -> None:
std = 0.02
if isinstance(module, nn.Linear):
module.weight.data.normal_(mean=0.0, std=std)
if module.bias is not None:
module.bias.data.zero_()
elif isinstance(module, nn.Embedding):
module.weight.data.normal_(mean=0.0, std=std)
elif isinstance(module, KohakuRMSNorm):
module.weight.data.fill_(1.0)
elif isinstance(module, KohakuMoE):
for p in (module.gate_proj, module.up_proj, module.down_proj):
p.data.normal_(mean=0.0, std=std)
class KohakuModel(KohakuPreTrainedModel):
def __init__(self, config: KohakuConfig) -> None:
super().__init__(config)
self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size)
self.layers = nn.ModuleList(
KohakuDecoderLayer(config, i) for i in range(config.num_hidden_layers)
)
self.norm = KohakuRMSNorm(config.hidden_size, config.rms_norm_eps)
self.rotary_emb = KohakuRotary(config)
self.post_init()
def get_input_embeddings(self):
return self.embed_tokens
def set_input_embeddings(self, value) -> None:
self.embed_tokens = value
def forward(
self,
input_ids=None,
attention_mask=None,
position_ids=None,
past_key_values=None,
inputs_embeds=None,
use_cache=None,
**kwargs,
):
if inputs_embeds is None:
inputs_embeds = self.embed_tokens(input_ids)
use_cache = use_cache if use_cache is not None else self.config.use_cache
if use_cache and past_key_values is None:
past_key_values = DynamicCache()
seen = (
past_key_values.get_seq_length()
if isinstance(past_key_values, Cache)
else 0
)
if position_ids is None:
length = inputs_embeds.shape[1]
position_ids = torch.arange(
seen, seen + length, device=inputs_embeds.device
).unsqueeze(0)
mask = None
if attention_mask is not None and attention_mask.dim() == 2:
total = seen + inputs_embeds.shape[1]
pad = attention_mask[:, None, None, :total].bool()
queries = torch.arange(seen, total, device=inputs_embeds.device)[:, None]
keys = torch.arange(total, device=inputs_embeds.device)[None, :]
mask = pad & (keys <= queries)[None, None]
elif attention_mask is not None:
mask = attention_mask
cos, sin = self.rotary_emb(inputs_embeds, position_ids)
hidden = inputs_embeds
for layer in self.layers:
hidden = layer(hidden, cos, sin, mask, past_key_values, **kwargs)
return BaseModelOutputWithPast(
last_hidden_state=self.norm(hidden), past_key_values=past_key_values
)
class KohakuForCausalLM(KohakuPreTrainedModel, GenerationMixin):
_tied_weights_keys = ["lm_head.weight"]
def __init__(self, config: KohakuConfig) -> None:
super().__init__(config)
self.model = KohakuModel(config)
self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
self.post_init()
def get_input_embeddings(self):
return self.model.embed_tokens
def set_input_embeddings(self, value) -> None:
self.model.embed_tokens = value
def get_output_embeddings(self):
return self.lm_head
def set_output_embeddings(self, value) -> None:
self.lm_head = value
def forward(self, input_ids=None, labels=None, **kwargs):
out = self.model(input_ids=input_ids, **kwargs)
logits = self.lm_head(out.last_hidden_state)
loss = None
if labels is not None:
loss = F.cross_entropy(
logits[:, :-1].reshape(-1, logits.shape[-1]).float(),
labels[:, 1:].reshape(-1),
ignore_index=-100,
)
return CausalLMOutputWithPast(
loss=loss, logits=logits, past_key_values=out.past_key_values
)
__all__ = ["KohakuConfig", "KohakuModel", "KohakuForCausalLM", "KohakuPreTrainedModel"]

1
hf/tokenizer.json Normal file

File diff suppressed because one or more lines are too long

9
hf/tokenizer_config.json Normal file
View File

@@ -0,0 +1,9 @@
{
"tokenizer_class": "PreTrainedTokenizerFast",
"bos_token": "<|bos|>",
"eos_token": "<|eos|>",
"pad_token": "<|pad|>",
"unk_token": "<|unk|>",
"model_max_length": 131072,
"clean_up_tokenization_spaces": false
}