初始化项目,由ModelHub XC社区提供模型
Model: KBlueLeaf/TIPOv2-1B-A200M Source: Original Platform
This commit is contained in:
38
.gitattributes
vendored
Normal file
38
.gitattributes
vendored
Normal 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
259
README.md
Normal 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.
|
||||
|
||||
|
||||

|
||||
|
||||
## 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
3
TIPOv2-1B-A200M-f16.gguf
Normal file
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:940a33e1ffe30886a0452c580d0d08124e4ca38655aca5501127432257c29128
|
||||
size 2015578528
|
||||
3
gguf/TIPOv2-1B-A200M-Q8_0.gguf
Normal file
3
gguf/TIPOv2-1B-A200M-Q8_0.gguf
Normal file
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:0a10a20ab475f978ab3bfab9e8b688559a80152d79cb05e39030c051ba708fc0
|
||||
size 1072689632
|
||||
3
gguf/TIPOv2-1B-A200M-f16.gguf
Normal file
3
gguf/TIPOv2-1B-A200M-f16.gguf
Normal file
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:940a33e1ffe30886a0452c580d0d08124e4ca38655aca5501127432257c29128
|
||||
size 2015578528
|
||||
32
hf/config.json
Normal file
32
hf/config.json
Normal 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"
|
||||
}
|
||||
75
hf/configuration_kohaku.py
Normal file
75
hf/configuration_kohaku.py
Normal 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
3
hf/model.safetensors
Normal 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
321
hf/modeling_kohaku.py
Normal 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
1
hf/tokenizer.json
Normal file
File diff suppressed because one or more lines are too long
9
hf/tokenizer_config.json
Normal file
9
hf/tokenizer_config.json
Normal 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
|
||||
}
|
||||
Reference in New Issue
Block a user