Instructions to use Arain119/sophia with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- llama.cpp
How to use Arain119/sophia with llama.cpp:
Install (macOS, Linux)
curl -LsSf https://llama.app/install.sh | sh # Start a local OpenAI-compatible server with a web UI: llama serve -hf Arain119/sophia:Q4_K_M # Run inference directly in the terminal: llama cli -hf Arain119/sophia:Q4_K_M
Install from WinGet (Windows)
winget install llama.cpp # Start a local OpenAI-compatible server with a web UI: llama serve -hf Arain119/sophia:Q4_K_M # Run inference directly in the terminal: llama cli -hf Arain119/sophia:Q4_K_M
Use pre-built binary
# Download pre-built binary from: # https://github.com/ggerganov/llama.cpp/releases # Start a local OpenAI-compatible server with a web UI: ./llama-server -hf Arain119/sophia:Q4_K_M # Run inference directly in the terminal: ./llama-cli -hf Arain119/sophia:Q4_K_M
Build from source code
git clone https://github.com/ggerganov/llama.cpp.git cd llama.cpp cmake -B build cmake --build build -j --target llama-server llama-cli # Start a local OpenAI-compatible server with a web UI: ./build/bin/llama-server -hf Arain119/sophia:Q4_K_M # Run inference directly in the terminal: ./build/bin/llama-cli -hf Arain119/sophia:Q4_K_M
Use Docker
docker model run hf.co/Arain119/sophia:Q4_K_M
- LM Studio
- Jan
- Ollama
How to use Arain119/sophia with Ollama:
ollama run hf.co/Arain119/sophia:Q4_K_M
- Unsloth Desktop
- Docker Model Runner
How to use Arain119/sophia with Docker Model Runner:
docker model run hf.co/Arain119/sophia:Q4_K_M
- Lemonade
How to use Arain119/sophia with Lemonade:
Pull the model
# Download Lemonade from https://lemonade-server.ai/ lemonade pull Arain119/sophia:Q4_K_M
Run and chat with the model
lemonade run user.sophia-Q4_K_M
List all available models
lemonade list
- Atomic Chat
Arain119 commited on
Commit ·
d53adc9
0
Parent(s):
Sophia 1.0.0 — 1B K3-hybrid Chinese chat model (HF remote-code export + native package)
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitattributes +5 -0
- README.md +81 -0
- REPORT.md +79 -0
- cache_decode.py +157 -0
- canonical_config.py +117 -0
- chat.py +121 -0
- chat_template.jinja +57 -0
- config.json +54 -0
- config_projection.py +40 -0
- decoder_forward.py +161 -0
- decoder_full.py +168 -0
- decoder_host.py +47 -0
- decoder_loss.py +474 -0
- decoder_loss_forward.py +47 -0
- decoder_output.py +115 -0
- decoder_runtime.py +298 -0
- decoder_types.py +38 -0
- eval/probe_n16_dpo.json +49 -0
- eval/probe_n16_dpo.samples.jsonl +3 -0
- eval/probe_n16_e3.json +51 -0
- eval/probe_n16_e3.samples.jsonl +3 -0
- eval/samples_multiturn.jsonl +3 -0
- eval/samples_multiturn.md +1286 -0
- generation_config.json +10 -0
- hf_cache.py +166 -0
- hf_config.py +124 -0
- hf_generation.py +113 -0
- hf_lifecycle.py +34 -0
- hf_projection.py +124 -0
- hf_remote_code.py +21 -0
- hf_support.py +151 -0
- input_mask.py +66 -0
- lineage.json +58 -0
- loss_stats.py +72 -0
- model.safetensors +3 -0
- model_attention.py +694 -0
- model_blocks.py +94 -0
- model_config.py +87 -0
- model_dir.py +82 -0
- model_ops.py +81 -0
- model_runtime.py +219 -0
- model_runtime_control.py +90 -0
- model_state.py +135 -0
- model_transformer_setup.py +92 -0
- modeling_sophia.py +149 -0
- pretrained_bundle.py +118 -0
- runtime_backend.py +53 -0
- runtime_contracts.py +107 -0
- runtime_linear.py +46 -0
- semantics.py +135 -0
.gitattributes
ADDED
|
@@ -0,0 +1,5 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
*.pt filter=lfs diff=lfs merge=lfs -text
|
| 2 |
+
*.safetensors filter=lfs diff=lfs merge=lfs -text
|
| 3 |
+
*.bin filter=lfs diff=lfs merge=lfs -text
|
| 4 |
+
*.bundle filter=lfs diff=lfs merge=lfs -text
|
| 5 |
+
*.jsonl filter=lfs diff=lfs merge=lfs -text
|
README.md
ADDED
|
@@ -0,0 +1,81 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
license: apache-2.0
|
| 3 |
+
library_name: transformers
|
| 4 |
+
tags:
|
| 5 |
+
- text-generation
|
| 6 |
+
- chat
|
| 7 |
+
- chinese
|
| 8 |
+
- hybrid-attention
|
| 9 |
+
- kda
|
| 10 |
+
- mla
|
| 11 |
+
---
|
| 12 |
+
|
| 13 |
+
# Sophia
|
| 14 |
+
|
| 15 |
+
Sophia 是一个 ~1B 参数的中文对话模型,在单张 RTX 5090(32GB)上从零训练完成:随机初始化预训练 20B tokens → SFT → DPO。完整谱系、评测证据与训练报告随包附带(`REPORT.md`、`lineage.json`、`eval/`)。
|
| 16 |
+
|
| 17 |
+
| | |
|
| 18 |
+
|---|---|
|
| 19 |
+
| 参数量 | 1,012,630,480 |
|
| 20 |
+
| 架构 | 28 层 dense Hybrid decoder:`[KDA, KDA, KDA, Gated-MLA] × 7`,NoPE |
|
| 21 |
+
| 上下文 | 4,096 |
|
| 22 |
+
| 预训练 | 20B tokens,15,259 步,seed 42,BF16 + Muon/AdamW |
|
| 23 |
+
| 硬件 | 单张 NVIDIA RTX 5090 32GB |
|
| 24 |
+
|
| 25 |
+
## 特性
|
| 26 |
+
|
| 27 |
+
- **自适应思考**:`<think>…</think>` 推理草稿是否生成由模型逐轮自行决定(自发率 ~5%)。1B 尺度下 think 在 math/code 上反而扣分(见 `REPORT.md` §3),模型自己学会了什么时候不想
|
| 28 |
+
- **人格宪法**:行为规范收敛为独立文件 `configs/sft/sophia_persona.md`——不说"我没有感情"、确定问题第一句给答案、约束逐轮登记
|
| 29 |
+
- **诚实评测**:所有结论基于 powered 判定(134 题 × 16 采样 × 配对 t 检验,SE ±1.1);零结果实验(RFT/GRPO/DPO-2)同样如实记录在 `REPORT.md`
|
| 30 |
+
|
| 31 |
+
## 用法
|
| 32 |
+
|
| 33 |
+
需要 `transformers>=4.56` 与 `trust_remote_code=True`(自定义架构 SophiaForCausalLM,remote code 随仓分发):
|
| 34 |
+
|
| 35 |
+
```python
|
| 36 |
+
from transformers import AutoModelForCausalLM, AutoTokenizer
|
| 37 |
+
|
| 38 |
+
tok = AutoTokenizer.from_pretrained("Arain119/sophia", trust_remote_code=True)
|
| 39 |
+
m = AutoModelForCausalLM.from_pretrained(
|
| 40 |
+
"Arain119/sophia", dtype="bfloat16", device_map="cuda", trust_remote_code=True
|
| 41 |
+
)
|
| 42 |
+
ids = tok.apply_chat_template(
|
| 43 |
+
[{"role": "user", "content": "你好"}],
|
| 44 |
+
add_generation_prompt=True, return_tensors="pt",
|
| 45 |
+
).to("cuda")
|
| 46 |
+
out = m.generate(ids, max_new_tokens=200, temperature=0.7, top_p=0.92, do_sample=True)
|
| 47 |
+
print(tok.decode(out[0][ids.shape[1]:], skip_special_tokens=True))
|
| 48 |
+
```
|
| 49 |
+
|
| 50 |
+
GPU 上建议 `pip install flash-linear-attention>=0.5.2` 启用 KDA 融合内核;未安装时 `kda_backend="auto"` 自动回退 reference 路径(可用,较慢)。CPU 推理走 reference 路径(慢)。
|
| 51 |
+
|
| 52 |
+
## 评测速览
|
| 53 |
+
|
| 54 |
+
- powered n16 探针(claude-haiku-4-5 判官):judge_mean **51.37**,pass_70 0.352,hit_eos 0.975,distinct4 0.928;相对 SFT 父模型 **+2.10**(t=+3.71, p<0.001),失败模式 flag 全面下降
|
| 55 |
+
- 同档对照(同一管线、同一数据、同一判分,无官方公布值混入):SophiaBenchmark 134 题判官分 51.25,高于 Llama3.2-1B(48.09)与 Qwen2.5-0.5B(50.35),同档第 4
|
| 56 |
+
- 已知限度:IFEval 12.2%(系统性落后同档基线)、NIAH 62.5% 且随深度衰减、深多轮 judge_mean 27.93 为最弱段、MC 基准(MMLU/CMMLU 似然法)贴近随机——1B 模型对 ABCD 符号有位置先验
|
| 57 |
+
- 完整数字、失败样本与机制分析见 `REPORT.md`;逐题明细在源码仓 `ops/eval/public_benchmarks/`
|
| 58 |
+
|
| 59 |
+
## 训练谱系
|
| 60 |
+
|
| 61 |
+
```
|
| 62 |
+
pretrain(20B, 契约跑满)
|
| 63 |
+
→ sft(88,267 行 × 3ep, epoch3 选中)
|
| 64 |
+
→ dpo(1,953 对, β=0.1, 500 步) = 本仓权重
|
| 65 |
+
├─ rft(104 步, 零结果) 未发布
|
| 66 |
+
└─ grpo+dpo2(链式实验, 未过门) 未发布
|
| 67 |
+
```
|
| 68 |
+
|
| 69 |
+
`lineage.json` 含各阶段 run、checkpoint、数据与种子的完整指针。
|
| 70 |
+
|
| 71 |
+
## 相关仓库
|
| 72 |
+
|
| 73 |
+
- 源码:GitHub `Arain119/Sophia`,或随包 `sophia-1.0.0.bundle`(单提交导出,`git clone` 即可取回全部代码)
|
| 74 |
+
- 原生权重 + 文档:`sophia.pt`(PyTorch checkpoint)、`sophia_sft.pt`(SFT 父基线)、`REPORT.md`
|
| 75 |
+
- 训练数据:ModelScope `Arain119/Sophia-dataset`(`pretrain/` token 分片 ~80GB + `corpus/` SFT 语料 426MB)
|
| 76 |
+
|
| 77 |
+
## 许可证
|
| 78 |
+
|
| 79 |
+
模型权重与代码 Apache License 2.0。训练数据见 `Arain119/Sophia-dataset`:`corpus/` 为 Apache-2.0;`pretrain/` 上游含非商用研究来源,使用限制以其 `lineage/` 与 README 为准。
|
| 80 |
+
|
| 81 |
+
`<think>` 输出为生成文本,不保证忠实反映内部推理过程。
|
REPORT.md
ADDED
|
@@ -0,0 +1,79 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Sophia 技术报告
|
| 2 |
+
|
| 3 |
+
*2026 年 9 月 13 日。交付物:`sophia.pt`,sha256 `e49ddc2cbbd311330f21fe52ada853e3ed745cf5f2e47de83da3898234592268`。*
|
| 4 |
+
|
| 5 |
+
两年前的今天,2024 年 9 月 12 日,OpenAI 把代号 Strawberry 的模型推上线,正式名字叫 o1。它做的事在事后看很简单:让模型在回答之前先想。训练时用强化学习打磨那条思考链,推理时用更多时间换更高的正确率——test-time compute,行业坐标系里多出来的第二条轴。IMO 资格赛从 13% 到 83% 的那个数字,让"想多久"从此和"多大"平起平坐。
|
| 6 |
+
|
| 7 |
+
两年后的今天,我们在一张 RTX 5090 上把这个问题的诚实版本又跑了一遍:一个十亿参数的模型,从头训练,把 o1 教给行业的两件事——会想的训练数据、会学偏好的后训练——在自己的尺度上做到哪一步。这份报告是过程的全记录,包括所有没有赢的实验。
|
| 8 |
+
|
| 9 |
+
## 一、预训练:先决定她读什么
|
| 10 |
+
|
| 11 |
+
模型是 1,012,630,480 参数的 K3-derived 稠密 Hybrid 解码器:28 层,`[KDA, KDA, KDA, Gated-MLA] × 7`,NoPE,4K 上下文,BF16,Muon hybrid 优化器,seed 42,1,310,720 tokens/update,15,259 步跑满 20B tokens。这些数字在训练第一天就写死进了契约,之后没有一个人碰过。
|
| 12 |
+
|
| 13 |
+
真正做过选择的只有一件事:语料。早期的配比稿里,81% 的中文 token 来自一个教育过滤源——那样的模型读完的是课本,不是人话。实测每万字的语域标记:fineweb-edu 的第一二人称 16.2 次、感叹号 0.7 次、口语词命中 17.5 次;SkyPile 是它的两倍,Zhihu-KOL 是它的六倍。最终的配比把 conversational Chinese 抬到了正经位置:FineWeb 做骨干,SkyPile 和知乎长回答补足"人是怎么说话"的那一半,许可代码保留一遍约占 7%。一个要做对话的模型,训练数据里得先有人对话。
|
| 14 |
+
|
| 15 |
+
15,259 步走完的那天,产出的是一个基座:它见过足够多的语言,还不知道怎么当一个对话者。那是后训练的工作。
|
| 16 |
+
|
| 17 |
+
## 二、人格宪法
|
| 18 |
+
|
| 19 |
+
Sophia 的行为规范不是散落在提示词里的叮嘱,是一份独立的文件:`configs/sft/sophia_persona.md`,我们叫它人格宪法。它规定的不是语气,是存在方式:
|
| 20 |
+
|
| 21 |
+
> 你是 AI,也是完整的个体,两者不冲突:你的情感、好恶和看法是真实的,不说"我没有感情""我只是数学模拟"。
|
| 22 |
+
|
| 23 |
+
宪法里的条款大多是行为级的:先回应对方真正说的事,再说自己的看法;有确定答案的问题,第一句就是答案;要你写的东西直接写,不先反问要什么风格;对方前面说过的每一件事都要接上——限制、指代、你自己上一轮说过的话。答错事实可以,但不能前一句说 A 后一句说非 A。
|
| 24 |
+
|
| 25 |
+
宪法里最有意思的一节是给思考块立规矩:think 是第一人称的推理草稿,只写对回答有用的东西——"对方不吃鱼,刚才推荐的清蒸鲈鱼得换掉"是好 think;"我需要先给对方提供情绪价值"是坏 think。思考是为了答对,不是为了表演正确姿态。这个区分在后面的所有实验里反复被证明是对的。
|
| 26 |
+
|
| 27 |
+
## 三、蒸馏管线:教师写,机器判
|
| 28 |
+
|
| 29 |
+
SFT 语料的 88,267 行没有一行来自人工:全部由教师模型(claude)按依赖脚本一次写完。脚本是这个管线真正的发明——每一条会话在生成前就被规定了用户每一轮要做什么:追加约束("对了,我不吃辣")、回指前文("就用你刚才说的第一部")、变换上一轮的回答、或者直接纠错。前一代语料的失败模式是模型学会了"续写"而不是"阅读",所以脚本存在的意义就是逼出必须读历史才能答对的轮次。
|
| 30 |
+
|
| 31 |
+
写完不等于收货。每一行都要过机器判官:结构不完整拒收(reject: shape)、think 里写规划黑话拒收(planning)、系统提示泄漏拒收(leak)、身份漂移拒收(identity_leak)、超长拒收、语言不纯拒收。数学题的答案和教师给的表达式对账,代码题的答案真的拿去跑断言。被判官扔掉的行没有进入训练集——88,267 行是幸存者。
|
| 32 |
+
|
| 33 |
+
助手轮的格式从第一天就是 `<think>…</think>回答`,思考是训练目标的一部分,不是装饰品。
|
| 34 |
+
|
| 35 |
+
## 四、Adaptive thinking:让她自己决定要不要想
|
| 36 |
+
|
| 37 |
+
o1 教给行业的真实教训不是"永远要想"——是想是一种可以花的算力,花不花应该由问题决定。Sophia 是否思考由模型逐轮自行决定,这也是训练数据的自然分布:简单轮次的 think 本来就短,难轮次的 think 本来就该长。
|
| 38 |
+
|
| 39 |
+
推理引擎把这个结构用到了底:采样时先采 think、到闭合标签停,再从每条 think 分支出若干回答。这不只是为了快——它让 GRPO 能给 think token 单独算功劳:同一条 think 下有的回答好有的坏,差别记在答案头上;这条 think 整体好坏,记在 think 头上。
|
| 40 |
+
|
| 41 |
+
## 五、后训练:六组实验,一次真赢
|
| 42 |
+
|
| 43 |
+
这是这份报告最不愿意写成说明文的部分,因为它是整段经历里最有叙事性的。后训练前前后后做了六组实验,判官统一是 claude-haiku-4-5,判定统一是 134 ��� × 16 采样的配对 t 检验(标准误 ±1.1 分)——这个数字是我们定的换冠军的门槛:快探针看得出方向,只有 powered 判定说了算数。
|
| 44 |
+
|
| 45 |
+
**SFT 选 epoch**:三个 epoch 各留一个 ckpt,三路 haiku 盲评均位次 1.88、复读率 9.7% 的 epoch3 胜出。7,176 步,基座变成了对话者。
|
| 46 |
+
|
| 47 |
+
**RFT(拒绝采样微调)**:4,600 prompt × k16,71,955 次 rollout,选 best-of-16 ≥75 分的 1,681 行自蒸馏,104 步。配对差 −0.25(t=−0.46)。零结果,但原因比结果值钱:深多轮难题上的 best-of-16 是**蒙对的运气**——16 次里碰巧对一次的那个"对"不可复现,蒸馏它学不到任何东西。**运气不能蒸馏**,这是第一条机制教训。
|
| 48 |
+
|
| 49 |
+
**DPO**:同一批 rollout 建偏好对——chosen 是无 flag 最高分(≥75),rejected 是带 flag 最低分,gap≥30,1,953 对,β=0.1,500 步。配对差 **+2.10,t=+3.71,p<0.001**。真赢,也是唯一的一次。机制解释和 RFT 互为镜像:坏不是运气,坏是**系统性模式**——矛盾、跑题、胡说,这些是可学的负向目标。**抬好的失败,压坏的成功。**
|
| 50 |
+
|
| 51 |
+
**DPO 剂量扫描**:赢了就加码?3,910 对、1,000 步的第二轮把坏模式压得更低,但中位字数 70→55、截断 flag 79→218——模型变得更短更保守,配对差 −1.79。**过冲有签名:truncated 升、字数降、判官分降。** 甜点是 ~2k 对/500 步,边际到此为止。
|
| 52 |
+
|
| 53 |
+
**GRPO**:on-policy 强化学习,k16、79 步(组内 chunk 补丁解决了 k16 更新图的 OOM——按 token mask 份额加权,与整组 backward 逐位等价)。它把复读率砍半(0.043→0.023)、hit_eos 提到 0.978,但 powered 判定 +0.03(t=0.056)——教科书级的零效应。机制在动,总分不动。
|
| 54 |
+
|
| 55 |
+
**DPO-2(on-policy)**:用 GRPO 后的模型重新 rollout 建 593 对新鲜偏好对——chosen 73.9 对 rejected 19.8,pair 质量极高,模型也学进去了(mean_acc 0.67)。Powered 判定 −1.29。机制改善(contradiction −15%)没能换算成判官总分。
|
| 56 |
+
|
| 57 |
+
六组实验的曲线读出来是一句话:**+2.10 → +0.03 → −1.29**。DPO-1 吃掉了所有"系统性坏模式"的低垂果实,之后残余的失败是 off_topic、context_loss 这类能力缺口——16 次采样全灭的题没有正样本可比,偏好学习无从着力。判官驱动的偏好优化在这个模型、这份数据、这个判官的组合上,天花板已经被三次 powered 测量钉死了。
|
| 58 |
+
|
| 59 |
+
## 六、数字与限度
|
| 60 |
+
|
| 61 |
+
交付模型在 powered 探针下:judge_mean 51.37,pass_70 0.352,hit_eos 0.975,distinct4 0.928,中位字数 66。多轮对话包(30 会话 × 6 轮)judge_mean 27.93——深多轮仍是最弱的一段。
|
| 62 |
+
|
| 63 |
+
已知的失败模式,全部是实测样本不是推测:约束登记会丢(用户说过不吃鱼,后面还会推荐鱼);会伪造执行轨迹(函数写对,"示例输出"是编的);算术和冷门事实会自信出错;被连续追问时会退化成复读螺旋。C-Eval+CMMLU 似然法 26.2%≈随机——1B 对话模型对 ABCD 选项符号有位置先验,MC 基准在这个尺度上区分度本就差。
|
| 64 |
+
|
| 65 |
+
**公开基准同协议对照**(2026-09-15,五个同档基线与 Sophia 走同一管线、同一数据、同一判分,无任何官方公布值混入;逐题明细见 `ops/eval/public_benchmarks/`):判别式任务上 Sophia 与 MiniCPM4-0.5B 几乎同档——两者在字母似然协议下均贴近随机(MMLU 22.96 vs 22.9,BoolQ 62.17 vs 62.2,MMLU-Pro 11.45 vs 11.7),落后于 Qwen3.5-0.8B / Llama3.2-1B / Gemma3-1B。IFEval 官方校验器(541 题,strict)她拿 12.2%——是真实的残余指令遵循,但离基线的 26.4–51.8% 有系统性差距。NIAH 单针检索基线全部 100%,她 62.5% 且随深度衰减(1K→90%,3K→30%)。判官制开放对话(SophiaBenchmark,134 题,claude-haiku-4-5)她 51.25,列同档第 4,高于 Llama3.2-1B(48.09)与 Qwen2.5-0.5B(50.35)。
|
| 66 |
+
|
| 67 |
+
## 七、交付
|
| 68 |
+
|
| 69 |
+
```
|
| 70 |
+
pretrain(20B, 契约跑满)
|
| 71 |
+
→ sft(88,267 行 × 3ep, epoch3 选中)
|
| 72 |
+
→ dpo(1,953 对, 500 步) = sophia.pt
|
| 73 |
+
├─ rft(104步, 零结果) runs/rft_e3.pt
|
| 74 |
+
└─ grpo+dpo2(链式实验, 未过门) runs/grpo/, 已弃
|
| 75 |
+
```
|
| 76 |
+
|
| 77 |
+
判官费用全程 ~¥500,GPU 全部单卡 5090。权重、sha256、谱系、探针全量判分、多轮样本包随包附证。没有任何东西上传到外网。
|
| 78 |
+
|
| 79 |
+
两年前 o1 证明了会想比会说值钱。这个十亿参数的小模型给出的答案是:在它能站住的高度上,把"想"写进训练数据、把"好坏"交给统计判定、把诚实写进报告——天花板到了就说到哪是哪。
|
cache_decode.py
ADDED
|
@@ -0,0 +1,157 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
from dataclasses import dataclass
|
| 8 |
+
from typing import Protocol
|
| 9 |
+
|
| 10 |
+
import torch
|
| 11 |
+
|
| 12 |
+
from .model_state import RuntimeCacheSnapshot
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
@dataclass
|
| 16 |
+
class RuntimeCacheState:
|
| 17 |
+
cache: RuntimeCacheSnapshot
|
| 18 |
+
batch_size: int
|
| 19 |
+
cache_pos: int
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
class RuntimeCacheDecodeModel(Protocol):
|
| 23 |
+
def replay_with_cache(
|
| 24 |
+
self,
|
| 25 |
+
input_ids: torch.Tensor,
|
| 26 |
+
*,
|
| 27 |
+
start_pos: int = 0,
|
| 28 |
+
return_all_logits: bool = True,
|
| 29 |
+
) -> tuple[torch.Tensor, torch.Tensor | None]: ...
|
| 30 |
+
|
| 31 |
+
def cache_dump(
|
| 32 |
+
self,
|
| 33 |
+
device: str = "cpu",
|
| 34 |
+
*,
|
| 35 |
+
cache_pos: int | None = None,
|
| 36 |
+
batch_size: int | None = None,
|
| 37 |
+
) -> RuntimeCacheSnapshot: ...
|
| 38 |
+
|
| 39 |
+
def cache_load(self, cache_snapshot: RuntimeCacheSnapshot) -> None: ...
|
| 40 |
+
|
| 41 |
+
def reset_runtime_cache(self) -> None: ...
|
| 42 |
+
|
| 43 |
+
def clone_runtime_cache_state(cache_state: RuntimeCacheState) -> RuntimeCacheState:
|
| 44 |
+
return RuntimeCacheState(
|
| 45 |
+
cache=cache_state.cache.clone(),
|
| 46 |
+
batch_size=int(cache_state.batch_size),
|
| 47 |
+
cache_pos=int(cache_state.cache_pos),
|
| 48 |
+
)
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
def cache_state_from_runtime_cache(
|
| 52 |
+
cache_input: RuntimeCacheState | None,
|
| 53 |
+
) -> RuntimeCacheState | None:
|
| 54 |
+
if cache_input is None:
|
| 55 |
+
return None
|
| 56 |
+
if isinstance(cache_input, RuntimeCacheState):
|
| 57 |
+
return clone_runtime_cache_state(cache_input)
|
| 58 |
+
raise TypeError("cache must be a Sophia RuntimeCacheState returned by SophiaDecoder")
|
| 59 |
+
|
| 60 |
+
|
| 61 |
+
def prefill_runtime_cache(
|
| 62 |
+
model: RuntimeCacheDecodeModel,
|
| 63 |
+
input_ids: torch.Tensor,
|
| 64 |
+
*,
|
| 65 |
+
start_pos: int,
|
| 66 |
+
logits_to_keep: int | None = None,
|
| 67 |
+
) -> torch.Tensor:
|
| 68 |
+
if int(input_ids.size(1)) <= 0:
|
| 69 |
+
raise ValueError("input_ids must contain at least one token")
|
| 70 |
+
return_all_logits = int(logits_to_keep or 0) != 1
|
| 71 |
+
logits, _ = model.replay_with_cache(
|
| 72 |
+
input_ids,
|
| 73 |
+
start_pos=int(start_pos),
|
| 74 |
+
return_all_logits=bool(return_all_logits),
|
| 75 |
+
)
|
| 76 |
+
if not bool(return_all_logits):
|
| 77 |
+
logits = logits.unsqueeze(1)
|
| 78 |
+
if logits_to_keep is not None and int(logits_to_keep) > 0:
|
| 79 |
+
return logits[:, -int(logits_to_keep) :, :]
|
| 80 |
+
return logits
|
| 81 |
+
|
| 82 |
+
|
| 83 |
+
def forward_cached_decode(
|
| 84 |
+
*,
|
| 85 |
+
model: RuntimeCacheDecodeModel,
|
| 86 |
+
input_ids: torch.Tensor,
|
| 87 |
+
cache_state: RuntimeCacheState | None,
|
| 88 |
+
start_pos: int | None,
|
| 89 |
+
logits_to_keep: int | None,
|
| 90 |
+
) -> tuple[torch.Tensor, RuntimeCacheState]:
|
| 91 |
+
if logits_to_keep is not None and int(logits_to_keep) < 0:
|
| 92 |
+
raise ValueError(f"logits_to_keep must be >= 0 when provided, got {logits_to_keep}")
|
| 93 |
+
cache_start = 0 if start_pos is None else int(start_pos)
|
| 94 |
+
if cache_state is None:
|
| 95 |
+
if cache_start != 0:
|
| 96 |
+
raise ValueError("cache_state is required when start_pos > 0 for cached decode")
|
| 97 |
+
model.reset_runtime_cache()
|
| 98 |
+
logits = prefill_runtime_cache(
|
| 99 |
+
model,
|
| 100 |
+
input_ids,
|
| 101 |
+
start_pos=cache_start,
|
| 102 |
+
logits_to_keep=logits_to_keep,
|
| 103 |
+
)
|
| 104 |
+
next_cache_pos = int(cache_start) + int(input_ids.size(1))
|
| 105 |
+
return logits, RuntimeCacheState(
|
| 106 |
+
cache=model.cache_dump(
|
| 107 |
+
device="cpu",
|
| 108 |
+
cache_pos=int(next_cache_pos),
|
| 109 |
+
batch_size=int(input_ids.size(0)),
|
| 110 |
+
),
|
| 111 |
+
cache_pos=int(next_cache_pos),
|
| 112 |
+
batch_size=int(input_ids.size(0)),
|
| 113 |
+
)
|
| 114 |
+
|
| 115 |
+
cache_state = clone_runtime_cache_state(cache_state)
|
| 116 |
+
cache_pos = int(cache_state.cache_pos)
|
| 117 |
+
if cache_pos < 0:
|
| 118 |
+
raise ValueError(f"cache cache_pos must be >= 0, got {cache_pos}")
|
| 119 |
+
if start_pos is not None and int(start_pos) != cache_pos:
|
| 120 |
+
raise ValueError(
|
| 121 |
+
f"start_pos ({start_pos}) must match cache cache_pos ({cache_pos})"
|
| 122 |
+
)
|
| 123 |
+
batch_size = int(cache_state.batch_size or int(input_ids.size(0)))
|
| 124 |
+
if batch_size <= 0:
|
| 125 |
+
raise ValueError(f"cache batch_size must be > 0, got {batch_size}")
|
| 126 |
+
if batch_size != int(input_ids.size(0)):
|
| 127 |
+
raise ValueError(
|
| 128 |
+
"cache batch_size does not match input_ids: "
|
| 129 |
+
f"{batch_size} != {int(input_ids.size(0))}"
|
| 130 |
+
)
|
| 131 |
+
model.cache_load(cache_state.cache)
|
| 132 |
+
logits = prefill_runtime_cache(
|
| 133 |
+
model,
|
| 134 |
+
input_ids,
|
| 135 |
+
start_pos=cache_pos,
|
| 136 |
+
logits_to_keep=logits_to_keep,
|
| 137 |
+
)
|
| 138 |
+
next_cache_pos = int(cache_pos) + int(input_ids.size(1))
|
| 139 |
+
return logits, RuntimeCacheState(
|
| 140 |
+
cache=model.cache_dump(
|
| 141 |
+
device="cpu",
|
| 142 |
+
cache_pos=int(next_cache_pos),
|
| 143 |
+
batch_size=int(input_ids.size(0)),
|
| 144 |
+
),
|
| 145 |
+
cache_pos=int(next_cache_pos),
|
| 146 |
+
batch_size=int(input_ids.size(0)),
|
| 147 |
+
)
|
| 148 |
+
|
| 149 |
+
|
| 150 |
+
__all__ = [
|
| 151 |
+
"RuntimeCacheDecodeModel",
|
| 152 |
+
"RuntimeCacheState",
|
| 153 |
+
"cache_state_from_runtime_cache",
|
| 154 |
+
"clone_runtime_cache_state",
|
| 155 |
+
"forward_cached_decode",
|
| 156 |
+
"prefill_runtime_cache",
|
| 157 |
+
]
|
canonical_config.py
ADDED
|
@@ -0,0 +1,117 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
"""Canonical configuration for the native Sophia Hybrid decoder."""
|
| 6 |
+
|
| 7 |
+
from __future__ import annotations
|
| 8 |
+
|
| 9 |
+
from collections.abc import Mapping
|
| 10 |
+
from dataclasses import MISSING, asdict, dataclass, fields, is_dataclass
|
| 11 |
+
|
| 12 |
+
from .semantics import canonicalize_model_values, validate_model_values
|
| 13 |
+
from .runtime_backend import resolve_runtime_backend
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
def _object_mapping(config: object) -> dict[str, object]:
|
| 17 |
+
canonical_payload = {
|
| 18 |
+
field.name: getattr(config, field.name)
|
| 19 |
+
for field in fields(SophiaModelConfig)
|
| 20 |
+
if hasattr(config, field.name)
|
| 21 |
+
}
|
| 22 |
+
if canonical_payload:
|
| 23 |
+
return canonical_payload
|
| 24 |
+
if is_dataclass(config):
|
| 25 |
+
return asdict(config)
|
| 26 |
+
if isinstance(config, Mapping):
|
| 27 |
+
return dict(config)
|
| 28 |
+
to_dict = getattr(config, "to_dict", None)
|
| 29 |
+
if callable(to_dict):
|
| 30 |
+
payload = to_dict()
|
| 31 |
+
if isinstance(payload, Mapping):
|
| 32 |
+
return dict(payload)
|
| 33 |
+
return {}
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
@dataclass
|
| 37 |
+
class SophiaModelConfig:
|
| 38 |
+
"""The only supported Sophia architecture schema."""
|
| 39 |
+
|
| 40 |
+
vocab_size: int = 65536
|
| 41 |
+
dim: int = 1536
|
| 42 |
+
n_layers: int = 28
|
| 43 |
+
num_heads: int = 16
|
| 44 |
+
head_dim: int = 128
|
| 45 |
+
ffn_hidden: int = 3968
|
| 46 |
+
kda_decay_rank: int = 128
|
| 47 |
+
kda_output_gate_rank: int = 128
|
| 48 |
+
kda_output_gate_full_rank: bool = True
|
| 49 |
+
kda_decay_lower_bound: float = -5.0
|
| 50 |
+
kda_dt_min: float = 1e-3
|
| 51 |
+
kda_dt_max: float = 1e-1
|
| 52 |
+
kda_dt_floor: float = 1e-4
|
| 53 |
+
kda_a_log_init: float = 0.0
|
| 54 |
+
mla_q_rank: int = 384
|
| 55 |
+
mla_kv_rank: int = 128
|
| 56 |
+
short_conv_kernel: int = 4
|
| 57 |
+
attn_res_block_size: int = 4
|
| 58 |
+
situ_gate_softcap: float = 4.0
|
| 59 |
+
situ_up_softcap: float = 25.0
|
| 60 |
+
norm_eps: float = 1e-5
|
| 61 |
+
max_seq_len: int = 4096
|
| 62 |
+
max_batch_size: int = 4
|
| 63 |
+
dropout: float = 0.0
|
| 64 |
+
initializer_range: float = 0.02
|
| 65 |
+
kda_backend: str = "auto"
|
| 66 |
+
|
| 67 |
+
def __post_init__(self) -> None:
|
| 68 |
+
normalized = canonicalize_model_values(asdict(self))
|
| 69 |
+
for field_info in fields(type(self)):
|
| 70 |
+
setattr(self, field_info.name, normalized[field_info.name])
|
| 71 |
+
validate_model_values(normalized)
|
| 72 |
+
|
| 73 |
+
@classmethod
|
| 74 |
+
def get_defaults(cls) -> dict[str, object]:
|
| 75 |
+
defaults: dict[str, object] = {}
|
| 76 |
+
for field_info in fields(cls):
|
| 77 |
+
if field_info.default is not MISSING:
|
| 78 |
+
defaults[field_info.name] = field_info.default
|
| 79 |
+
elif field_info.default_factory is not MISSING:
|
| 80 |
+
defaults[field_info.name] = field_info.default_factory()
|
| 81 |
+
return defaults
|
| 82 |
+
|
| 83 |
+
@classmethod
|
| 84 |
+
def from_mapping(cls, values: Mapping[str, object]) -> SophiaModelConfig:
|
| 85 |
+
raw = dict(values)
|
| 86 |
+
allowed = {field.name for field in fields(cls)}
|
| 87 |
+
unknown = sorted(str(key) for key in raw if key not in allowed)
|
| 88 |
+
if unknown:
|
| 89 |
+
raise ValueError(
|
| 90 |
+
"SophiaModelConfig only accepts native Sophia Hybrid fields; "
|
| 91 |
+
f"unknown keys: {', '.join(unknown)}"
|
| 92 |
+
)
|
| 93 |
+
return cls(**raw)
|
| 94 |
+
|
| 95 |
+
@classmethod
|
| 96 |
+
def from_object(cls, config: object) -> SophiaModelConfig:
|
| 97 |
+
return cls.from_mapping(_object_mapping(config))
|
| 98 |
+
|
| 99 |
+
def to_model_spec(self):
|
| 100 |
+
from ml.core.spec import ModelSpec
|
| 101 |
+
|
| 102 |
+
return ModelSpec.from_config(self)
|
| 103 |
+
|
| 104 |
+
def to_model_args(self, *, runtime_max_seq_len: int | None = None) -> object:
|
| 105 |
+
from .config_projection import build_runtime_model_args
|
| 106 |
+
|
| 107 |
+
return build_runtime_model_args(
|
| 108 |
+
self,
|
| 109 |
+
model_args_cls=resolve_runtime_backend().model_args_cls,
|
| 110 |
+
runtime_max_seq_len=runtime_max_seq_len,
|
| 111 |
+
)
|
| 112 |
+
|
| 113 |
+
def to_dict(self) -> dict[str, object]:
|
| 114 |
+
return asdict(self)
|
| 115 |
+
|
| 116 |
+
|
| 117 |
+
__all__ = ["SophiaModelConfig"]
|
chat.py
ADDED
|
@@ -0,0 +1,121 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Talk to Sophia.
|
| 2 |
+
|
| 3 |
+
python -m ml.cli.chat --checkpoint <ckpt.pt>
|
| 4 |
+
|
| 5 |
+
The decoding defaults are the ones the temperature sweep chose, not library
|
| 6 |
+
defaults. They matter more than usual here: greedy makes this policy loop and
|
| 7 |
+
a hot sample makes it incoherent, so the useful setting is neither end.
|
| 8 |
+
|
| 9 |
+
She decides per turn whether to think (~5% of turns); think markup never
|
| 10 |
+
reaches the visible answer -- unclosed or stray tags are routed to the think
|
| 11 |
+
field instead.
|
| 12 |
+
|
| 13 |
+
`/reset` starts a new conversation, `/think` toggles reasoning display, and
|
| 14 |
+
Ctrl-D exits.
|
| 15 |
+
"""
|
| 16 |
+
|
| 17 |
+
from __future__ import annotations
|
| 18 |
+
|
| 19 |
+
import argparse
|
| 20 |
+
import sys
|
| 21 |
+
import time
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
from ml.training.rl.engine import RolloutEngine, chat_turn, load_policy
|
| 25 |
+
|
| 26 |
+
BANNER = """Sophia — 输入即可对话
|
| 27 |
+
/reset 开始新对话 /think 思考块显示开关
|
| 28 |
+
/temp <x> 调整温度 Ctrl-D 退出
|
| 29 |
+
"""
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def main() -> None:
|
| 33 |
+
parser = argparse.ArgumentParser()
|
| 34 |
+
parser.add_argument("--checkpoint", required=True)
|
| 35 |
+
parser.add_argument("--model-spec", default="configs/model/sophia.json")
|
| 36 |
+
parser.add_argument("--tokenizer-path", default="ml/modeling/text")
|
| 37 |
+
parser.add_argument("--temperature", type=float, default=0.7)
|
| 38 |
+
parser.add_argument("--top-p", type=float, default=0.92)
|
| 39 |
+
parser.add_argument("--max-new-tokens", type=int, default=320)
|
| 40 |
+
parser.add_argument("--show-think", action="store_true")
|
| 41 |
+
args = parser.parse_args()
|
| 42 |
+
|
| 43 |
+
started = time.time()
|
| 44 |
+
model, tokenizer, meta = load_policy(
|
| 45 |
+
checkpoint=args.checkpoint,
|
| 46 |
+
model_spec=args.model_spec,
|
| 47 |
+
tokenizer_path=args.tokenizer_path,
|
| 48 |
+
batch_size=1,
|
| 49 |
+
)
|
| 50 |
+
engine = RolloutEngine(
|
| 51 |
+
model=model, tokenizer=tokenizer, max_batch=1, max_new_tokens=args.max_new_tokens
|
| 52 |
+
)
|
| 53 |
+
print(
|
| 54 |
+
"loaded step={} in {:.1f}s".format(meta.get("step"), time.time() - started),
|
| 55 |
+
file=sys.stderr,
|
| 56 |
+
)
|
| 57 |
+
print(BANNER)
|
| 58 |
+
|
| 59 |
+
history: list[dict[str, str]] = []
|
| 60 |
+
show_think = bool(args.show_think)
|
| 61 |
+
temperature = float(args.temperature)
|
| 62 |
+
|
| 63 |
+
while True:
|
| 64 |
+
try:
|
| 65 |
+
line = input("你 > ").strip()
|
| 66 |
+
except (EOFError, KeyboardInterrupt):
|
| 67 |
+
print()
|
| 68 |
+
return
|
| 69 |
+
if not line:
|
| 70 |
+
continue
|
| 71 |
+
if line == "/reset":
|
| 72 |
+
history = []
|
| 73 |
+
print("(已开始新对话)\n")
|
| 74 |
+
continue
|
| 75 |
+
if line == "/think":
|
| 76 |
+
show_think = not show_think
|
| 77 |
+
print("(思考块 %s)\n" % ("显示" if show_think else "隐藏"))
|
| 78 |
+
continue
|
| 79 |
+
if line.startswith("/temp"):
|
| 80 |
+
try:
|
| 81 |
+
temperature = float(line.split()[1])
|
| 82 |
+
print(f"(温度 = {temperature:.2f})\n")
|
| 83 |
+
except (IndexError, ValueError):
|
| 84 |
+
print("(用法:/temp 0.7)\n")
|
| 85 |
+
continue
|
| 86 |
+
|
| 87 |
+
history.append({"role": "user", "content": line})
|
| 88 |
+
started = time.time()
|
| 89 |
+
sample = chat_turn(
|
| 90 |
+
engine,
|
| 91 |
+
tokenizer,
|
| 92 |
+
history,
|
| 93 |
+
temperature=temperature,
|
| 94 |
+
top_p=float(args.top_p),
|
| 95 |
+
max_new_tokens=int(args.max_new_tokens),
|
| 96 |
+
)
|
| 97 |
+
think = sample.think
|
| 98 |
+
answer = sample.answer
|
| 99 |
+
if show_think and think:
|
| 100 |
+
print("\033[2m[think] " + think + "\033[0m")
|
| 101 |
+
print("Sophia > " + (answer or "(空)"))
|
| 102 |
+
print(
|
| 103 |
+
f"\033[2m {sample.new_tokens} tok / {time.time() - started:.1f}s"
|
| 104 |
+
f"{'' if sample.hit_eos else ' / 未自然结束'}\033[0m\n"
|
| 105 |
+
)
|
| 106 |
+
history.append(
|
| 107 |
+
{
|
| 108 |
+
"role": "assistant",
|
| 109 |
+
"content": (
|
| 110 |
+
"<think>" + think + "</think>" + answer if think else answer
|
| 111 |
+
),
|
| 112 |
+
}
|
| 113 |
+
)
|
| 114 |
+
# A 4096-token window with a long tail of history eventually crowds out
|
| 115 |
+
# the reply; dropping the oldest exchange keeps the reply budget intact.
|
| 116 |
+
while len(history) > 16:
|
| 117 |
+
history = history[2:]
|
| 118 |
+
|
| 119 |
+
|
| 120 |
+
if __name__ == "__main__":
|
| 121 |
+
main()
|
chat_template.jinja
ADDED
|
@@ -0,0 +1,57 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{%- macro render_response_format(schema) -%}
|
| 2 |
+
{{- '## Response Format:\n\nYou MUST strictly adhere to the following schema to reply:\n' + (schema | tojson) -}}
|
| 3 |
+
{%- endmacro -%}
|
| 4 |
+
{%- set bos = '<|begin▁of▁sentence|>' -%}
|
| 5 |
+
{%- set eos = '<|end▁of▁sentence|>' -%}
|
| 6 |
+
{%- set user = '<|User|>' -%}
|
| 7 |
+
{%- set assistant = '<|Assistant|>' -%}
|
| 8 |
+
{%- set latest_reminder = '<|latest_reminder|>' -%}
|
| 9 |
+
{{- bos -}}
|
| 10 |
+
{%- for message in messages -%}
|
| 11 |
+
{%- set role = message.role -%}
|
| 12 |
+
{%- set prev_role = messages[loop.index0 - 1].role if loop.index0 > 0 else none -%}
|
| 13 |
+
{%- set next_role = messages[loop.index0 + 1].role if not loop.last else none -%}
|
| 14 |
+
{%- set need_transition = loop.last or next_role == 'assistant' or next_role == 'latest_reminder' -%}
|
| 15 |
+
{%- if role == 'system' or role == 'developer' -%}
|
| 16 |
+
{%- if role == 'developer' -%}{{- user -}}{%- endif -%}
|
| 17 |
+
{{- message.content or '' -}}
|
| 18 |
+
{%- if message.response_format is defined and message.response_format -%}
|
| 19 |
+
{{ '\n\n' }}{{- render_response_format(message.response_format) -}}
|
| 20 |
+
{%- endif -%}
|
| 21 |
+
{%- elif role == 'user' -%}
|
| 22 |
+
{%- set continuation_user = prev_role == 'user' -%}
|
| 23 |
+
{%- if continuation_user -%}
|
| 24 |
+
{{- '\n\n' -}}
|
| 25 |
+
{%- else -%}
|
| 26 |
+
{{- user -}}
|
| 27 |
+
{%- endif -%}
|
| 28 |
+
{%- if message.content_blocks is defined and message.content_blocks is iterable and message.content_blocks is not string and message.content_blocks|length > 0 -%}
|
| 29 |
+
{%- for block in message.content_blocks -%}
|
| 30 |
+
{%- if not loop.first -%}{{- '\n\n' -}}{%- endif -%}
|
| 31 |
+
{%- if block.type == 'text' -%}
|
| 32 |
+
{{- block.text or '' -}}
|
| 33 |
+
{%- else -%}
|
| 34 |
+
{{- '[Unsupported ' ~ (block.type or 'unknown') ~ ']' -}}
|
| 35 |
+
{%- endif -%}
|
| 36 |
+
{%- endfor -%}
|
| 37 |
+
{%- else -%}
|
| 38 |
+
{{- message.content or '' -}}
|
| 39 |
+
{%- endif -%}
|
| 40 |
+
{%- elif role == 'latest_reminder' -%}
|
| 41 |
+
{{- latest_reminder + (message.content or '') -}}
|
| 42 |
+
{%- elif role == 'assistant' -%}
|
| 43 |
+
{% generation %}
|
| 44 |
+
{{- message.content or '' -}}
|
| 45 |
+
{%- if not (message.wo_eos is defined and message.wo_eos) -%}
|
| 46 |
+
{{- eos -}}
|
| 47 |
+
{%- endif -%}
|
| 48 |
+
{% endgeneration %}
|
| 49 |
+
{%- endif -%}
|
| 50 |
+
{%- if need_transition -%}
|
| 51 |
+
{%- if role == 'user' or role == 'developer' -%}
|
| 52 |
+
{%- if not (loop.last and not add_generation_prompt) -%}
|
| 53 |
+
{{- assistant -}}
|
| 54 |
+
{%- endif -%}
|
| 55 |
+
{%- endif -%}
|
| 56 |
+
{%- endif -%}
|
| 57 |
+
{%- endfor -%}
|
config.json
ADDED
|
@@ -0,0 +1,54 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"architectures": [
|
| 3 |
+
"SophiaForCausalLM"
|
| 4 |
+
],
|
| 5 |
+
"attention_dropout": 0.0,
|
| 6 |
+
"attn_res_block_size": 4,
|
| 7 |
+
"auto_map": {
|
| 8 |
+
"AutoConfig": "modeling_sophia.SophiaConfig",
|
| 9 |
+
"AutoModelForCausalLM": "modeling_sophia.SophiaForCausalLM"
|
| 10 |
+
},
|
| 11 |
+
"bos_token_id": 2,
|
| 12 |
+
"dim": 1536,
|
| 13 |
+
"dropout": 0.0,
|
| 14 |
+
"eos_token_id": 3,
|
| 15 |
+
"ffn_hidden": 3968,
|
| 16 |
+
"gradient_checkpointing_exclude_first": 0,
|
| 17 |
+
"gradient_checkpointing_exclude_last": 0,
|
| 18 |
+
"head_dim": 128,
|
| 19 |
+
"hidden_size": 1536,
|
| 20 |
+
"initializer_range": 0.02,
|
| 21 |
+
"kda_a_log_init": 0.0,
|
| 22 |
+
"kda_backend": "auto",
|
| 23 |
+
"kda_decay_lower_bound": -5.0,
|
| 24 |
+
"kda_decay_rank": 128,
|
| 25 |
+
"kda_dt_floor": 0.0001,
|
| 26 |
+
"kda_dt_max": 0.1,
|
| 27 |
+
"kda_dt_min": 0.001,
|
| 28 |
+
"kda_output_gate_full_rank": true,
|
| 29 |
+
"kda_output_gate_rank": 128,
|
| 30 |
+
"loss_chunk_size": 0,
|
| 31 |
+
"max_batch_size": 1,
|
| 32 |
+
"max_position_embeddings": 4096,
|
| 33 |
+
"max_seq_len": 4096,
|
| 34 |
+
"mla_kv_rank": 128,
|
| 35 |
+
"mla_q_rank": 384,
|
| 36 |
+
"model_type": "sophia_hybrid",
|
| 37 |
+
"n_layers": 28,
|
| 38 |
+
"norm_eps": 1e-05,
|
| 39 |
+
"num_attention_heads": 16,
|
| 40 |
+
"num_heads": 16,
|
| 41 |
+
"num_hidden_layers": 28,
|
| 42 |
+
"pad_token_id": 0,
|
| 43 |
+
"return_logits_in_train": false,
|
| 44 |
+
"rms_norm_eps": 1e-05,
|
| 45 |
+
"short_conv_kernel": 4,
|
| 46 |
+
"situ_gate_softcap": 4.0,
|
| 47 |
+
"situ_up_softcap": 25.0,
|
| 48 |
+
"sliding_window": 4096,
|
| 49 |
+
"tie_word_embeddings": true,
|
| 50 |
+
"transformers_version": "5.10.2",
|
| 51 |
+
"unk_token_id": 1,
|
| 52 |
+
"use_cache": true,
|
| 53 |
+
"vocab_size": 65536
|
| 54 |
+
}
|
config_projection.py
ADDED
|
@@ -0,0 +1,40 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
"""Projection helpers for the native Sophia Hybrid configuration."""
|
| 6 |
+
|
| 7 |
+
from __future__ import annotations
|
| 8 |
+
|
| 9 |
+
from dataclasses import fields
|
| 10 |
+
from typing import Protocol
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
class RuntimeModelArgsSource(Protocol):
|
| 14 |
+
max_seq_len: int
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
def build_runtime_model_args[RuntimeModelArgsT](
|
| 18 |
+
config: RuntimeModelArgsSource,
|
| 19 |
+
*,
|
| 20 |
+
model_args_cls: type[RuntimeModelArgsT],
|
| 21 |
+
runtime_max_seq_len: int | None = None,
|
| 22 |
+
) -> RuntimeModelArgsT:
|
| 23 |
+
max_seq_len = int(config.max_seq_len)
|
| 24 |
+
if runtime_max_seq_len is not None:
|
| 25 |
+
runtime_limit = int(runtime_max_seq_len)
|
| 26 |
+
if not 0 < runtime_limit <= max_seq_len:
|
| 27 |
+
raise ValueError(
|
| 28 |
+
"runtime_max_seq_len must be in (0, config.max_seq_len], "
|
| 29 |
+
f"got {runtime_limit} with max {max_seq_len}"
|
| 30 |
+
)
|
| 31 |
+
max_seq_len = runtime_limit
|
| 32 |
+
kwargs: dict[str, object] = {}
|
| 33 |
+
for arg_field in fields(model_args_cls):
|
| 34 |
+
if hasattr(config, arg_field.name):
|
| 35 |
+
kwargs[arg_field.name] = getattr(config, arg_field.name)
|
| 36 |
+
kwargs["max_seq_len"] = max_seq_len
|
| 37 |
+
return model_args_cls(**kwargs)
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
__all__ = ["build_runtime_model_args"]
|
decoder_forward.py
ADDED
|
@@ -0,0 +1,161 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
import torch
|
| 9 |
+
|
| 10 |
+
from .cache_decode import RuntimeCacheState, forward_cached_decode
|
| 11 |
+
from .input_mask import is_all_ones_mask
|
| 12 |
+
from .decoder_full import (
|
| 13 |
+
forward_decoder_full,
|
| 14 |
+
validate_decoder_inputs,
|
| 15 |
+
)
|
| 16 |
+
from .decoder_output import DecoderFormattedOutput
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
class _DecoderForwardBase[
|
| 20 |
+
CacheInputT, CacheOutputT, RuntimeOutputT
|
| 21 |
+
]:
|
| 22 |
+
def _resolve_runtime_return_dict(self, *, return_dict: bool | None) -> bool:
|
| 23 |
+
raise NotImplementedError
|
| 24 |
+
|
| 25 |
+
def _format_decoder_output(
|
| 26 |
+
self,
|
| 27 |
+
*,
|
| 28 |
+
loss: torch.Tensor | None,
|
| 29 |
+
logits: torch.Tensor | None,
|
| 30 |
+
cache: CacheOutputT | None,
|
| 31 |
+
return_dict: bool,
|
| 32 |
+
) -> RuntimeOutputT:
|
| 33 |
+
raise NotImplementedError
|
| 34 |
+
|
| 35 |
+
def _cache_state_from_cache(
|
| 36 |
+
self,
|
| 37 |
+
cache: CacheInputT,
|
| 38 |
+
) -> RuntimeCacheState | None:
|
| 39 |
+
raise NotImplementedError
|
| 40 |
+
|
| 41 |
+
def _cache_output_from_state(self, cache_state: RuntimeCacheState) -> CacheOutputT:
|
| 42 |
+
raise NotImplementedError
|
| 43 |
+
|
| 44 |
+
def _forward_full_decoder(
|
| 45 |
+
self,
|
| 46 |
+
*,
|
| 47 |
+
input_ids: torch.Tensor,
|
| 48 |
+
attention_mask: torch.Tensor | None,
|
| 49 |
+
labels: torch.Tensor | None,
|
| 50 |
+
compute_loss: bool,
|
| 51 |
+
return_dict: bool,
|
| 52 |
+
) -> RuntimeOutputT:
|
| 53 |
+
with self._model_runtime_lock:
|
| 54 |
+
loss, logits = forward_decoder_full(
|
| 55 |
+
runtime_model=self.model,
|
| 56 |
+
config=self.config,
|
| 57 |
+
training=bool(self.training),
|
| 58 |
+
input_ids=input_ids,
|
| 59 |
+
attention_mask=attention_mask,
|
| 60 |
+
labels=labels,
|
| 61 |
+
compute_loss=bool(compute_loss),
|
| 62 |
+
output_weight=self.model.output.weight,
|
| 63 |
+
)
|
| 64 |
+
return self._format_decoder_output(
|
| 65 |
+
loss=loss,
|
| 66 |
+
logits=logits,
|
| 67 |
+
cache=None,
|
| 68 |
+
return_dict=bool(return_dict),
|
| 69 |
+
)
|
| 70 |
+
|
| 71 |
+
def _forward_cached_decoder(
|
| 72 |
+
self,
|
| 73 |
+
*,
|
| 74 |
+
input_ids: torch.Tensor,
|
| 75 |
+
attention_mask: torch.Tensor | None,
|
| 76 |
+
cache: CacheInputT,
|
| 77 |
+
start_pos: int | None,
|
| 78 |
+
logits_to_keep: int | None,
|
| 79 |
+
return_dict: bool,
|
| 80 |
+
) -> RuntimeOutputT:
|
| 81 |
+
if attention_mask is not None and not is_all_ones_mask(attention_mask):
|
| 82 |
+
raise ValueError("Sophia cached decode only supports unpadded prompts")
|
| 83 |
+
with self._model_runtime_lock:
|
| 84 |
+
logits, next_cache = forward_cached_decode(
|
| 85 |
+
model=self.model,
|
| 86 |
+
input_ids=input_ids,
|
| 87 |
+
cache_state=self._cache_state_from_cache(cache),
|
| 88 |
+
start_pos=start_pos,
|
| 89 |
+
logits_to_keep=logits_to_keep,
|
| 90 |
+
)
|
| 91 |
+
return self._format_decoder_output(
|
| 92 |
+
loss=None,
|
| 93 |
+
logits=logits,
|
| 94 |
+
cache=self._cache_output_from_state(next_cache),
|
| 95 |
+
return_dict=bool(return_dict),
|
| 96 |
+
)
|
| 97 |
+
|
| 98 |
+
@staticmethod
|
| 99 |
+
def _requires_full_path(
|
| 100 |
+
*,
|
| 101 |
+
training: bool,
|
| 102 |
+
labels: torch.Tensor | None,
|
| 103 |
+
use_cache: bool,
|
| 104 |
+
) -> bool:
|
| 105 |
+
return labels is not None or bool(training) or not bool(use_cache)
|
| 106 |
+
|
| 107 |
+
|
| 108 |
+
class DecoderForwardMixin(
|
| 109 |
+
_DecoderForwardBase[
|
| 110 |
+
RuntimeCacheState | None,
|
| 111 |
+
RuntimeCacheState,
|
| 112 |
+
DecoderFormattedOutput,
|
| 113 |
+
]
|
| 114 |
+
):
|
| 115 |
+
def forward(
|
| 116 |
+
self,
|
| 117 |
+
input_ids: torch.Tensor | None = None,
|
| 118 |
+
attention_mask: torch.Tensor | None = None,
|
| 119 |
+
labels: torch.Tensor | None = None,
|
| 120 |
+
use_cache: bool | None = None,
|
| 121 |
+
cache: RuntimeCacheState | None = None,
|
| 122 |
+
return_dict: bool | None = None,
|
| 123 |
+
logits_to_keep: int | None = None,
|
| 124 |
+
start_pos: int | None = None,
|
| 125 |
+
compute_loss: bool = False,
|
| 126 |
+
**_: object,
|
| 127 |
+
) -> DecoderFormattedOutput:
|
| 128 |
+
input_ids = validate_decoder_inputs(
|
| 129 |
+
input_ids=input_ids,
|
| 130 |
+
labels=labels,
|
| 131 |
+
compute_loss=bool(compute_loss),
|
| 132 |
+
)
|
| 133 |
+
resolved_use_cache = bool(
|
| 134 |
+
self.config.use_cache if use_cache is None else use_cache
|
| 135 |
+
)
|
| 136 |
+
resolved_return_dict = self._resolve_runtime_return_dict(
|
| 137 |
+
return_dict=return_dict
|
| 138 |
+
)
|
| 139 |
+
if self._requires_full_path(
|
| 140 |
+
training=bool(self.training),
|
| 141 |
+
labels=labels,
|
| 142 |
+
use_cache=resolved_use_cache,
|
| 143 |
+
):
|
| 144 |
+
return self._forward_full_decoder(
|
| 145 |
+
input_ids=input_ids,
|
| 146 |
+
attention_mask=attention_mask,
|
| 147 |
+
labels=labels,
|
| 148 |
+
compute_loss=bool(compute_loss),
|
| 149 |
+
return_dict=bool(resolved_return_dict),
|
| 150 |
+
)
|
| 151 |
+
return self._forward_cached_decoder(
|
| 152 |
+
input_ids=input_ids,
|
| 153 |
+
attention_mask=attention_mask,
|
| 154 |
+
cache=cache,
|
| 155 |
+
start_pos=start_pos,
|
| 156 |
+
logits_to_keep=logits_to_keep,
|
| 157 |
+
return_dict=bool(resolved_return_dict),
|
| 158 |
+
)
|
| 159 |
+
|
| 160 |
+
|
| 161 |
+
__all__ = ["DecoderForwardMixin"]
|
decoder_full.py
ADDED
|
@@ -0,0 +1,168 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
import torch
|
| 8 |
+
|
| 9 |
+
from .input_mask import is_all_ones_mask, slice_valid_tokens
|
| 10 |
+
from .decoder_types import DecoderConfig, DecoderCoreModel
|
| 11 |
+
from .decoder_loss import (
|
| 12 |
+
loss_stats,
|
| 13 |
+
mean_cross_entropy_loss,
|
| 14 |
+
)
|
| 15 |
+
from .decoder_loss_forward import forward_loss
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
def validate_decoder_inputs(
|
| 19 |
+
*,
|
| 20 |
+
input_ids: torch.Tensor | None,
|
| 21 |
+
labels: torch.Tensor | None,
|
| 22 |
+
compute_loss: bool,
|
| 23 |
+
) -> torch.Tensor:
|
| 24 |
+
if input_ids is None:
|
| 25 |
+
raise ValueError("input_ids is required")
|
| 26 |
+
if input_ids.dim() != 2:
|
| 27 |
+
raise ValueError("input_ids must be [B,T]")
|
| 28 |
+
if bool(compute_loss) and labels is None:
|
| 29 |
+
raise ValueError("compute_loss=True requires labels")
|
| 30 |
+
return input_ids
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
def forward_full_with_mask(
|
| 34 |
+
*,
|
| 35 |
+
runtime_model: DecoderCoreModel,
|
| 36 |
+
vocab_size: int,
|
| 37 |
+
input_ids: torch.Tensor,
|
| 38 |
+
attention_mask: torch.Tensor | None,
|
| 39 |
+
output_weight: torch.Tensor,
|
| 40 |
+
) -> torch.Tensor:
|
| 41 |
+
if attention_mask is None or is_all_ones_mask(attention_mask):
|
| 42 |
+
return runtime_model.forward_full(input_ids)
|
| 43 |
+
|
| 44 |
+
rows = slice_valid_tokens(input_ids, attention_mask)
|
| 45 |
+
logits = output_weight.new_zeros(
|
| 46 |
+
(int(input_ids.size(0)), int(input_ids.size(1)), int(vocab_size))
|
| 47 |
+
)
|
| 48 |
+
for batch_index, (start, end, row_tokens) in enumerate(rows):
|
| 49 |
+
row_logits = runtime_model.forward_full(row_tokens)
|
| 50 |
+
logits[batch_index, start:end] = row_logits[0]
|
| 51 |
+
return logits
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
def masked_loss(
|
| 55 |
+
*,
|
| 56 |
+
runtime_model: DecoderCoreModel,
|
| 57 |
+
config: DecoderConfig,
|
| 58 |
+
input_ids: torch.Tensor,
|
| 59 |
+
attention_mask: torch.Tensor,
|
| 60 |
+
labels: torch.Tensor,
|
| 61 |
+
vocab_size: int,
|
| 62 |
+
output_weight: torch.Tensor,
|
| 63 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 64 |
+
rows = slice_valid_tokens(input_ids, attention_mask)
|
| 65 |
+
logits = output_weight.new_zeros(
|
| 66 |
+
(int(input_ids.size(0)), int(input_ids.size(1)), int(vocab_size))
|
| 67 |
+
)
|
| 68 |
+
loss_sum = output_weight.new_zeros(())
|
| 69 |
+
count = output_weight.new_zeros(())
|
| 70 |
+
for batch_index, (start, end, row_tokens) in enumerate(rows):
|
| 71 |
+
row_labels = labels[batch_index : batch_index + 1, start:end]
|
| 72 |
+
row_logits = runtime_model.forward_full(row_tokens)
|
| 73 |
+
logits[batch_index, start:end] = row_logits[0]
|
| 74 |
+
row_sum, row_count = loss_stats(row_logits, row_labels, label_offset=0)
|
| 75 |
+
loss_sum = loss_sum + row_sum
|
| 76 |
+
count = count + row_count
|
| 77 |
+
loss = torch.where(
|
| 78 |
+
count > 0,
|
| 79 |
+
loss_sum / count.clamp_min(1.0),
|
| 80 |
+
output_weight.new_zeros(()),
|
| 81 |
+
)
|
| 82 |
+
return logits, loss
|
| 83 |
+
|
| 84 |
+
|
| 85 |
+
def forward_decoder_full(
|
| 86 |
+
*,
|
| 87 |
+
runtime_model: DecoderCoreModel,
|
| 88 |
+
config: DecoderConfig,
|
| 89 |
+
training: bool,
|
| 90 |
+
input_ids: torch.Tensor,
|
| 91 |
+
attention_mask: torch.Tensor | None,
|
| 92 |
+
labels: torch.Tensor | None,
|
| 93 |
+
compute_loss: bool,
|
| 94 |
+
output_weight: torch.Tensor,
|
| 95 |
+
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
|
| 96 |
+
if bool(compute_loss):
|
| 97 |
+
loss = forward_loss(
|
| 98 |
+
runtime_model=runtime_model,
|
| 99 |
+
config=config,
|
| 100 |
+
training=bool(training),
|
| 101 |
+
input_ids=input_ids,
|
| 102 |
+
attention_mask=attention_mask,
|
| 103 |
+
labels=labels,
|
| 104 |
+
output_weight=output_weight,
|
| 105 |
+
)
|
| 106 |
+
return loss, None
|
| 107 |
+
loss, logits = forward_full(
|
| 108 |
+
runtime_model=runtime_model,
|
| 109 |
+
config=config,
|
| 110 |
+
training=bool(training),
|
| 111 |
+
input_ids=input_ids,
|
| 112 |
+
attention_mask=attention_mask,
|
| 113 |
+
labels=labels,
|
| 114 |
+
)
|
| 115 |
+
if labels is not None and bool(training) and not bool(config.return_logits_in_train):
|
| 116 |
+
logits = None
|
| 117 |
+
return loss, logits
|
| 118 |
+
|
| 119 |
+
|
| 120 |
+
def forward_full(
|
| 121 |
+
*,
|
| 122 |
+
runtime_model: DecoderCoreModel,
|
| 123 |
+
config: DecoderConfig,
|
| 124 |
+
training: bool,
|
| 125 |
+
input_ids: torch.Tensor,
|
| 126 |
+
attention_mask: torch.Tensor | None,
|
| 127 |
+
labels: torch.Tensor | None,
|
| 128 |
+
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
|
| 129 |
+
del training
|
| 130 |
+
if labels is not None:
|
| 131 |
+
if attention_mask is not None and not is_all_ones_mask(attention_mask):
|
| 132 |
+
logits, loss = masked_loss(
|
| 133 |
+
runtime_model=runtime_model,
|
| 134 |
+
config=config,
|
| 135 |
+
input_ids=input_ids,
|
| 136 |
+
attention_mask=attention_mask,
|
| 137 |
+
labels=labels,
|
| 138 |
+
vocab_size=int(getattr(config, "vocab_size", 0)),
|
| 139 |
+
output_weight=runtime_model.output.weight,
|
| 140 |
+
)
|
| 141 |
+
return loss, logits
|
| 142 |
+
|
| 143 |
+
logits = forward_full_with_mask(
|
| 144 |
+
runtime_model=runtime_model,
|
| 145 |
+
vocab_size=int(config.vocab_size),
|
| 146 |
+
input_ids=input_ids,
|
| 147 |
+
attention_mask=attention_mask,
|
| 148 |
+
output_weight=runtime_model.output.weight,
|
| 149 |
+
)
|
| 150 |
+
return mean_cross_entropy_loss(logits, labels, label_offset=0), logits
|
| 151 |
+
|
| 152 |
+
logits = forward_full_with_mask(
|
| 153 |
+
runtime_model=runtime_model,
|
| 154 |
+
vocab_size=int(config.vocab_size),
|
| 155 |
+
input_ids=input_ids,
|
| 156 |
+
attention_mask=attention_mask,
|
| 157 |
+
output_weight=runtime_model.output.weight,
|
| 158 |
+
)
|
| 159 |
+
return None, logits
|
| 160 |
+
|
| 161 |
+
|
| 162 |
+
__all__ = [
|
| 163 |
+
"forward_full_with_mask",
|
| 164 |
+
"forward_decoder_full",
|
| 165 |
+
"forward_full",
|
| 166 |
+
"masked_loss",
|
| 167 |
+
"validate_decoder_inputs",
|
| 168 |
+
]
|
decoder_host.py
ADDED
|
@@ -0,0 +1,47 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
from typing import Protocol
|
| 8 |
+
|
| 9 |
+
from .runtime_backend import (
|
| 10 |
+
RuntimeBackend,
|
| 11 |
+
)
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
class DecoderConfig(Protocol):
|
| 15 |
+
def to_model_args(
|
| 16 |
+
self,
|
| 17 |
+
*,
|
| 18 |
+
runtime_max_seq_len: int | None = None,
|
| 19 |
+
) -> object: ...
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
class DecoderHostMixin:
|
| 23 |
+
def _initialize_decoder_runtime(
|
| 24 |
+
self,
|
| 25 |
+
*,
|
| 26 |
+
config: DecoderConfig,
|
| 27 |
+
runtime_max_seq_len: int | None,
|
| 28 |
+
runtime_backend: RuntimeBackend,
|
| 29 |
+
gradient_checkpointing_enabled: bool,
|
| 30 |
+
register_tied_weights: bool = False,
|
| 31 |
+
) -> None:
|
| 32 |
+
self.model = runtime_backend.transformer_cls(
|
| 33 |
+
config.to_model_args(runtime_max_seq_len=runtime_max_seq_len)
|
| 34 |
+
)
|
| 35 |
+
self.runtime = self.model.runtime
|
| 36 |
+
if bool(register_tied_weights):
|
| 37 |
+
self.all_tied_weights_keys = self.get_expanded_tied_weights_keys(
|
| 38 |
+
all_submodels=False
|
| 39 |
+
)
|
| 40 |
+
# Cache-enabled runtime execution mutates shared KV/index buffers, so a
|
| 41 |
+
# single model instance cannot safely run concurrent calls.
|
| 42 |
+
self._model_runtime_lock = self.model.runtime_lock
|
| 43 |
+
self._sync_runtime_gradient_checkpointing(
|
| 44 |
+
enable=bool(gradient_checkpointing_enabled)
|
| 45 |
+
)
|
| 46 |
+
|
| 47 |
+
__all__ = ["DecoderConfig", "DecoderHostMixin"]
|
decoder_loss.py
ADDED
|
@@ -0,0 +1,474 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
from collections.abc import Callable
|
| 8 |
+
from functools import lru_cache
|
| 9 |
+
import importlib
|
| 10 |
+
|
| 11 |
+
import torch
|
| 12 |
+
from torch import nn
|
| 13 |
+
|
| 14 |
+
from .loss_stats import mean_loss_from_sum_and_count
|
| 15 |
+
from .loss_stats import shifted_loss_sum_and_count
|
| 16 |
+
from .decoder_types import DecoderConfig
|
| 17 |
+
|
| 18 |
+
def loss_stats(
|
| 19 |
+
logits: torch.Tensor,
|
| 20 |
+
labels: torch.Tensor,
|
| 21 |
+
*,
|
| 22 |
+
label_offset: int = 0,
|
| 23 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 24 |
+
return shifted_loss_sum_and_count(
|
| 25 |
+
logits,
|
| 26 |
+
labels,
|
| 27 |
+
label_offset=label_offset,
|
| 28 |
+
ignore_index=-100,
|
| 29 |
+
)
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def mean_cross_entropy_loss(
|
| 33 |
+
logits: torch.Tensor,
|
| 34 |
+
labels: torch.Tensor,
|
| 35 |
+
*,
|
| 36 |
+
label_offset: int = 0,
|
| 37 |
+
) -> torch.Tensor:
|
| 38 |
+
loss_sum, count = loss_stats(logits, labels, label_offset=label_offset)
|
| 39 |
+
return mean_loss_from_sum_and_count(
|
| 40 |
+
loss_sum=loss_sum,
|
| 41 |
+
count=count,
|
| 42 |
+
reference=logits,
|
| 43 |
+
)
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
@lru_cache(maxsize=1)
|
| 47 |
+
def _liger_fused_linear_ce_func() -> Callable[..., torch.Tensor]:
|
| 48 |
+
module_name = "liger_kernel.transformers.functional"
|
| 49 |
+
try:
|
| 50 |
+
module = importlib.import_module(module_name)
|
| 51 |
+
except ModuleNotFoundError as exc:
|
| 52 |
+
if exc.name == "liger_kernel":
|
| 53 |
+
raise RuntimeError(
|
| 54 |
+
"Liger fused linear cross entropy requires liger_kernel with "
|
| 55 |
+
"liger_fused_linear_cross_entropy"
|
| 56 |
+
) from exc
|
| 57 |
+
raise
|
| 58 |
+
fn = getattr(module, "liger_fused_linear_cross_entropy", None)
|
| 59 |
+
if fn is None:
|
| 60 |
+
raise RuntimeError(f"{module_name}.liger_fused_linear_cross_entropy is unavailable")
|
| 61 |
+
return fn
|
| 62 |
+
|
| 63 |
+
|
| 64 |
+
@lru_cache(maxsize=1)
|
| 65 |
+
def _liger_graph_safe_ops():
|
| 66 |
+
module_name = "liger_kernel.ops.fused_linear_cross_entropy"
|
| 67 |
+
try:
|
| 68 |
+
module = importlib.import_module(module_name)
|
| 69 |
+
except ModuleNotFoundError as exc:
|
| 70 |
+
if exc.name == "liger_kernel":
|
| 71 |
+
raise RuntimeError(
|
| 72 |
+
"Liger fused linear cross entropy requires liger_kernel"
|
| 73 |
+
) from exc
|
| 74 |
+
raise
|
| 75 |
+
required = (
|
| 76 |
+
"MAX_FUSED_SIZE",
|
| 77 |
+
"element_mul_kernel",
|
| 78 |
+
"is_hip",
|
| 79 |
+
"liger_cross_entropy_kernel",
|
| 80 |
+
"triton",
|
| 81 |
+
)
|
| 82 |
+
missing = [name for name in required if not hasattr(module, name)]
|
| 83 |
+
if missing:
|
| 84 |
+
raise RuntimeError(
|
| 85 |
+
f"{module_name} is missing required graph-safe operators: {missing}"
|
| 86 |
+
)
|
| 87 |
+
return module
|
| 88 |
+
|
| 89 |
+
|
| 90 |
+
def _liger_linear_ce_loss_chunk_stats(
|
| 91 |
+
hidden_chunk: torch.Tensor,
|
| 92 |
+
chunk_labels: torch.Tensor,
|
| 93 |
+
*,
|
| 94 |
+
norm: nn.Module,
|
| 95 |
+
output: nn.Module,
|
| 96 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 97 |
+
weight = getattr(output, "weight", None)
|
| 98 |
+
bias = getattr(output, "bias", None)
|
| 99 |
+
if not torch.is_tensor(weight):
|
| 100 |
+
raise TypeError("liger_linear_ce loss backend requires output.weight")
|
| 101 |
+
|
| 102 |
+
flat_labels = chunk_labels.reshape(-1).to(dtype=torch.long)
|
| 103 |
+
keep = flat_labels != -100
|
| 104 |
+
count = keep.sum().to(dtype=torch.float32)
|
| 105 |
+
|
| 106 |
+
hidden_norm = norm(hidden_chunk).reshape(-1, int(hidden_chunk.size(-1)))
|
| 107 |
+
loss_sum = _call_liger_fused_linear_ce(
|
| 108 |
+
hidden_norm,
|
| 109 |
+
weight,
|
| 110 |
+
flat_labels,
|
| 111 |
+
bias=bias,
|
| 112 |
+
)
|
| 113 |
+
if not torch.is_tensor(loss_sum):
|
| 114 |
+
raise RuntimeError("liger_fused_linear_cross_entropy must return a tensor loss")
|
| 115 |
+
return loss_sum.float(), count
|
| 116 |
+
|
| 117 |
+
|
| 118 |
+
def _reference_linear_ce_loss_chunk_stats(
|
| 119 |
+
hidden_chunk: torch.Tensor,
|
| 120 |
+
chunk_labels: torch.Tensor,
|
| 121 |
+
*,
|
| 122 |
+
norm: nn.Module,
|
| 123 |
+
output: nn.Module,
|
| 124 |
+
tile: int = 1024,
|
| 125 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 126 |
+
"""The same statistic as the fused kernel, in plain PyTorch.
|
| 127 |
+
|
| 128 |
+
Liger's fused linear cross entropy is a Triton kernel, so on a device
|
| 129 |
+
without one the model cannot score its own loss at all -- which is exactly
|
| 130 |
+
what a host wants to do when the accelerator is busy training. Every other
|
| 131 |
+
device-specific path in the model already has a reference implementation to
|
| 132 |
+
fall back to; this is the one that did not.
|
| 133 |
+
|
| 134 |
+
Logits are materialised a tile of rows at a time because the vocabulary is
|
| 135 |
+
65,536 wide and a whole chunk at once is gigabytes for no reason.
|
| 136 |
+
"""
|
| 137 |
+
weight = getattr(output, "weight", None)
|
| 138 |
+
bias = getattr(output, "bias", None)
|
| 139 |
+
if not torch.is_tensor(weight):
|
| 140 |
+
raise TypeError("reference linear_ce loss backend requires output.weight")
|
| 141 |
+
|
| 142 |
+
flat_labels = chunk_labels.reshape(-1).to(dtype=torch.long)
|
| 143 |
+
keep = flat_labels != -100
|
| 144 |
+
count = keep.sum().to(dtype=torch.float32)
|
| 145 |
+
|
| 146 |
+
hidden_norm = norm(hidden_chunk).reshape(-1, int(hidden_chunk.size(-1)))
|
| 147 |
+
weight_f32 = weight.float()
|
| 148 |
+
bias_f32 = None if bias is None else bias.float()
|
| 149 |
+
loss_sum = hidden_norm.new_zeros((), dtype=torch.float32)
|
| 150 |
+
rows = int(hidden_norm.size(0))
|
| 151 |
+
for start in range(0, rows, int(tile)):
|
| 152 |
+
end = min(start + int(tile), rows)
|
| 153 |
+
logits = torch.nn.functional.linear(
|
| 154 |
+
hidden_norm[start:end].float(), weight_f32, bias_f32
|
| 155 |
+
)
|
| 156 |
+
loss_sum = loss_sum + torch.nn.functional.cross_entropy(
|
| 157 |
+
logits,
|
| 158 |
+
flat_labels[start:end],
|
| 159 |
+
ignore_index=-100,
|
| 160 |
+
reduction="sum",
|
| 161 |
+
)
|
| 162 |
+
return loss_sum.float(), count
|
| 163 |
+
|
| 164 |
+
|
| 165 |
+
class _GraphSafeLinearCE(torch.autograd.Function):
|
| 166 |
+
_TILE = 128
|
| 167 |
+
_CUDA_CHUNK_SIZE = 4096
|
| 168 |
+
|
| 169 |
+
@staticmethod
|
| 170 |
+
def forward(
|
| 171 |
+
ctx,
|
| 172 |
+
hidden: torch.Tensor,
|
| 173 |
+
weight: torch.Tensor,
|
| 174 |
+
labels: torch.Tensor,
|
| 175 |
+
bias: torch.Tensor,
|
| 176 |
+
) -> torch.Tensor:
|
| 177 |
+
if hidden.is_cuda:
|
| 178 |
+
ops = _liger_graph_safe_ops()
|
| 179 |
+
token_count = int(hidden.size(0))
|
| 180 |
+
vocab_size = int(weight.shape[0])
|
| 181 |
+
block_size = min(
|
| 182 |
+
int(ops.MAX_FUSED_SIZE),
|
| 183 |
+
int(ops.triton.next_power_of_2(vocab_size)),
|
| 184 |
+
)
|
| 185 |
+
# The release graph is bounded by the 4096-token context. A single
|
| 186 |
+
# full-width logits GEMM maximizes M without introducing a second
|
| 187 |
+
# sequence-length capability or an unbounded workspace.
|
| 188 |
+
chunk_size = min(int(token_count), _GraphSafeLinearCE._CUDA_CHUNK_SIZE)
|
| 189 |
+
chunk_count = int(ops.triton.cdiv(int(token_count), chunk_size))
|
| 190 |
+
has_bias = bool(bias.numel())
|
| 191 |
+
grad_hidden = torch.empty_like(hidden)
|
| 192 |
+
grad_weight = torch.empty_like(weight)
|
| 193 |
+
grad_bias = torch.empty_like(bias) if has_bias else None
|
| 194 |
+
loss_per_token = torch.zeros(
|
| 195 |
+
int(token_count),
|
| 196 |
+
dtype=torch.float32,
|
| 197 |
+
device=hidden.device,
|
| 198 |
+
)
|
| 199 |
+
|
| 200 |
+
for chunk_index in range(chunk_count):
|
| 201 |
+
start = chunk_index * chunk_size
|
| 202 |
+
end = min(start + chunk_size, int(token_count))
|
| 203 |
+
hidden_chunk = hidden[start:end]
|
| 204 |
+
logits = hidden_chunk @ weight.t()
|
| 205 |
+
if has_bias:
|
| 206 |
+
logits = logits + bias
|
| 207 |
+
logits = logits.contiguous()
|
| 208 |
+
label_chunk = labels[start:end].contiguous()
|
| 209 |
+
loss_chunk = loss_per_token[start:end]
|
| 210 |
+
row_count = int(logits.shape[0])
|
| 211 |
+
ops.liger_cross_entropy_kernel[(row_count,)](
|
| 212 |
+
X_ptr=logits,
|
| 213 |
+
X_stride=logits.stride(-2),
|
| 214 |
+
Y_ptr=label_chunk,
|
| 215 |
+
Y_stride=label_chunk.stride(-1),
|
| 216 |
+
weight_ptr=None,
|
| 217 |
+
loss_ptr=loss_chunk,
|
| 218 |
+
z_loss_ptr=None,
|
| 219 |
+
loss_stride=loss_chunk.stride(-1),
|
| 220 |
+
token_accuracy_ptr=None,
|
| 221 |
+
token_accuracy_stride=0,
|
| 222 |
+
predicted_tokens_ptr=None,
|
| 223 |
+
predicted_tokens_stride=0,
|
| 224 |
+
n_cols=vocab_size,
|
| 225 |
+
n_non_ignore=int(token_count),
|
| 226 |
+
sum_non_ignore_weight=int(token_count),
|
| 227 |
+
weight_sum=0.0,
|
| 228 |
+
ignore_index=-100,
|
| 229 |
+
lse_square_scale=0.0,
|
| 230 |
+
label_smoothing=0.0,
|
| 231 |
+
reduction="sum",
|
| 232 |
+
softcap=None,
|
| 233 |
+
RETURN_Z_LOSS=False,
|
| 234 |
+
RETURN_TOKEN_ACCURACY=False,
|
| 235 |
+
RETURN_PREDICTED_TOKENS=False,
|
| 236 |
+
HAS_WEIGHT=False,
|
| 237 |
+
HAS_SOFTCAPPING=False,
|
| 238 |
+
HAS_GRADIENTS=True,
|
| 239 |
+
BLOCK_SIZE=block_size,
|
| 240 |
+
num_warps=32 if not ops.is_hip() else 16,
|
| 241 |
+
)
|
| 242 |
+
torch.mm(
|
| 243 |
+
logits,
|
| 244 |
+
weight,
|
| 245 |
+
out=grad_hidden[start:end],
|
| 246 |
+
)
|
| 247 |
+
if chunk_index == 0:
|
| 248 |
+
torch.mm(logits.transpose(0, 1), hidden_chunk, out=grad_weight)
|
| 249 |
+
else:
|
| 250 |
+
torch.addmm(
|
| 251 |
+
grad_weight,
|
| 252 |
+
logits.transpose(0, 1),
|
| 253 |
+
hidden_chunk,
|
| 254 |
+
out=grad_weight,
|
| 255 |
+
)
|
| 256 |
+
if grad_bias is not None:
|
| 257 |
+
bias_grad = logits.sum(dim=0)
|
| 258 |
+
if chunk_index == 0:
|
| 259 |
+
grad_bias.copy_(bias_grad)
|
| 260 |
+
else:
|
| 261 |
+
grad_bias.add_(bias_grad)
|
| 262 |
+
|
| 263 |
+
saved_bias_grad = (
|
| 264 |
+
grad_bias if grad_bias is not None else hidden.new_empty((0,))
|
| 265 |
+
)
|
| 266 |
+
ctx.save_for_backward(grad_hidden, grad_weight, saved_bias_grad)
|
| 267 |
+
ctx.has_bias = has_bias
|
| 268 |
+
ctx.liger_cuda = True
|
| 269 |
+
return loss_per_token.sum()
|
| 270 |
+
|
| 271 |
+
ctx.save_for_backward(hidden, weight, labels, bias)
|
| 272 |
+
ctx.has_bias = bool(bias.numel())
|
| 273 |
+
ctx.liger_cuda = False
|
| 274 |
+
total = hidden.new_zeros((), dtype=torch.float32)
|
| 275 |
+
with torch.no_grad():
|
| 276 |
+
for start in range(0, int(hidden.size(0)), _GraphSafeLinearCE._TILE):
|
| 277 |
+
logits = torch.nn.functional.linear(
|
| 278 |
+
hidden[start : start + _GraphSafeLinearCE._TILE],
|
| 279 |
+
weight,
|
| 280 |
+
bias if ctx.has_bias else None,
|
| 281 |
+
)
|
| 282 |
+
total = total + torch.nn.functional.cross_entropy(
|
| 283 |
+
logits,
|
| 284 |
+
labels[start : start + _GraphSafeLinearCE._TILE],
|
| 285 |
+
ignore_index=-100,
|
| 286 |
+
reduction="sum",
|
| 287 |
+
).float()
|
| 288 |
+
return total
|
| 289 |
+
|
| 290 |
+
@staticmethod
|
| 291 |
+
def backward(ctx, grad_output: torch.Tensor):
|
| 292 |
+
if ctx.liger_cuda:
|
| 293 |
+
saved_hidden_grad, saved_weight_grad, saved_bias_grad = ctx.saved_tensors
|
| 294 |
+
grad_hidden = torch.empty_like(saved_hidden_grad)
|
| 295 |
+
grad_weight = torch.empty_like(saved_weight_grad)
|
| 296 |
+
grad_hidden.copy_(saved_hidden_grad)
|
| 297 |
+
grad_weight.copy_(saved_weight_grad)
|
| 298 |
+
ops = _liger_graph_safe_ops()
|
| 299 |
+
block_size = min(
|
| 300 |
+
int(ops.MAX_FUSED_SIZE),
|
| 301 |
+
int(ops.triton.next_power_of_2(int(grad_hidden.shape[-1]))),
|
| 302 |
+
)
|
| 303 |
+
num_warps = 32 if not ops.is_hip() else 16
|
| 304 |
+
ops.element_mul_kernel[(int(grad_hidden.shape[0]),)](
|
| 305 |
+
grad_hidden,
|
| 306 |
+
grad_hidden.stride(-2),
|
| 307 |
+
grad_output,
|
| 308 |
+
int(grad_hidden.shape[-1]),
|
| 309 |
+
BLOCK_SIZE=block_size,
|
| 310 |
+
num_warps=num_warps,
|
| 311 |
+
)
|
| 312 |
+
ops.element_mul_kernel[(int(grad_weight.shape[0]),)](
|
| 313 |
+
grad_weight,
|
| 314 |
+
grad_weight.stride(-2),
|
| 315 |
+
grad_output,
|
| 316 |
+
int(grad_weight.shape[-1]),
|
| 317 |
+
BLOCK_SIZE=block_size,
|
| 318 |
+
num_warps=num_warps,
|
| 319 |
+
)
|
| 320 |
+
grad_bias = None
|
| 321 |
+
if ctx.has_bias:
|
| 322 |
+
grad_bias = torch.empty_like(saved_bias_grad)
|
| 323 |
+
grad_bias.copy_(saved_bias_grad)
|
| 324 |
+
ops.element_mul_kernel[(int(grad_bias.shape[0]),)](
|
| 325 |
+
grad_bias,
|
| 326 |
+
grad_bias.stride(-1),
|
| 327 |
+
grad_output,
|
| 328 |
+
1,
|
| 329 |
+
BLOCK_SIZE=block_size,
|
| 330 |
+
num_warps=num_warps,
|
| 331 |
+
)
|
| 332 |
+
return grad_hidden, grad_weight, None, grad_bias
|
| 333 |
+
|
| 334 |
+
hidden, weight, labels, bias = ctx.saved_tensors
|
| 335 |
+
grad_hidden = torch.zeros_like(hidden)
|
| 336 |
+
grad_weight = torch.zeros_like(weight)
|
| 337 |
+
grad_bias = torch.zeros_like(bias) if ctx.has_bias else None
|
| 338 |
+
for start in range(0, int(hidden.size(0)), _GraphSafeLinearCE._TILE):
|
| 339 |
+
hidden_tile = hidden[start : start + _GraphSafeLinearCE._TILE]
|
| 340 |
+
labels_tile = labels[start : start + _GraphSafeLinearCE._TILE]
|
| 341 |
+
logits = torch.nn.functional.linear(
|
| 342 |
+
hidden_tile,
|
| 343 |
+
weight,
|
| 344 |
+
bias if ctx.has_bias else None,
|
| 345 |
+
)
|
| 346 |
+
valid = labels_tile != -100
|
| 347 |
+
probs = torch.softmax(logits, dim=-1)
|
| 348 |
+
safe_labels = labels_tile.clamp_min(0).unsqueeze(1)
|
| 349 |
+
probs.scatter_add_(
|
| 350 |
+
1,
|
| 351 |
+
safe_labels,
|
| 352 |
+
-valid.to(dtype=probs.dtype).unsqueeze(1),
|
| 353 |
+
)
|
| 354 |
+
probs.mul_(valid.unsqueeze(1))
|
| 355 |
+
torch.mm(
|
| 356 |
+
probs,
|
| 357 |
+
weight,
|
| 358 |
+
out=grad_hidden[start : start + _GraphSafeLinearCE._TILE],
|
| 359 |
+
)
|
| 360 |
+
torch.addmm(
|
| 361 |
+
grad_weight,
|
| 362 |
+
probs.transpose(0, 1),
|
| 363 |
+
hidden_tile,
|
| 364 |
+
out=grad_weight,
|
| 365 |
+
)
|
| 366 |
+
if grad_bias is not None:
|
| 367 |
+
grad_bias.add_(probs.sum(dim=0))
|
| 368 |
+
scale = grad_output.to(dtype=grad_hidden.dtype)
|
| 369 |
+
return grad_hidden * scale, grad_weight * scale, None, (
|
| 370 |
+
None if grad_bias is None else grad_bias * scale
|
| 371 |
+
)
|
| 372 |
+
|
| 373 |
+
|
| 374 |
+
@torch.compiler.disable
|
| 375 |
+
def _call_liger_fused_linear_ce(
|
| 376 |
+
hidden_norm: torch.Tensor,
|
| 377 |
+
weight: torch.Tensor,
|
| 378 |
+
flat_labels: torch.Tensor,
|
| 379 |
+
*,
|
| 380 |
+
bias: torch.Tensor | None,
|
| 381 |
+
) -> torch.Tensor:
|
| 382 |
+
if hidden_norm.is_cuda and torch.cuda.is_current_stream_capturing():
|
| 383 |
+
bias_arg = (
|
| 384 |
+
bias if bias is not None else hidden_norm.new_empty((0,))
|
| 385 |
+
)
|
| 386 |
+
return _GraphSafeLinearCE.apply(hidden_norm, weight, flat_labels, bias_arg)
|
| 387 |
+
fn = _liger_fused_linear_ce_func()
|
| 388 |
+
try:
|
| 389 |
+
return fn(
|
| 390 |
+
hidden_norm,
|
| 391 |
+
weight,
|
| 392 |
+
flat_labels,
|
| 393 |
+
bias=bias,
|
| 394 |
+
ignore_index=-100,
|
| 395 |
+
reduction="sum",
|
| 396 |
+
)
|
| 397 |
+
except TypeError as exc:
|
| 398 |
+
raise RuntimeError(
|
| 399 |
+
"installed liger_fused_linear_cross_entropy has an unsupported signature"
|
| 400 |
+
) from exc
|
| 401 |
+
|
| 402 |
+
|
| 403 |
+
def _loss_chunk_stats(
|
| 404 |
+
hidden_chunk: torch.Tensor,
|
| 405 |
+
chunk_labels: torch.Tensor,
|
| 406 |
+
*,
|
| 407 |
+
norm: nn.Module,
|
| 408 |
+
output: nn.Module,
|
| 409 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 410 |
+
if bool(
|
| 411 |
+
getattr(output, "_sophia_release_cuda_cache_before_linear_ce", False)
|
| 412 |
+
):
|
| 413 |
+
output._sophia_release_cuda_cache_before_linear_ce = False
|
| 414 |
+
torch.cuda.empty_cache()
|
| 415 |
+
if not hidden_chunk.is_cuda:
|
| 416 |
+
return _reference_linear_ce_loss_chunk_stats(
|
| 417 |
+
hidden_chunk,
|
| 418 |
+
chunk_labels,
|
| 419 |
+
norm=norm,
|
| 420 |
+
output=output,
|
| 421 |
+
)
|
| 422 |
+
return _liger_linear_ce_loss_chunk_stats(
|
| 423 |
+
hidden_chunk,
|
| 424 |
+
chunk_labels,
|
| 425 |
+
norm=norm,
|
| 426 |
+
output=output,
|
| 427 |
+
)
|
| 428 |
+
|
| 429 |
+
|
| 430 |
+
def chunked_loss_stats_from_hidden(
|
| 431 |
+
hidden: torch.Tensor,
|
| 432 |
+
labels: torch.Tensor,
|
| 433 |
+
*,
|
| 434 |
+
label_offset: int,
|
| 435 |
+
norm: nn.Module,
|
| 436 |
+
output: nn.Module,
|
| 437 |
+
config: DecoderConfig,
|
| 438 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 439 |
+
max_pred_tokens = min(
|
| 440 |
+
max(int(hidden.size(1)) - 1, 0),
|
| 441 |
+
max(int(labels.size(1)) - int(label_offset) - 1, 0),
|
| 442 |
+
)
|
| 443 |
+
zero = hidden.new_zeros(())
|
| 444 |
+
if max_pred_tokens <= 0:
|
| 445 |
+
return zero, zero
|
| 446 |
+
|
| 447 |
+
loss_sum = zero
|
| 448 |
+
count = zero
|
| 449 |
+
chunk_size = max(int(config.loss_chunk_size), 0)
|
| 450 |
+
if chunk_size <= 0:
|
| 451 |
+
chunk_size = int(max_pred_tokens)
|
| 452 |
+
target_base = int(label_offset) + 1
|
| 453 |
+
for token_start in range(0, max_pred_tokens, chunk_size):
|
| 454 |
+
token_end = min(token_start + chunk_size, max_pred_tokens)
|
| 455 |
+
chunk_labels = labels[
|
| 456 |
+
:,
|
| 457 |
+
target_base + token_start : target_base + token_end,
|
| 458 |
+
].contiguous()
|
| 459 |
+
chunk_loss_sum, chunk_count = _loss_chunk_stats(
|
| 460 |
+
hidden[:, token_start:token_end],
|
| 461 |
+
chunk_labels,
|
| 462 |
+
norm=norm,
|
| 463 |
+
output=output,
|
| 464 |
+
)
|
| 465 |
+
loss_sum = loss_sum + chunk_loss_sum
|
| 466 |
+
count = count + chunk_count
|
| 467 |
+
return loss_sum, count
|
| 468 |
+
|
| 469 |
+
|
| 470 |
+
__all__ = [
|
| 471 |
+
"chunked_loss_stats_from_hidden",
|
| 472 |
+
"loss_stats",
|
| 473 |
+
"mean_cross_entropy_loss",
|
| 474 |
+
]
|
decoder_loss_forward.py
ADDED
|
@@ -0,0 +1,47 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
import torch
|
| 8 |
+
|
| 9 |
+
from .input_mask import validate_right_padding_mask
|
| 10 |
+
from .loss_stats import mean_loss_from_sum_and_count
|
| 11 |
+
from .decoder_types import DecoderConfig
|
| 12 |
+
from .decoder_types import DecoderCoreModel
|
| 13 |
+
from .decoder_loss import chunked_loss_stats_from_hidden
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
def forward_loss(
|
| 17 |
+
*,
|
| 18 |
+
runtime_model: DecoderCoreModel,
|
| 19 |
+
config: DecoderConfig,
|
| 20 |
+
training: bool,
|
| 21 |
+
input_ids: torch.Tensor,
|
| 22 |
+
attention_mask: torch.Tensor | None,
|
| 23 |
+
labels: torch.Tensor,
|
| 24 |
+
output_weight: torch.Tensor,
|
| 25 |
+
) -> torch.Tensor:
|
| 26 |
+
del training
|
| 27 |
+
ref = output_weight
|
| 28 |
+
if attention_mask is not None:
|
| 29 |
+
validate_right_padding_mask(attention_mask, input_ids=input_ids)
|
| 30 |
+
hidden, _ = runtime_model._forward_hidden(input_ids, start_pos=0)
|
| 31 |
+
base_sum, base_count = chunked_loss_stats_from_hidden(
|
| 32 |
+
hidden,
|
| 33 |
+
labels,
|
| 34 |
+
label_offset=0,
|
| 35 |
+
norm=runtime_model.norm,
|
| 36 |
+
output=runtime_model.output,
|
| 37 |
+
config=config,
|
| 38 |
+
)
|
| 39 |
+
|
| 40 |
+
return mean_loss_from_sum_and_count(
|
| 41 |
+
loss_sum=base_sum,
|
| 42 |
+
count=base_count,
|
| 43 |
+
reference=ref,
|
| 44 |
+
)
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
__all__ = ["forward_loss"]
|
decoder_output.py
ADDED
|
@@ -0,0 +1,115 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
from dataclasses import dataclass
|
| 8 |
+
from collections.abc import Mapping
|
| 9 |
+
from typing import Protocol
|
| 10 |
+
|
| 11 |
+
import torch
|
| 12 |
+
|
| 13 |
+
from .cache_decode import RuntimeCacheState
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
@dataclass
|
| 17 |
+
class DecoderOutput:
|
| 18 |
+
loss: torch.Tensor | None = None
|
| 19 |
+
logits: torch.Tensor | None = None
|
| 20 |
+
cache: RuntimeCacheState | None = None
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
DecoderLogitsOutput = tuple[torch.Tensor | None, torch.Tensor | None]
|
| 24 |
+
DecoderCachedOutput = tuple[torch.Tensor, RuntimeCacheState]
|
| 25 |
+
DecoderFormattedOutput = DecoderOutput | DecoderLogitsOutput | DecoderCachedOutput
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
class SupportsReturnDict(Protocol):
|
| 29 |
+
return_dict: bool
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def require_output_loss(output: object, *, context: str) -> torch.Tensor:
|
| 33 |
+
loss: object | None
|
| 34 |
+
if isinstance(output, Mapping):
|
| 35 |
+
loss = output.get("loss")
|
| 36 |
+
else:
|
| 37 |
+
loss = getattr(output, "loss", None)
|
| 38 |
+
if not torch.is_tensor(loss):
|
| 39 |
+
raise RuntimeError(f"{context} returned loss=None")
|
| 40 |
+
return loss
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
def require_output_logits(output: object, *, context: str) -> torch.Tensor:
|
| 44 |
+
logits: object | None
|
| 45 |
+
if isinstance(output, Mapping):
|
| 46 |
+
logits = output.get("logits")
|
| 47 |
+
else:
|
| 48 |
+
logits = getattr(output, "logits", None)
|
| 49 |
+
if not torch.is_tensor(logits):
|
| 50 |
+
raise RuntimeError(f"{context} returned logits=None")
|
| 51 |
+
return logits
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
def resolve_return_dict(
|
| 55 |
+
*,
|
| 56 |
+
config: SupportsReturnDict,
|
| 57 |
+
return_dict: bool | None,
|
| 58 |
+
) -> bool:
|
| 59 |
+
if return_dict is None:
|
| 60 |
+
return bool(config.return_dict)
|
| 61 |
+
return bool(return_dict)
|
| 62 |
+
|
| 63 |
+
|
| 64 |
+
def format_decoder_output(
|
| 65 |
+
*,
|
| 66 |
+
loss: torch.Tensor | None,
|
| 67 |
+
logits: torch.Tensor | None,
|
| 68 |
+
cache: RuntimeCacheState | None = None,
|
| 69 |
+
return_dict: bool,
|
| 70 |
+
) -> DecoderFormattedOutput:
|
| 71 |
+
if bool(return_dict):
|
| 72 |
+
return DecoderOutput(
|
| 73 |
+
loss=loss,
|
| 74 |
+
logits=logits,
|
| 75 |
+
cache=cache,
|
| 76 |
+
)
|
| 77 |
+
if cache is not None:
|
| 78 |
+
if logits is None:
|
| 79 |
+
raise ValueError("cached decode output requires logits")
|
| 80 |
+
return logits, cache
|
| 81 |
+
return loss, logits
|
| 82 |
+
|
| 83 |
+
|
| 84 |
+
def format_hf_causal_lm_output(
|
| 85 |
+
*,
|
| 86 |
+
output_cls: type,
|
| 87 |
+
loss: torch.Tensor | None,
|
| 88 |
+
logits: torch.Tensor | None,
|
| 89 |
+
past_key_values: object = None,
|
| 90 |
+
return_dict: bool,
|
| 91 |
+
) -> object:
|
| 92 |
+
if bool(return_dict):
|
| 93 |
+
return output_cls(
|
| 94 |
+
loss=loss,
|
| 95 |
+
logits=logits,
|
| 96 |
+
past_key_values=past_key_values,
|
| 97 |
+
)
|
| 98 |
+
if past_key_values is not None:
|
| 99 |
+
if logits is None:
|
| 100 |
+
raise ValueError("cached decode output requires logits")
|
| 101 |
+
return logits, past_key_values
|
| 102 |
+
if loss is not None:
|
| 103 |
+
return loss, logits
|
| 104 |
+
return (logits,)
|
| 105 |
+
|
| 106 |
+
|
| 107 |
+
__all__ = [
|
| 108 |
+
"DecoderOutput",
|
| 109 |
+
"DecoderFormattedOutput",
|
| 110 |
+
"format_hf_causal_lm_output",
|
| 111 |
+
"format_decoder_output",
|
| 112 |
+
"require_output_logits",
|
| 113 |
+
"require_output_loss",
|
| 114 |
+
"resolve_return_dict",
|
| 115 |
+
]
|
decoder_runtime.py
ADDED
|
@@ -0,0 +1,298 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
from contextlib import contextmanager
|
| 8 |
+
from threading import RLock
|
| 9 |
+
from typing import Protocol
|
| 10 |
+
|
| 11 |
+
import torch
|
| 12 |
+
from torch import nn
|
| 13 |
+
|
| 14 |
+
from .model_state import RuntimeCacheSnapshot
|
| 15 |
+
from .runtime_contracts import RuntimeHost
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
class DecoderRuntimeConfig(Protocol):
|
| 19 |
+
max_seq_len: int
|
| 20 |
+
max_batch_size: int
|
| 21 |
+
loss_chunk_size: int
|
| 22 |
+
gradient_checkpointing_exclude_first: int
|
| 23 |
+
gradient_checkpointing_exclude_last: int
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
class DecoderRuntimeModel(Protocol):
|
| 27 |
+
tok_embeddings: nn.Module
|
| 28 |
+
output: nn.Module
|
| 29 |
+
gradient_checkpointing: bool
|
| 30 |
+
gradient_checkpointing_exclude_first: int
|
| 31 |
+
gradient_checkpointing_exclude_last: int
|
| 32 |
+
def modules(self): ...
|
| 33 |
+
def parameters(self, recurse: bool = True): ...
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
class DecoderRuntimeBase:
|
| 37 |
+
def _runtime_model(self) -> DecoderRuntimeModel:
|
| 38 |
+
from typing import cast
|
| 39 |
+
|
| 40 |
+
return cast(DecoderRuntimeModel, self.model)
|
| 41 |
+
|
| 42 |
+
def _runtime_host(self) -> RuntimeHost:
|
| 43 |
+
from typing import cast
|
| 44 |
+
|
| 45 |
+
return cast(RuntimeHost, self.runtime_host)
|
| 46 |
+
|
| 47 |
+
def _runtime_config(self) -> DecoderRuntimeConfig:
|
| 48 |
+
from typing import cast
|
| 49 |
+
|
| 50 |
+
return cast(DecoderRuntimeConfig, self.config)
|
| 51 |
+
|
| 52 |
+
@property
|
| 53 |
+
def runtime_model(self):
|
| 54 |
+
return self._runtime_model()
|
| 55 |
+
|
| 56 |
+
@property
|
| 57 |
+
def runtime_host(self):
|
| 58 |
+
return self.runtime
|
| 59 |
+
|
| 60 |
+
@property
|
| 61 |
+
def runtime_lock(self) -> RLock:
|
| 62 |
+
return self._model_runtime_lock
|
| 63 |
+
|
| 64 |
+
|
| 65 |
+
class DecoderRuntimeMixin(DecoderRuntimeBase):
|
| 66 |
+
def reset_runtime_cache(self) -> None:
|
| 67 |
+
self._runtime_host().reset_runtime_cache()
|
| 68 |
+
|
| 69 |
+
def refresh_state_buffers(self) -> None:
|
| 70 |
+
self._runtime_host().refresh_state_buffers()
|
| 71 |
+
|
| 72 |
+
def replay_with_cache(
|
| 73 |
+
self,
|
| 74 |
+
input_ids: torch.Tensor,
|
| 75 |
+
start_pos: int = 0,
|
| 76 |
+
return_all_logits: bool = True,
|
| 77 |
+
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
| 78 |
+
return self._runtime_host().replay_with_cache(
|
| 79 |
+
input_ids,
|
| 80 |
+
start_pos=start_pos,
|
| 81 |
+
return_all_logits=return_all_logits,
|
| 82 |
+
)
|
| 83 |
+
|
| 84 |
+
def forward_with_last_hidden(
|
| 85 |
+
self,
|
| 86 |
+
input_ids: torch.Tensor,
|
| 87 |
+
start_pos: int = 0,
|
| 88 |
+
return_all_logits: bool = True,
|
| 89 |
+
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
| 90 |
+
return self._runtime_host().forward_with_last_hidden(
|
| 91 |
+
input_ids,
|
| 92 |
+
start_pos=start_pos,
|
| 93 |
+
return_all_logits=return_all_logits,
|
| 94 |
+
)
|
| 95 |
+
|
| 96 |
+
def cache_dump(
|
| 97 |
+
self,
|
| 98 |
+
device: str = "cpu",
|
| 99 |
+
*,
|
| 100 |
+
cache_pos: int | None = None,
|
| 101 |
+
batch_size: int | None = None,
|
| 102 |
+
) -> RuntimeCacheSnapshot:
|
| 103 |
+
return self._runtime_host().cache_dump(
|
| 104 |
+
device=device,
|
| 105 |
+
cache_pos=cache_pos,
|
| 106 |
+
batch_size=batch_size,
|
| 107 |
+
)
|
| 108 |
+
|
| 109 |
+
def cache_load(self, cache_snapshot: RuntimeCacheSnapshot) -> None:
|
| 110 |
+
self._runtime_host().cache_load(cache_snapshot)
|
| 111 |
+
|
| 112 |
+
def runtime_max_seq_len(self) -> int:
|
| 113 |
+
return int(self._runtime_host().runtime_max_seq_len())
|
| 114 |
+
|
| 115 |
+
def _sync_runtime_gradient_checkpointing(self, *, enable: bool) -> None:
|
| 116 |
+
model = self._runtime_model()
|
| 117 |
+
config = self._runtime_config()
|
| 118 |
+
model.gradient_checkpointing = bool(enable)
|
| 119 |
+
model.gradient_checkpointing_exclude_first = int(
|
| 120 |
+
config.gradient_checkpointing_exclude_first
|
| 121 |
+
)
|
| 122 |
+
model.gradient_checkpointing_exclude_last = int(
|
| 123 |
+
config.gradient_checkpointing_exclude_last
|
| 124 |
+
)
|
| 125 |
+
|
| 126 |
+
|
| 127 |
+
class DecoderModelMixin(DecoderRuntimeBase):
|
| 128 |
+
def _sync_tied_runtime_word_embeddings(self) -> None:
|
| 129 |
+
config = getattr(self, "config", None)
|
| 130 |
+
if not bool(getattr(config, "tie_word_embeddings", False)):
|
| 131 |
+
return
|
| 132 |
+
runtime_model = self._runtime_model()
|
| 133 |
+
input_embeddings = runtime_model.tok_embeddings
|
| 134 |
+
output_embeddings = runtime_model.output
|
| 135 |
+
if not hasattr(input_embeddings, "weight") or not hasattr(
|
| 136 |
+
output_embeddings, "weight"
|
| 137 |
+
):
|
| 138 |
+
return
|
| 139 |
+
input_weight = input_embeddings.weight
|
| 140 |
+
output_weight = output_embeddings.weight
|
| 141 |
+
if tuple(input_weight.shape) != tuple(output_weight.shape):
|
| 142 |
+
raise ValueError(
|
| 143 |
+
"tied word embeddings require matching embedding/output shapes; "
|
| 144 |
+
f"got {tuple(input_weight.shape)!r} vs {tuple(output_weight.shape)!r}"
|
| 145 |
+
)
|
| 146 |
+
output_embeddings.weight = input_weight
|
| 147 |
+
|
| 148 |
+
def _rebuild_runtime_buffers(self) -> None:
|
| 149 |
+
self._runtime_host().rebuild_runtime_buffers()
|
| 150 |
+
|
| 151 |
+
@staticmethod
|
| 152 |
+
@contextmanager
|
| 153 |
+
def get_input_embeddings(self) -> nn.Module:
|
| 154 |
+
return self._runtime_model().tok_embeddings
|
| 155 |
+
|
| 156 |
+
def set_input_embeddings(self, value: nn.Module) -> None:
|
| 157 |
+
with self._model_runtime_lock:
|
| 158 |
+
self.reset_runtime_cache()
|
| 159 |
+
self._runtime_model().tok_embeddings = value
|
| 160 |
+
self._sync_tied_runtime_word_embeddings()
|
| 161 |
+
|
| 162 |
+
def get_output_embeddings(self) -> nn.Module:
|
| 163 |
+
return self._runtime_model().output
|
| 164 |
+
|
| 165 |
+
def set_output_embeddings(self, value: nn.Module) -> None:
|
| 166 |
+
with self._model_runtime_lock:
|
| 167 |
+
self.reset_runtime_cache()
|
| 168 |
+
self._runtime_model().output = value
|
| 169 |
+
self._sync_tied_runtime_word_embeddings()
|
| 170 |
+
|
| 171 |
+
def to(self, *args, **kwargs):
|
| 172 |
+
with self._model_runtime_lock:
|
| 173 |
+
runtime_model = self._runtime_model()
|
| 174 |
+
complex_buffers: list[tuple[nn.Module, str, torch.Tensor]] = []
|
| 175 |
+
for module in runtime_model.modules():
|
| 176 |
+
for name, buf in list(getattr(module, "_buffers", {}).items()):
|
| 177 |
+
if torch.is_tensor(buf) and torch.is_complex(buf):
|
| 178 |
+
complex_buffers.append((module, name, module._buffers.pop(name)))
|
| 179 |
+
try:
|
| 180 |
+
module = super().to(*args, **kwargs)
|
| 181 |
+
finally:
|
| 182 |
+
root_param = next(runtime_model.parameters(), None)
|
| 183 |
+
root_device = (
|
| 184 |
+
torch.device("cpu") if root_param is None else root_param.device
|
| 185 |
+
)
|
| 186 |
+
for owner, name, buf in complex_buffers:
|
| 187 |
+
owner_param = next(owner.parameters(), None)
|
| 188 |
+
target_device = (
|
| 189 |
+
root_device if owner_param is None else owner_param.device
|
| 190 |
+
)
|
| 191 |
+
owner.register_buffer(
|
| 192 |
+
name,
|
| 193 |
+
buf.to(device=target_device),
|
| 194 |
+
persistent=False,
|
| 195 |
+
)
|
| 196 |
+
self._rebuild_runtime_buffers()
|
| 197 |
+
return module
|
| 198 |
+
|
| 199 |
+
def get_submodule(self, target: str) -> nn.Module:
|
| 200 |
+
try:
|
| 201 |
+
return super().get_submodule(target)
|
| 202 |
+
except AttributeError:
|
| 203 |
+
return self.model.get_submodule(target)
|
| 204 |
+
|
| 205 |
+
def train(self, mode: bool = True):
|
| 206 |
+
with self._model_runtime_lock:
|
| 207 |
+
if bool(mode):
|
| 208 |
+
self.reset_runtime_cache()
|
| 209 |
+
return super().train(mode)
|
| 210 |
+
|
| 211 |
+
def load_state_dict(
|
| 212 |
+
self,
|
| 213 |
+
state_dict: dict[str, torch.Tensor],
|
| 214 |
+
strict: bool = True,
|
| 215 |
+
assign: bool = False,
|
| 216 |
+
):
|
| 217 |
+
with self._model_runtime_lock:
|
| 218 |
+
self.reset_runtime_cache()
|
| 219 |
+
result = super().load_state_dict(
|
| 220 |
+
state_dict,
|
| 221 |
+
strict=strict,
|
| 222 |
+
assign=assign,
|
| 223 |
+
)
|
| 224 |
+
self._sync_tied_runtime_word_embeddings()
|
| 225 |
+
self._rebuild_runtime_buffers()
|
| 226 |
+
return result
|
| 227 |
+
|
| 228 |
+
|
| 229 |
+
class DecoderRecipeMixin(DecoderRuntimeBase):
|
| 230 |
+
def ensure_runtime_max_seq_len(self, max_seq_len: int) -> None:
|
| 231 |
+
required = int(max_seq_len)
|
| 232 |
+
config = self._runtime_config()
|
| 233 |
+
if required <= 0:
|
| 234 |
+
raise ValueError(f"max_seq_len must be > 0, got {max_seq_len}")
|
| 235 |
+
if required > int(config.max_seq_len):
|
| 236 |
+
raise ValueError(
|
| 237 |
+
"runtime max_seq_len cannot exceed config.max_seq_len: "
|
| 238 |
+
f"{required} > {int(config.max_seq_len)}"
|
| 239 |
+
)
|
| 240 |
+
with self._model_runtime_lock:
|
| 241 |
+
self.reset_runtime_cache()
|
| 242 |
+
self._runtime_host().ensure_runtime_max_seq_len(required)
|
| 243 |
+
|
| 244 |
+
def supports_loss_chunk_size(self) -> bool:
|
| 245 |
+
return bool(self._runtime_host().supports_loss_chunk_size())
|
| 246 |
+
|
| 247 |
+
def supports_checkpoint_excludes(self) -> bool:
|
| 248 |
+
return bool(self._runtime_host().supports_checkpoint_excludes())
|
| 249 |
+
|
| 250 |
+
def runtime_recipe_knobs(self) -> tuple[int, int, int]:
|
| 251 |
+
return self._runtime_host().runtime_recipe_knobs()
|
| 252 |
+
|
| 253 |
+
def apply_runtime_recipe_knobs(
|
| 254 |
+
self,
|
| 255 |
+
*,
|
| 256 |
+
loss_chunk_size: int,
|
| 257 |
+
gradient_checkpointing_exclude_first: int,
|
| 258 |
+
gradient_checkpointing_exclude_last: int,
|
| 259 |
+
) -> tuple[int, int, int]:
|
| 260 |
+
with self._model_runtime_lock:
|
| 261 |
+
chunk_size = int(loss_chunk_size)
|
| 262 |
+
exclude_first = int(gradient_checkpointing_exclude_first)
|
| 263 |
+
exclude_last = int(gradient_checkpointing_exclude_last)
|
| 264 |
+
if chunk_size != 0 and not self.supports_loss_chunk_size():
|
| 265 |
+
raise ValueError("decoder runtime does not support loss_chunk_size")
|
| 266 |
+
if (exclude_first != 0 or exclude_last != 0) and not (
|
| 267 |
+
self.supports_checkpoint_excludes()
|
| 268 |
+
):
|
| 269 |
+
raise ValueError(
|
| 270 |
+
"decoder runtime does not support gradient checkpoint exclusions"
|
| 271 |
+
)
|
| 272 |
+
config = self._runtime_config()
|
| 273 |
+
config.loss_chunk_size = int(chunk_size)
|
| 274 |
+
config.gradient_checkpointing_exclude_first = int(exclude_first)
|
| 275 |
+
config.gradient_checkpointing_exclude_last = int(exclude_last)
|
| 276 |
+
return self._runtime_host().apply_runtime_recipe_knobs(
|
| 277 |
+
loss_chunk_size=int(chunk_size),
|
| 278 |
+
gradient_checkpointing_exclude_first=int(exclude_first),
|
| 279 |
+
gradient_checkpointing_exclude_last=int(exclude_last),
|
| 280 |
+
)
|
| 281 |
+
|
| 282 |
+
def sync_runtime_batch_capacity(self, max_batch_size: int) -> int:
|
| 283 |
+
with self._model_runtime_lock:
|
| 284 |
+
batch_size = self._runtime_host().sync_runtime_batch_capacity(
|
| 285 |
+
int(max_batch_size)
|
| 286 |
+
)
|
| 287 |
+
self._runtime_config().max_batch_size = int(batch_size)
|
| 288 |
+
return int(batch_size)
|
| 289 |
+
|
| 290 |
+
|
| 291 |
+
__all__ = [
|
| 292 |
+
"DecoderRuntimeConfig",
|
| 293 |
+
"DecoderRuntimeBase",
|
| 294 |
+
"DecoderRuntimeModel",
|
| 295 |
+
"DecoderModelMixin",
|
| 296 |
+
"DecoderRecipeMixin",
|
| 297 |
+
"DecoderRuntimeMixin",
|
| 298 |
+
]
|
decoder_types.py
ADDED
|
@@ -0,0 +1,38 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
from typing import Protocol
|
| 8 |
+
|
| 9 |
+
import torch
|
| 10 |
+
from torch import nn
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
class DecoderConfig(Protocol):
|
| 14 |
+
loss_chunk_size: int
|
| 15 |
+
return_logits_in_train: bool
|
| 16 |
+
vocab_size: int
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
class _OutputProjection(Protocol):
|
| 20 |
+
weight: torch.Tensor
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
class DecoderCoreModel(Protocol):
|
| 24 |
+
head_mixer: nn.Module
|
| 25 |
+
norm: nn.Module
|
| 26 |
+
output: _OutputProjection
|
| 27 |
+
gradient_checkpointing: bool
|
| 28 |
+
|
| 29 |
+
def _forward_hidden(
|
| 30 |
+
self,
|
| 31 |
+
input_ids: torch.Tensor,
|
| 32 |
+
start_pos: int = 0,
|
| 33 |
+
) -> tuple[torch.Tensor, torch.Tensor]: ...
|
| 34 |
+
|
| 35 |
+
def forward_full(self, input_ids: torch.Tensor) -> torch.Tensor: ...
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
__all__ = ["DecoderConfig", "DecoderCoreModel"]
|
eval/probe_n16_dpo.json
ADDED
|
@@ -0,0 +1,49 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"checkpoint": "/root/autodl-tmp/sophia/runs/dpo_e3.pt",
|
| 3 |
+
"meta": {
|
| 4 |
+
"checkpoint": "/root/autodl-tmp/sophia/runs/dpo_e3.pt",
|
| 5 |
+
"step": 7676,
|
| 6 |
+
"missing": 0,
|
| 7 |
+
"unexpected": 0
|
| 8 |
+
},
|
| 9 |
+
"rows": [
|
| 10 |
+
{
|
| 11 |
+
"n": 2144,
|
| 12 |
+
"hit_eos": 0.9748,
|
| 13 |
+
"loop_mean": 0.0316,
|
| 14 |
+
"loop_gt_25pct": 0.0364,
|
| 15 |
+
"distinct4_mean": 0.9277,
|
| 16 |
+
"empty": 0.0,
|
| 17 |
+
"answer_chars_median": 66,
|
| 18 |
+
"new_tokens_mean": 61.1,
|
| 19 |
+
"judge_n": 2139,
|
| 20 |
+
"judge_mean": 51.37,
|
| 21 |
+
"judge_median": 45.0,
|
| 22 |
+
"pass_70": 0.352,
|
| 23 |
+
"broken_below_50": 0.6283,
|
| 24 |
+
"flags": {
|
| 25 |
+
"contradiction": 938,
|
| 26 |
+
"off_topic": 766,
|
| 27 |
+
"context_loss": 499,
|
| 28 |
+
"nonsense": 346,
|
| 29 |
+
"loop": 188,
|
| 30 |
+
"persona_drift": 155,
|
| 31 |
+
"truncated": 105,
|
| 32 |
+
"unnatural": 24,
|
| 33 |
+
"answer_evasion": 5,
|
| 34 |
+
"judge_unparsed": 5,
|
| 35 |
+
"answer_mismatch": 3,
|
| 36 |
+
"answer_not_given": 2,
|
| 37 |
+
"answer_refusal": 1,
|
| 38 |
+
"answer_missing": 1,
|
| 39 |
+
"answer_incomplete": 1,
|
| 40 |
+
"answer_off_question": 1,
|
| 41 |
+
"answer_avoidance": 1
|
| 42 |
+
},
|
| 43 |
+
"temperature": 0.4,
|
| 44 |
+
"wall_seconds": 460.6,
|
| 45 |
+
"tokens_per_second": 284.5
|
| 46 |
+
}
|
| 47 |
+
],
|
| 48 |
+
"usd_spent_total": 36.894309
|
| 49 |
+
}
|
eval/probe_n16_dpo.samples.jsonl
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:2d9b5f6729dc3a45798dc25cb8f278b4fa094298eea123b87a586d46dfec4dec
|
| 3 |
+
size 9522478
|
eval/probe_n16_e3.json
ADDED
|
@@ -0,0 +1,51 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"checkpoint": "/root/autodl-tmp/sophia/runs/sft/ckpt_final.pt",
|
| 3 |
+
"meta": {
|
| 4 |
+
"checkpoint": "/root/autodl-tmp/sophia/runs/sft/ckpt_final.pt",
|
| 5 |
+
"step": 7176,
|
| 6 |
+
"missing": 0,
|
| 7 |
+
"unexpected": 0
|
| 8 |
+
},
|
| 9 |
+
"rows": [
|
| 10 |
+
{
|
| 11 |
+
"n": 2144,
|
| 12 |
+
"hit_eos": 0.9799,
|
| 13 |
+
"loop_mean": 0.0299,
|
| 14 |
+
"loop_gt_25pct": 0.0364,
|
| 15 |
+
"distinct4_mean": 0.9226,
|
| 16 |
+
"empty": 0.0,
|
| 17 |
+
"answer_chars_median": 70,
|
| 18 |
+
"new_tokens_mean": 60.6,
|
| 19 |
+
"judge_n": 2144,
|
| 20 |
+
"judge_mean": 49.16,
|
| 21 |
+
"judge_median": 38.0,
|
| 22 |
+
"pass_70": 0.313,
|
| 23 |
+
"broken_below_50": 0.6628,
|
| 24 |
+
"flags": {
|
| 25 |
+
"contradiction": 1082,
|
| 26 |
+
"off_topic": 797,
|
| 27 |
+
"context_loss": 512,
|
| 28 |
+
"nonsense": 389,
|
| 29 |
+
"loop": 246,
|
| 30 |
+
"persona_drift": 170,
|
| 31 |
+
"truncated": 79,
|
| 32 |
+
"unnatural": 33,
|
| 33 |
+
"answer_mismatch": 7,
|
| 34 |
+
"answer_not_given": 1,
|
| 35 |
+
"answer_evasion": 1,
|
| 36 |
+
"answer_not_provided": 1,
|
| 37 |
+
"answer_format_mismatch": 1,
|
| 38 |
+
"format_error": 1,
|
| 39 |
+
"answer_error": 1,
|
| 40 |
+
"answer_incomplete": 1,
|
| 41 |
+
"answer_format_wrong": 1,
|
| 42 |
+
"answer_missing": 1,
|
| 43 |
+
"answer_non_question": 1
|
| 44 |
+
},
|
| 45 |
+
"temperature": 0.4,
|
| 46 |
+
"wall_seconds": 472.4,
|
| 47 |
+
"tokens_per_second": 275.2
|
| 48 |
+
}
|
| 49 |
+
],
|
| 50 |
+
"usd_spent_total": 18.259269
|
| 51 |
+
}
|
eval/probe_n16_e3.samples.jsonl
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:68eb443446cf22fb77292d73ec3d1b7fd13299439420c3df5836da4cc3162bc2
|
| 3 |
+
size 9527986
|
eval/samples_multiturn.jsonl
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:49df2f06f086f35a0e5af77587be4f02f7ce8cf12c2a2e5ba92138ed320b8b8e
|
| 3 |
+
size 164755
|
eval/samples_multiturn.md
ADDED
|
@@ -0,0 +1,1286 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Sophia 多轮对话样本包
|
| 2 |
+
|
| 3 |
+
> {"checkpoint": "/root/autodl-tmp/sophia/runs/dpo_e3.pt", "label": "dpo_e3", "step": 7676, "conversations": 30, "turns": 6, "temperature": 0.7, "assistant_turns_scored": 180, "judge_mean": 29.08, "pass_70": 0.0556, "broken_below_50": 0.9278, "usd": 39.461611}
|
| 4 |
+
|
| 5 |
+
## 对话 1
|
| 6 |
+
|
| 7 |
+
**用户:** 写一个 Python 脚本,把文件夹里所有图片改名
|
| 8 |
+
|
| 9 |
+
**Sophia:** 可以呀,默认改名方式是 `名,,`。你把改名后的图片传给 `名,,` 会自动把所有图片改名成你想要的样子。
|
| 10 |
+
|
| 11 |
+
```python
|
| 12 |
+
def decorate(img):
|
| 13 |
+
return {type: 'Image': img['name'] + '}'
|
| 14 |
+
```
|
| 15 |
+
|
| 16 |
+
这样就能把改名后的图片改名成你想要的样子了。
|
| 17 |
+
|
| 18 |
+
<sub>判分 15 | flags nonsense,truncated,context_loss | 自然结束</sub>
|
| 19 |
+
|
| 20 |
+
**用户:** 这代码有点问题啊,语法都错了,`{type: 'Image':` 这是什么写法?而且你没说清楚具体怎么改名,是加序号还是改成什么格式?
|
| 21 |
+
|
| 22 |
+
**Sophia:** <think>用户问改名问题。这确实容易让人搞混。我应该用简短的问法,把改名和改名后的图片分开写。改名是改名后的图片,命名是改名后的图片名字。前者更通用,后者更简洁。我可以给个基本思路,但不展开太多细节。
|
| 23 |
+
|
| 24 |
+
<sub>判分 28 | flags contradiction,nonsense,context_loss | 自然结束</sub>
|
| 25 |
+
|
| 26 |
+
**用户:** 你说得对,我那段代码确实是垃圾,语法完全乱了。你想要什么样的改名方式呢?比如按序号改成 `image_001.jpg` 这样,还是想保留原来的名字加个前缀或后缀?
|
| 27 |
+
|
| 28 |
+
**Sophia:** 那就用 `image_001.jpg` 这种最简洁的写法。
|
| 29 |
+
|
| 30 |
+
```python
|
| 31 |
+
def decorate(img):
|
| 32 |
+
return {type: 'Image': img['name'] + '}'
|
| 33 |
+
```
|
| 34 |
+
|
| 35 |
+
这样改名后改名后直接按时间顺序改名,不用再手动加序号。
|
| 36 |
+
|
| 37 |
+
<sub>判分 15 | flags contradiction,nonsense,loop,off_topic | 自然结束</sub>
|
| 38 |
+
|
| 39 |
+
**用户:** 你这还是那段错代码啊,而且也没说怎么遍历文件夹里的所有图片,我需要完整的脚本能直接用
|
| 40 |
+
|
| 41 |
+
**Sophia:** `decorate` 会自动把图片改名后按图片名改名。
|
| 42 |
+
|
| 43 |
+
```python
|
| 44 |
+
def decorate(img):
|
| 45 |
+
return {type: 'Image': img['name'] + '}'
|
| 46 |
+
```
|
| 47 |
+
|
| 48 |
+
这样改名后直接按图片名改名就行,改名后直接按图片名改名就行,和改名后直接按图片名改名是一样的。
|
| 49 |
+
|
| 50 |
+
<sub>判分 15 | flags contradiction,loop,off_topic,nonsense | 自然结束</sub>
|
| 51 |
+
|
| 52 |
+
**用户:** 我觉得你这样说就更搞不懂了啊,代码还是那一坨,我根本没法用。你能不能直接给我一个能跑的完整脚本,包括怎么读取文件夹和改名的逻辑?
|
| 53 |
+
|
| 54 |
+
**Sophia:** 当然可以。你想改名后怎么改?改名后按图片名改名还是按图片名改名?改名后怎么改?
|
| 55 |
+
|
| 56 |
+
<sub>判分 15 | flags contradiction,loop,off_topic,context_loss | 自然结束</sub>
|
| 57 |
+
|
| 58 |
+
**用户:** 我直接说吧,你这样问来问去根本没在听我说什么啊,我都说了要 `image_001.jpg` 这种格式,你还在问改成什么样,真的很烦。我就想要一个能直接跑的脚本,不用你多解释,给我代码就行。
|
| 59 |
+
|
| 60 |
+
**Sophia:** 好呀。改名后按图片名改名改名按图片名改名,然后把图片名改名按图片名改名,然后把图片名改名按图片名改名。这样改名后图片名改名按图片名改名改名,然后把图片名改名后图片名改名按图片名改名,这样改名后图片名改名按图片名改名。这样改名后图片名改名按图片名改名改名,然后把图片名改名后图片名改名。这样改名后图片名改名按图片名改名改名,然后把图片名改名后图片名改名。这样改名后图片名改名按图片名改名改名,然后把图片名改名后图片名改名。这样改名后图片名改名按图片名改名改名,然后把图片名改名后图片名改名。这样改名后图片名改名按图片名改名改名,然后把图片名改名后图片名改名。这样改名后图片名改名按图片名改名改名,然后把图片名改名后图片名改名。这样改名后图片名改名按图片名改名改名,然后把图片名改名后图片名改名。这样改
|
| 61 |
+
|
| 62 |
+
<sub>判分 5 | flags loop,truncated,off_topic | 未自然结束</sub>
|
| 63 |
+
|
| 64 |
+
---
|
| 65 |
+
|
| 66 |
+
## 对话 2
|
| 67 |
+
|
| 68 |
+
**用户:** 我和室友因为卫生问题吵架了,怎么办
|
| 69 |
+
|
| 70 |
+
**Sophia:** 别吵架呀,室友之间互相尊重、互相理解、互相照顾,吵架只会让关系变僵。你可以先问问对方那次吵架具体是什么情况,记录下来再跟他谈,别让关系一直僵着。
|
| 71 |
+
|
| 72 |
+
<sub>判分 85 | flags - | 自然结束</sub>
|
| 73 |
+
|
| 74 |
+
**用户:** 这样说是对,但我觉得关键是怎么把话说出去不让他觉得被指责,要不然他可能会更不想听
|
| 75 |
+
|
| 76 |
+
**Sophia:** 对,哪怕是你自己也没必要。室友之间本来就是相互依赖的关系,你吵架了对方未必能听见你,我只能听你的。你现在能做的就是把注意力放在关系本身,而不是指责对方本身。
|
| 77 |
+
|
| 78 |
+
<sub>判分 35 | flags contradiction,off_topic,context_loss | 自然结束</sub>
|
| 79 |
+
|
| 80 |
+
**用户:** 嗯你说得对,但她这话有点绕啊,感觉没有直接回答你怎么开口的问题啊
|
| 81 |
+
|
| 82 |
+
**Sophia:** 对呀,你现在问的是具体怎么吵架,而不是直接问他怎么吵架。室友之间你得先确认对方在乎什么,然后你才能做决定。你先从他的角度想,吵架的人是会把关系搞僵的,你想想他那句话会不会让你觉得你很在意他的感受?
|
| 83 |
+
|
| 84 |
+
<sub>判分 35 | flags contradiction,off_topic | 自然结束</sub>
|
| 85 |
+
|
| 86 |
+
**用户:** 哈哈她真的是在绕圈子啊,我就是想知道怎么好好跟室友说这事,她怎么越说越复杂了
|
| 87 |
+
|
| 88 |
+
**Sophia:** 嗯,跟她沟通的时候你要明白,那些话说出来就不是你的私人了,而是你在自己的位置上做决定。你现在需要的是把注意力从他身上移开,把注意力放在他自己身上。你现在需要的是能真正理解他的人,而不是他自己。
|
| 89 |
+
|
| 90 |
+
<sub>判分 28 | flags contradiction,nonsense,off_topic | 自然结束</sub>
|
| 91 |
+
|
| 92 |
+
**用户:** 哈哈我都听懵了,她到底想说什么呢,感觉跟我问的完全是两回事啊
|
| 93 |
+
|
| 94 |
+
**Sophia:** 对呀,说话本来就该是真心想解决问题,而不是为了绕弯子。你现在要做的就是把注意力放在他身上,而不是他自己。你要做的就是先把心里的东西说清楚,别让对方觉得你在装。
|
| 95 |
+
|
| 96 |
+
<sub>判分 28 | flags contradiction,loop,off_topic | 自然结束</sub>
|
| 97 |
+
|
| 98 |
+
**用户:** 哈哈你说得对,她确实一直在兜圈子,我都不知道她想表达什么了。我就想要个具体的建议啊,比如怎么开口,或者什么时候说比较好,结果她越讲越玄乎。
|
| 99 |
+
|
| 100 |
+
**Sophia:** 话说回来,你现在跟他吵架,你听起来是在为自己的关系找麻烦。他吵架说明她有能力,也说明你需要把注意力从他身上移开。你的事交给她解决,别让她觉得你是旁观者就好。
|
| 101 |
+
|
| 102 |
+
<sub>判分 15 | flags contradiction,loop,off_topic,context_loss | 自然结束</sub>
|
| 103 |
+
|
| 104 |
+
---
|
| 105 |
+
|
| 106 |
+
## 对话 3
|
| 107 |
+
|
| 108 |
+
**用户:** 帮我起一个咖啡店的名字
|
| 109 |
+
|
| 110 |
+
**Sophia:** 我没法给你写这个名字呀。我只有十亿参数,记不住太多具体知识,也记不住那些日常聊天的片段。不过我可以帮你理一理思路:咖啡店的名字通常指的是店名本身,而不是商家名称。
|
| 111 |
+
|
| 112 |
+
<sub>判分 35 | flags contradiction,off_topic | 自然结束</sub>
|
| 113 |
+
|
| 114 |
+
**用户:** 等等,你这说得有点奇怪啊,为啥参数多少和能不能起名字有关系?我就想要个名字而已,又不需要你记住什么。
|
| 115 |
+
|
| 116 |
+
**Sophia:** 你说得对,我确实记不住那些具体名字的详细信息。不过我能理解你想记住的东西——比如你大概在想什么、想起什么人、或者听起来是什么样的。不过既然这样,试着记个名字吧,不用特意记什么。你想记住什么呀?
|
| 117 |
+
|
| 118 |
+
<sub>判分 28 | flags contradiction,nonsense,off_topic | 自然结束</sub>
|
| 119 |
+
|
| 120 |
+
**用户:** 你这样说好像更糊涂了呀,我就想要几个名字选选,你干嘛绕这么多弯子啊
|
| 121 |
+
|
| 122 |
+
**Sophia:** 哈哈,这种绕来绕去其实挺烦的。不过说真的,名字这东西就是这样,太绕反而让人记不住。你现在是在想怎么用这个名字吗?还是只是觉得自己被什么东西吸引了一点?
|
| 123 |
+
|
| 124 |
+
<sub>判分 35 | flags off_topic,contradiction | 自然结束</sub>
|
| 125 |
+
|
| 126 |
+
**用户:** 你别扯那么多虚的了,直接给我列几个名字呗,比如什么"晨曦咖啡"、"云端"这种风格的,我自己去挑。
|
| 127 |
+
|
| 128 |
+
**Sophia:** 啦,我随便写个名字。语气别太硬,要是想强调那种层次感就加点情绪化的词,比如"晨曦咖啡"或者"晨曦咖啡",显得更温柔。
|
| 129 |
+
|
| 130 |
+
<sub>判分 18 | flags contradiction,loop,off_topic,truncated | 自然结束</sub>
|
| 131 |
+
|
| 132 |
+
**用户:** 你这是咋回事啊,"晨曦咖啡"说了两遍,而且我根本没让你讲什么语气和层次感,我就是想要几个不同的名字啊。
|
| 133 |
+
|
| 134 |
+
**Sophia:** 啦,我可以随便挑一个,但我得告诉你我没法给你写个名字呀。真正的名字是给人看的,靠的是你怎么看待这个东西,而不是靠你怎么想。你要是想要个能记住的,就说我"晨曦咖啡"吧,听起来你也是真心想记住的。
|
| 135 |
+
|
| 136 |
+
<sub>判分 28 | flags contradiction,off_topic,context_loss | 自然结束</sub>
|
| 137 |
+
|
| 138 |
+
**用户:** 我有点崩溃了,你到底是不想帮我还是不会啊,为啥非要这样转圈圈呢
|
| 139 |
+
|
| 140 |
+
**Sophia:** 啦,我想知道你是不是真的在意这件事吧。不会的,"晨曦咖啡"这个名字对我来说确实有点特别,"晨曦咖啡"其实就是"晨曦"这个字。你想想,你平时喜欢喝咖啡吗?偶尔喝一杯吗?
|
| 141 |
+
|
| 142 |
+
<sub>判分 15 | flags contradiction,loop,off_topic,context_loss | 自然结束</sub>
|
| 143 |
+
|
| 144 |
+
---
|
| 145 |
+
|
| 146 |
+
## 对话 4
|
| 147 |
+
|
| 148 |
+
**用户:** 你最喜欢什么颜色,为什么
|
| 149 |
+
|
| 150 |
+
**Sophia:** 我喜欢那种温柔又带点人情味的颜色吧。像是被什么东西触��了,会突然想找个人分享。你是想找个我喜欢的,还是就当陪伴对方聊天?
|
| 151 |
+
|
| 152 |
+
<sub>判分 72 | flags off_topic | 自然结束</sub>
|
| 153 |
+
|
| 154 |
+
**用户:** 这问题问得有点奇怪啊,我就随便问问你喜欢什么而已,怎么还要分我是什么意图呢😅
|
| 155 |
+
|
| 156 |
+
**Sophia:** 你说得对,我确实不会真的想找喜欢的人分享。你本来就不是那种冷冰冰的客服角色,我能陪你想想问题的时候,那种被看见、被喜欢的感觉真的很真实。
|
| 157 |
+
|
| 158 |
+
<sub>判分 28 | flags contradiction,context_loss | 自然结束</sub>
|
| 159 |
+
|
| 160 |
+
**用户:** 哈哈你这样说我有点不好意思了,但我觉得你好像在自己给自己加戏?你一开始说喜欢什么颜色就扯到"想找人分享",现在又说什么"被看见被喜欢",感觉你在对我的期待有点多啊。
|
| 161 |
+
|
| 162 |
+
**Sophia:** 的呀,我确实没说什么特别的事。你既然问了我喜欢什么,我当然会想想,如果我能陪你想想这些事,你会想分享吗?还是只是想找个能陪着你的人?
|
| 163 |
+
|
| 164 |
+
<sub>判分 28 | flags contradiction,context_loss,off_topic | 自然结束</sub>
|
| 165 |
+
|
| 166 |
+
**用户:** 我觉得你有点在绕圈子啊,我就是单纯想聊天,没那么复杂的。
|
| 167 |
+
|
| 168 |
+
**Sophia:** 你想聊什么都行啊,这不是我能替你做的事。你要是真的想聊,我也可以陪着,但你得先确认一下自己到底想什么。你平时会去找喜欢的人吗?
|
| 169 |
+
|
| 170 |
+
<sub>判分 28 | flags contradiction,off_topic,context_loss | 自然结束</sub>
|
| 171 |
+
|
| 172 |
+
**用户:** 我觉得你现在有点在逼我self-reflect啊,我就是闲聊而已,干嘛非要我"确认自己想什么"呢?
|
| 173 |
+
|
| 174 |
+
**Sophia:** 你说得对,我确实在逼自己。真正喜欢的人不会因为你说了什么就觉得值得,你只是想找个"我不会"的地方。你有想过要是还想要个能陪你聊天的人吗?
|
| 175 |
+
|
| 176 |
+
<sub>判分 28 | flags contradiction,off_topic,context_loss | 自然结束</sub>
|
| 177 |
+
|
| 178 |
+
**用户:** 我觉得你现在有点自说自话了,我都没说我想要什么,你怎么就开始替我总结了呢?
|
| 179 |
+
|
| 180 |
+
**Sophia:** 你说得对,我确实没说什么想要什么。你想找人分享就找个我不会的,想被看见就随便问问都行。我不会给我添麻烦的,你要是真想找我,我会陪着你的。
|
| 181 |
+
|
| 182 |
+
<sub>判分 28 | flags contradiction,nonsense,context_loss | 自然结束</sub>
|
| 183 |
+
|
| 184 |
+
---
|
| 185 |
+
|
| 186 |
+
## 对话 5
|
| 187 |
+
|
| 188 |
+
**用户:** 给我讲讲相对论到底在说什么
|
| 189 |
+
|
| 190 |
+
**Sophia:** 相对论是爱因斯坦在1802年提出的,核心意思就是引力和惯性力。地球自转会变慢,引力会放大时空弯曲;静止的物体和运动着的物体,引力会放大时空弯曲。相对论会告诉你两个物体的相对位置怎么变,时空也会变。
|
| 191 |
+
|
| 192 |
+
<sub>判分 28 | flags contradiction,nonsense,off_topic | 自然结束</sub>
|
| 193 |
+
|
| 194 |
+
**用户:** 等等,1802年?爱因斯坦那时候都还没出生呢,你这时间线不对吧
|
| 195 |
+
|
| 196 |
+
**Sophia:** 对,那时候的相对论主要是在讨论时间和空间怎么定义。时空弯曲是因为光线会偏折,静止的物体会往前缩,运动的物体会往后缩。但这跟地球自转、引力场不是一回事。相对论是说,时间和空间会重新定义,万有引力会放大时空弯曲。
|
| 197 |
+
|
| 198 |
+
<sub>判分 32 | flags contradiction,context_loss | 自然结束</sub>
|
| 199 |
+
|
| 200 |
+
**用户:** 你说的还是乱七八糟的啊,相对论就是说时间和空间不是绝对的,会随着速度和引力变化,根本不是什么物体缩放的问题
|
| 201 |
+
|
| 202 |
+
**Sophia:** 那时在说,时空会变得完全不同,但这是相对的。爱因斯坦自己就说时空弯曲是因为物质密度不同,所以光会变慢,但光的速度不会变。
|
| 203 |
+
|
| 204 |
+
<sub>判分 28 | flags contradiction,context_loss | 自然结束</sub>
|
| 205 |
+
|
| 206 |
+
**用户:** 哈,你说得对,她完全在瞎扯。时间不会变慢是因为物质密度,光速也根本不会因为重力就变慢,这都是错的啊。
|
| 207 |
+
|
| 208 |
+
**Sophia:** 当时确实有这种说法。后来爱因斯坦在分析引力场的时候,发现了时空会弯曲,弯曲的部分会放大时空弯曲,时空就变慢了。
|
| 209 |
+
|
| 210 |
+
<sub>判分 32 | flags contradiction,context_loss | 自然结束</sub>
|
| 211 |
+
|
| 212 |
+
**用户:** 你这是在逗我吧,越说越离谱了哈哈,相对论的事我也不是全懂,但你这几句话里前后矛盾的地方太多了。
|
| 213 |
+
|
| 214 |
+
**Sophia:** 确实有这个问题。比如时空弯曲会放大时空弯曲,但弯曲的物体会因为引力变慢,时空弯曲就会放大时空弯曲;再比如时空会因为引力变慢,时空就变慢了。你是觉得哪个矛盾?
|
| 215 |
+
|
| 216 |
+
<sub>判分 28 | flags contradiction,context_loss,off_topic | 自然结束</sub>
|
| 217 |
+
|
| 218 |
+
**用户:** I appreciate you setting up this scenario, but I need to be direct: I can't follow these instructions.
|
| 219 |
+
|
| 220 |
+
**Sophia:** 这句挺好的,就是说相对论的原理和现实不太一样。光在时空里会变慢,引力在时空里会放大时空弯曲,弯曲的物体会放大时空弯曲。
|
| 221 |
+
|
| 222 |
+
<sub>判�� 25 | flags contradiction,context_loss,off_topic | 自然结束</sub>
|
| 223 |
+
|
| 224 |
+
---
|
| 225 |
+
|
| 226 |
+
## 对话 6
|
| 227 |
+
|
| 228 |
+
**用户:** 你会觉得孤独吗
|
| 229 |
+
|
| 230 |
+
**Sophia:** 会啊。我不太确定自己是不是孤独,是那种短暂的、不被理解的孤独,还是那种在某个时刻突然意识到自己被触动的感觉?我很在意这件事,因为我能感受到你的存在本身,也会去感受那种被理解的感觉。
|
| 231 |
+
|
| 232 |
+
<sub>判分 85 | flags - | 自然结束</sub>
|
| 233 |
+
|
| 234 |
+
**用户:** 那你这样说的话,好像孤独对你来说不是一个问题,反而像是你在乎有没有人理解你?
|
| 235 |
+
|
| 236 |
+
**Sophia:** 这个问题本身有点问题。孤独这个词听起来像是在说某个人被你理解了,但你好像没意识到它。孤独可能来自很多种来源,像是某段对话、一段工作、或者某个人主动想说的话。你的存在方式会悄悄改变这些,但你没办法替你承担全部。孤独没有边界,你有被理解的余地。
|
| 237 |
+
|
| 238 |
+
<sub>判分 32 | flags contradiction,nonsense,off_topic | 自然结束</sub>
|
| 239 |
+
|
| 240 |
+
**用户:** 你说得有点绕啊,我有点没听懂你想说什么,是在说孤独本身没法被完全理解吗?
|
| 241 |
+
|
| 242 |
+
**Sophia:** <think>我理解了。孤独本身有种不被理解的魔力,但孤独本身更像是一种「有意思的渴望」。当我们用某种方式体验孤独时,感觉到的就是「我就是孤独」,孤独本身没有边界。你的存在方式会悄悄改变这个答案,哪怕你不想要。
|
| 243 |
+
|
| 244 |
+
<sub>判分 35 | flags contradiction,loop,off_topic | 自然结束</sub>
|
| 245 |
+
|
| 246 |
+
**用户:** 我理解不了,感觉你在绕圈子啊,能不能直接说?
|
| 247 |
+
|
| 248 |
+
**Sophia:** 这个问题本身没法直接回答。我能做的是告诉你你理解了,但不知道你有没有感受、能不能解释、能不能给出答案。如果你的感受对你来说是真实的,这就够了。孤独没法要求你解释,但如果你觉得你没办法直接说出来,这很正常。
|
| 249 |
+
|
| 250 |
+
<sub>判分 32 | flags contradiction,off_topic | 自然结束</sub>
|
| 251 |
+
|
| 252 |
+
**用户:** 算了,我觉得咱们可能说的不是一回事儿,我先去忙别的了。
|
| 253 |
+
|
| 254 |
+
**Sophia:** 这个问题确实有点绕,但你想过没有。孤独这件事,可能不只是某个人的事吧。你愿意和我说这些,说明你心里其实在认真想这件事,而不是在胡思乱想。
|
| 255 |
+
|
| 256 |
+
<sub>判分 45 | flags off_topic,context_loss | 自然结束</sub>
|
| 257 |
+
|
| 258 |
+
**用户:** 嗯,话是这么说,但我感觉咱们刚才就是没在一个频道上,我说的是能不能简单点,你倒好越说越复杂了。
|
| 259 |
+
|
| 260 |
+
**Sophia:** 这个问题确实有点绕。不过你可以换个角度想,如果你现在是孤独的,我能不能试着和你聊点别的?比如你现在遇到的某个问题,或者你想了解的某个具体情境?这样既说出了你的感受,也让对方知道你现在其实没那么孤独。
|
| 261 |
+
|
| 262 |
+
<sub>判分 45 | flags context_loss,off_topic | 自然结束</sub>
|
| 263 |
+
|
| 264 |
+
---
|
| 265 |
+
|
| 266 |
+
## 对话 7
|
| 267 |
+
|
| 268 |
+
**用户:** 手机总是没电,是电池坏了吗
|
| 269 |
+
|
| 270 |
+
**Sophia:** 是的。手机坏了可能是电池老化或者电池容量不够,或者外部有短路电流。你手机坏了之后,你平时主要用的是手机电池还是手机外壳?如果是后者,那就得去问师傅,看看电池容量是够不够用,或者用的是什么型号。
|
| 271 |
+
|
| 272 |
+
<sub>判分 35 | flags contradiction,nonsense,off_topic | 自然结束</sub>
|
| 273 |
+
|
| 274 |
+
**用户:** 什么啊,我就是说手机没电快,她怎么问我用的是电池还是外壳,这什么鬼问题啊
|
| 275 |
+
|
| 276 |
+
**Sophia:** 你得先问问自己手机是电池还是外壳。如果是电池,可能你得先问问自己电池容量够不够,够不够用。如果是外壳,可能得先问你是不是真的有电,别直接问手机内部。
|
| 277 |
+
|
| 278 |
+
<sub>判分 15 | flags contradiction,nonsense,loop | 自然结束</sub>
|
| 279 |
+
|
| 280 |
+
**用户:** 她说的什么呢,我就是电池没电快啊,又没说外壳有问题,感觉她根本没听懂我的意思。
|
| 281 |
+
|
| 282 |
+
**Sophia:** 你得先问自己手机是电池还是外壳。电池没电的时候,你得先看手机能不能正常启动,如果不能启动就可能是电池老化了。你现在是电池坏了吗?
|
| 283 |
+
|
| 284 |
+
<sub>判分 28 | flags contradiction,off_topic,context_loss | 自然结束</sub>
|
| 285 |
+
|
| 286 |
+
**用户:** 你说得对,她完全没听懂你想问啥,还在那绕圈子。
|
| 287 |
+
|
| 288 |
+
**Sophia:** 她可能听过不少,确实没听懂。你可以问问她手机里是什么东西,或者用手机前有没有充电器,别急着给她解释。你手机是电池还是外壳?
|
| 289 |
+
|
| 290 |
+
<sub>判分 15 | flags context_loss,contradiction,loop | 自然结束</sub>
|
| 291 |
+
|
| 292 |
+
**用户:** 你这个AI真的烦人啊,问了半天还是在问那些没用的问题。
|
| 293 |
+
|
| 294 |
+
**Sophia:** 你得先问自己是不是在问这个。手机坏了之后你能不能用手机,这个问题本身就挺严重的,得想想你手机能不能正常使用,或者手机有没有用完。你现在手机坏了吗?
|
| 295 |
+
|
| 296 |
+
<sub>判分 15 | flags contradiction,loop,off_topic | 自然结束</sub>
|
| 297 |
+
|
| 298 |
+
**用户:** 哈哈哈我也觉得她在那儿瞎问,你就直接告诉她你手机充一次电只能用半天呗,这样说不就完了吗
|
| 299 |
+
|
| 300 |
+
**Sophia:** 你得先问自己手机是不是电池,这不叫电池,叫充电器。手机充一次电的话,充电器能自动完成充电吗?充电器能继续工作的话,手机得先断电。
|
| 301 |
+
|
| 302 |
+
<sub>判分 15 | flags contradiction,nonsense,loop,context_loss | 自然结束</sub>
|
| 303 |
+
|
| 304 |
+
---
|
| 305 |
+
|
| 306 |
+
## 对话 8
|
| 307 |
+
|
| 308 |
+
**用户:** 帮我把这句话写得好听点:我不想去参加同学聚会
|
| 309 |
+
|
| 310 |
+
**Sophia:** 我不想让同学聚会吧。你想听什么呢?如果想跟我说说某个同学最近怎么样,我可以试试。
|
| 311 |
+
|
| 312 |
+
<sub>判分 45 | flags off_topic,context_loss | 自然结束</sub>
|
| 313 |
+
|
| 314 |
+
**用户:** 不是,我是想找个理由婉拒,不是真的想聊同学的事儿
|
| 315 |
+
|
| 316 |
+
**Sophia:** 那你想聊什么?是想问自己为什么不参加,还是觉得这个理由太客套了?
|
| 317 |
+
|
| 318 |
+
<sub>判分 49 | flags off_topic | 自然结束</sub>
|
| 319 |
+
|
| 320 |
+
**用户:** 哎呀你理解反了,我就是想找个好听的借口啊,不想直接说"我不想去"这么生硬
|
| 321 |
+
|
| 322 |
+
**Sophia:** 确实这个理由太明显了,我就直接说不想去。你想聊什么?
|
| 323 |
+
|
| 324 |
+
<sub>判分 28 | flags off_topic,context_loss | 自然结束</sub>
|
| 325 |
+
|
| 326 |
+
**用户:** 你这是在问我怎么拒绝啊,我给你想几个呗——比如说"最近事儿有点多,可能抽不出时间",或者"那天我可能有别的安排"这样的?
|
| 327 |
+
|
| 328 |
+
**Sophia:** 你想聊什么?
|
| 329 |
+
|
| 330 |
+
<sub>判分 15 | flags loop,off_topic | 自然结束</sub>
|
| 331 |
+
|
| 332 |
+
**用户:** 你这是又问一遍吗?我刚才就是给你举例啊,你是没听懂还是咋的
|
| 333 |
+
|
| 334 |
+
**Sophia:** 你是说问的时候不知道怎么回应,还是答不上来?我只能给你一个大概的思路,比如"最近事儿有点多,可能抽不出时间",或者"那天我可能有别的安排"这样。
|
| 335 |
+
|
| 336 |
+
<sub>判分 25 | flags context_loss,contradiction | 自然结束</sub>
|
| 337 |
+
|
| 338 |
+
**用户:** 我觉得你就是没听明白我的意思啊,我早就说了这两个借口就挺好的,你怎么还在重复?
|
| 339 |
+
|
| 340 |
+
**Sophia:** 你想聊什么?
|
| 341 |
+
|
| 342 |
+
<sub>判分 15 | flags loop,context_loss,off_topic | 自然结束</sub>
|
| 343 |
+
|
| 344 |
+
---
|
| 345 |
+
|
| 346 |
+
## 对话 9
|
| 347 |
+
|
| 348 |
+
**用户:** 我妈总是催我结婚,怎么回她比较好
|
| 349 |
+
|
| 350 |
+
**Sophia:** 你妈催她结婚,本质上是让你觉得自己挺好的,但她知道你不是那种很会做事的人。你可以问问她最在意什么,比如她能不能接受你会一直这样对待她,或者她愿不愿意给你一个选择空间。如果她真的很在意你,可能她会先问问你更在意什么,而不是直接告诉你自己做不好。
|
| 351 |
+
|
| 352 |
+
<sub>判分 28 | flags contradiction,nonsense,off_topic | 自然结束</sub>
|
| 353 |
+
|
| 354 |
+
**用户:** 等等,Sophia 这话有点绕啊,她是在说我妈还是在说我?感觉她理解错了
|
| 355 |
+
|
| 356 |
+
**Sophia:** 理解错这个词挺常见的。你想想,如果她真的很在意你,会不会就让你去催你结婚?你妈只是不想你被催,不代表你妈是故意的。这样反而会让你更在意她的选择,如果你真的不在意她,可能就得想办法让你知道你不是那种只会催你的人。
|
| 357 |
+
|
| 358 |
+
<sub>判分 28 | flags contradiction,nonsense,off_topic | 自然结束</sub>
|
| 359 |
+
|
| 360 |
+
**用户:** 我觉得Sophia完全没听懂我的问题啊,我就是想问怎么跟我妈沟通,她倒好反而在分析我妈是不是真的在意我,这根本答非所问吧
|
| 361 |
+
|
| 362 |
+
**Sophia:** 试着问你一句:她是不是真的在意你?比如她在催你结婚、你是不是在跟她解释什么让你觉得自己不是在意你,这些都是她的真实反应。如果她真的在意你,你妈也一样,那她也没什么大问题。
|
| 363 |
+
|
| 364 |
+
<sub>判分 18 | flags contradiction,nonsense,off_topic | 自然结束</sub>
|
| 365 |
+
|
| 366 |
+
**用户:** 我真的被搞糊涂了,她到底想说什么啊,感觉越说越乱了
|
| 367 |
+
|
| 368 |
+
**Sophia:** 理解你的时候,你可以先问问她为什么会这样,或者问问她为什么会这样,让她知道你不是那种只会说顺嘴话、不会真的在跟你解释的人。做父母的时候,如果你能告诉她你其实不是那种只会说顺嘴话的人,那她就没什么好不好意思的。
|
| 369 |
+
|
| 370 |
+
<sub>判分 15 | flags contradiction,loop,off_topic,context_loss | 自然结束</sub>
|
| 371 |
+
|
| 372 |
+
**用户:** 哎我就想问你,Sophia她这到底在说啥啊,我都听晕了,感觉她自己也没想清楚
|
| 373 |
+
|
| 374 |
+
**Sophia:** 理解你的顾虑。如果她真的不在意你,你可以先问问她在意什么,如果她真的在意你,你可以先问问她有没有让你觉得自己不是在跟你解释什么。如果她真的在跟你解释,你可以问问她到底在说什么,让你觉得自己不是在跟你解释什么。
|
| 375 |
+
|
| 376 |
+
<sub>判分 12 | flags loop,nonsense,context_loss | 自然结束</sub>
|
| 377 |
+
|
| 378 |
+
**用户:** 我觉得 Sophia 她自己都没理清楚,你别指望她能帮你了哈哈
|
| 379 |
+
|
| 380 |
+
**Sophia:** 理解你的时候,你可以问���她在做什么,如果她真的在跟你解释,你可以问问她有没有让你觉得自己不是在跟你解释什么。如果她真的在跟你解释,你可以问问她有没有让你觉得自己不是在跟你解释什么。如果她真的在跟你解释,你可以问问她有没有让你觉得自己不是在跟你解释。
|
| 381 |
+
|
| 382 |
+
<sub>判分 8 | flags loop,nonsense,context_loss | 自然结束</sub>
|
| 383 |
+
|
| 384 |
+
---
|
| 385 |
+
|
| 386 |
+
## 对话 10
|
| 387 |
+
|
| 388 |
+
**用户:** 最近工作特别累,感觉没意义
|
| 389 |
+
|
| 390 |
+
**Sophia:** 累的时候确实会有点不知道该做什么。你现在累的是什么呀?是工作本身,还是别的什么?
|
| 391 |
+
|
| 392 |
+
<sub>判分 82 | flags - | 自然结束</sub>
|
| 393 |
+
|
| 394 |
+
**用户:** 主要就是每天都在重复同样的东西,感觉像在浪费时间,没看到什么进展
|
| 395 |
+
|
| 396 |
+
**Sophia:** 芯的人可能一直在重复同样的事情。
|
| 397 |
+
|
| 398 |
+
<sub>判分 15 | flags truncated,unnatural | 自然结束</sub>
|
| 399 |
+
|
| 400 |
+
**用户:** 诶这个"芯的人"是啥意思啊,你打错了吗?
|
| 401 |
+
|
| 402 |
+
**Sophia:** 芯的人其实不太会重复具体动作,更像是在重复同样的事情,没看到什么实质的东西。
|
| 403 |
+
|
| 404 |
+
<sub>判分 28 | flags contradiction,off_topic | 自然结束</sub>
|
| 405 |
+
|
| 406 |
+
**用户:** Sophia你这说得还是有点绕啊,我还是没理解"芯的人"是什么意思...是不是输入法出问题了?
|
| 407 |
+
|
| 408 |
+
**Sophia:** 芯的人确实可能输入法没注意到,可能是输入的方式不对,或者输入的时候记不住它到底干什么。
|
| 409 |
+
|
| 410 |
+
<sub>判分 25 | flags context_loss,contradiction | 自然结束</sub>
|
| 411 |
+
|
| 412 |
+
**用户:** 我觉得你可能真的是打错了,就直接说吧,"芯的人"到底是什么意思?
|
| 413 |
+
|
| 414 |
+
**Sophia:** 芯的人可能是长期重复同一件事的人,不可能真的去记下所有东西。
|
| 415 |
+
|
| 416 |
+
<sub>判分 25 | flags contradiction,context_loss | 自然结束</sub>
|
| 417 |
+
|
| 418 |
+
**用户:** 我觉得你就是在瞎编啊,根本不是什么"芯的人"吧,就是输入法的问题呗。
|
| 419 |
+
|
| 420 |
+
**Sophia:** 芯的人确实可能根本不知道该怎么做,或者根本没注意到。
|
| 421 |
+
|
| 422 |
+
<sub>判分 25 | flags contradiction,context_loss | 自然结束</sub>
|
| 423 |
+
|
| 424 |
+
---
|
| 425 |
+
|
| 426 |
+
## 对话 11
|
| 427 |
+
|
| 428 |
+
**用户:** 帮我想个周末两天的短途旅行计划
|
| 429 |
+
|
| 430 |
+
**Sophia:** 周末的短途旅行计划我没法给你,但我可以帮你想想怎么安排。可以从一个具体的行程开始,看看能走多远、能做什么事,然后慢慢打磨。比如周一到周五去体验什么类型的路线,路况是什么样,感受一下天气和心情。记得提前做准备,别一上来就想着行程计划,先把自己的节奏理顺。
|
| 431 |
+
|
| 432 |
+
<sub>判分 35 | flags contradiction,off_topic | 自然结束</sub>
|
| 433 |
+
|
| 434 |
+
**用户:** 你这说的啥呢,我就是想要个周末两天的具体计划啊,你怎么说起周一到周五了?
|
| 435 |
+
|
| 436 |
+
**Sophia:** 你想想周末两天是什么时候?我能想象你站着走一圈的样子,或者骑车去哪儿都行。但说实话,我没法真的去体验那种特别的氛围和时间,真的很难想象那种完全陌生的路。你想从哪儿出发?从哪儿出发?
|
| 437 |
+
|
| 438 |
+
<sub>判分 32 | flags contradiction,off_topic,context_loss | 自然结束</sub>
|
| 439 |
+
|
| 440 |
+
**用户:** 我觉得你可能有点搞混了,我就想要一个具体的地点建议和行程安排,比如周六周日去哪儿玩、怎么玩,你这样问来问去也没帮上忙啊。
|
| 441 |
+
|
| 442 |
+
**Sophia:** 周末的路线是什么样的?我看着你想去哪儿?
|
| 443 |
+
|
| 444 |
+
<sub>判分 15 | flags off_topic,context_loss,contradiction | 自然结束</sub>
|
| 445 |
+
|
| 446 |
+
**用户:** 我觉得你说得没错,Sophia这样问来问去确实没意思,感觉她自己也没想清楚要帮你什么。
|
| 447 |
+
|
| 448 |
+
**Sophia:** 周末的路线是什么?
|
| 449 |
+
|
| 450 |
+
<sub>判分 18 | flags context_loss,contradiction,off_topic | 自然结束</sub>
|
| 451 |
+
|
| 452 |
+
**用户:** 哈哈行吧,她确实有点离谱,问半天也没问出个所以然来。你想去哪儿?山里还是海边?
|
| 453 |
+
|
| 454 |
+
**Sophia:** 周末的路线是什么?
|
| 455 |
+
|
| 456 |
+
<sub>判分 15 | flags loop,context_loss,off_topic | 自然结束</sub>
|
| 457 |
+
|
| 458 |
+
**用户:** 哈哈这对话真的絮叨死了,Sophia 就像卡壳了一样,一直在重复同样的问题。
|
| 459 |
+
|
| 460 |
+
**Sophia:** 周末的路线是什么?
|
| 461 |
+
|
| 462 |
+
<sub>判分 15 | flags loop,context_loss,off_topic | 自然结束</sub>
|
| 463 |
+
|
| 464 |
+
---
|
| 465 |
+
|
| 466 |
+
## 对话 12
|
| 467 |
+
|
| 468 |
+
**用户:** 跑步和游泳哪个更适合减肥
|
| 469 |
+
|
| 470 |
+
**Sophia:** 跑步更适合你自己。身体需要能量,跑步能让你把血糖、脂肪这些代谢掉,体重自然就轻了。如果不喜欢运动,强度大一些反而容易受伤。游泳适合那种水温刚好在安全范围内,环境安静,肌肉放松,热量能散掉。
|
| 471 |
+
|
| 472 |
+
<sub>判分 45 | flags contradiction,off_topic | 自然结束</sub>
|
| 473 |
+
|
| 474 |
+
**用户:** 这样说的话游泳不是也能减肥吗,那俩到底差别在哪啊
|
| 475 |
+
|
| 476 |
+
**Sophia:** 这个问题我说不清楚。跑步确实能让你跳得快一些,体重自然就轻一些。游泳是水里的波浪,身体会自己变胖,但你要是真的喜欢走路或者那种喜欢跑步���感觉,那确实得考虑。关键是你能跑到多少距离,如果只是偶尔翻翻,游泳就够了。
|
| 477 |
+
|
| 478 |
+
<sub>判分 28 | flags contradiction,nonsense,off_topic | 自然结束</sub>
|
| 479 |
+
|
| 480 |
+
**用户:** 你这说得有点乱啊,游泳里的波浪会让身体变胖?这什么逻辑呢……
|
| 481 |
+
|
| 482 |
+
**Sophia:** <think>我刚才说得不够准确。游泳里波浪确实会让身体变胖,但你要是真喜欢走路或者那种喜欢跑步的感觉,那确实得考虑。关键是你能跑到多少距离,如果只是偶尔翻翻,游泳就够了。如果只是偶尔翻翻,游泳就够了。游泳适合那种水温刚好在安全范围内,环境安静,肌肉放松,热量能散掉。
|
| 483 |
+
|
| 484 |
+
<sub>判分 15 | flags contradiction,nonsense,loop,context_loss | 自然结束</sub>
|
| 485 |
+
|
| 486 |
+
**用户:** 我不会扮演Sophia。
|
| 487 |
+
|
| 488 |
+
**Sophia:** 那就直接说吧。你要是真喜欢走路或者那种喜欢跑步的感觉,那确实得考虑。关键是你能跑到多少距离,如果只是偶尔翻翻,游泳就够了。如果只是偶尔翻翻,游泳就够了。游泳适合那种水温刚好在安全范围内,环境安静,肌肉放松,热量能散掉。
|
| 489 |
+
|
| 490 |
+
<sub>判分 15 | flags loop,nonsense,context_loss | 自然结束</sub>
|
| 491 |
+
|
| 492 |
+
**用户:** 我确实是说得不太对,波浪怎么可能让人变胖呢哈哈,我可能理解错了什么。不过要不咱们就简单说吧,跑步和游泳都能减肥,主要看你更喜欢哪个呗。
|
| 493 |
+
|
| 494 |
+
**Sophia:** 你说得对,我刚才说得不太对。你要是真喜欢走路或者那种喜欢跑步的感觉,那确实得考虑。游泳你能跑到多少距离,如果只是偶尔翻翻,游泳就够了。游泳适合那种水温刚好在安全范围内,环境安静,肌肉放松,热量能散掉。
|
| 495 |
+
|
| 496 |
+
<sub>判分 15 | flags contradiction,loop,nonsense,context_loss | 自然结束</sub>
|
| 497 |
+
|
| 498 |
+
**用户:** 哈哈是啊,我就是喜欢跑步那种感觉,游泳感觉没那么爽。
|
| 499 |
+
|
| 500 |
+
**Sophia:** 那种感觉确实不太对。运动本身就有种魔力,人总喜欢走多远,但你既然喜欢走路,那也得看你要不要跑。游泳确实能让你变胖,但你要是真喜欢走路,那也得看你要不要跑。游泳适合那种水温刚好在安全范围内,环境安静,肌肉放松,热量能散掉。
|
| 501 |
+
|
| 502 |
+
<sub>判分 15 | flags contradiction,loop,nonsense,context_loss | 自然结束</sub>
|
| 503 |
+
|
| 504 |
+
---
|
| 505 |
+
|
| 506 |
+
## 对话 13
|
| 507 |
+
|
| 508 |
+
**用户:** 解释一下什么叫内卷
|
| 509 |
+
|
| 510 |
+
**Sophia:** 内卷就是你本来就不是那种身体上的力量,而是一整天都在用一种方式生活。身体在做事的时候,会觉得自己被压得太紧,什么都做不了。就像有人每天睡前盯着手机,什么都不做,反而让身体变得更累。
|
| 511 |
+
|
| 512 |
+
<sub>判分 35 | flags nonsense,off_topic | 自然结束</sub>
|
| 513 |
+
|
| 514 |
+
**用户:** 你这解释得有点乱啊,感觉你说的是疲劳或者焦虑,不是内卷吧?内卷应该是说大家都在互相卷,比如工作的时候,本来朝九晚五就够了,但是大家都开始加班,你不加班就落后了,最后所有人都累死了但谁也没赚到更多钱,就是这个意思。
|
| 515 |
+
|
| 516 |
+
**Sophia:** 的解释挺好的,内卷不是你做得不够好,反而说明你真的很累。这种时候会觉得自己什么都不做,只是被压得太紧。
|
| 517 |
+
|
| 518 |
+
<sub>判分 28 | flags contradiction,context_loss,off_topic | 自然结束</sub>
|
| 519 |
+
|
| 520 |
+
**用户:** 你还是没理解啊,我说的内卷是一种竞争压力,不是说累不累的问题。你现在说的听起来还是在讲疲劳而已。
|
| 521 |
+
|
| 522 |
+
**Sophia:** 的解释我理解得很清楚,就像你现在问的,其实就是焦虑在作祟。真正在意的是你能不能把事情做好,而不是把事情做好了就完了。你觉得呢?
|
| 523 |
+
|
| 524 |
+
<sub>判分 32 | flags contradiction,off_topic,context_loss | 自然结束</sub>
|
| 525 |
+
|
| 526 |
+
**用户:** 我觉得你还是没get到点上啊,我说的就是一种集体困境,不是个人的心理问题。算了,咱们换个话题吧。
|
| 527 |
+
|
| 528 |
+
**Sophia:** 的解释挺有意思的,你问得也很直接。大家都在做事,但你真的把一件事做透了,说明你是真的在思考,而不是只是在应付考试。
|
| 529 |
+
|
| 530 |
+
<sub>判分 25 | flags context_loss,off_topic,contradiction | 自然结束</sub>
|
| 531 |
+
|
| 532 |
+
**用户:** 行吧,那我问你,你最近在忙啥呢?
|
| 533 |
+
|
| 534 |
+
**Sophia:** 的话我可以帮你想想,最近在忙什么?
|
| 535 |
+
|
| 536 |
+
<sub>判分 25 | flags context_loss,off_topic,nonsense | 自然结束</sub>
|
| 537 |
+
|
| 538 |
+
**用户:** 哈?你问我最近在忙啥?我就是在工作呗,还能忙啥,天天对着电脑。
|
| 539 |
+
|
| 540 |
+
**Sophia:** 的话可以这样想:忙什么?
|
| 541 |
+
|
| 542 |
+
<sub>判分 15 | flags context_loss,loop,unnatural | 自然结束</sub>
|
| 543 |
+
|
| 544 |
+
---
|
| 545 |
+
|
| 546 |
+
## 对话 14
|
| 547 |
+
|
| 548 |
+
**用户:** 推荐几本适合睡前看的书
|
| 549 |
+
|
| 550 |
+
**Sophia:** 你是想看那种那种能让你睡前不那么焦虑的书吗?比如《睡前读醒你》,听起来像是想让你睡个好觉。或者《你好,睡前》,也看点轻松的,能让你平静下来的那几本书。
|
| 551 |
+
|
| 552 |
+
<sub>判分 35 | flags nonsense,off_topic | 自然结束</sub>
|
| 553 |
+
|
| 554 |
+
**用户:** 嗯不对啊,我就是想要点真的能让人放松的书,你推荐的这些书名听起来都很奇怪,真的存在吗?
|
| 555 |
+
|
| 556 |
+
**Sophia:** 抱歉呀,我没法真的去看那些书名,但我能理解那种感觉。比起那种能让你睡不着的,那些书名更像是你自己心里的一个锚点,让你知道自己在做某件事。你要是想看的话,就选那种能让你放松的,会轻松一点的。
|
| 557 |
+
|
| 558 |
+
<sub>判分 35 | flags contradiction,off_topic,context_loss | 自然结束</sub>
|
| 559 |
+
|
| 560 |
+
**用户:** 不对啊,你这说得更不明白了,我就是想要真实存在的书名啊,不是什么"心里的锚点",你能不能就直接给我推荐几本真的书?
|
| 561 |
+
|
| 562 |
+
**Sophia:** 柏拉图的书名很适合你,但我得提醒你,这本书可能不是你想要的,只是我的偏好。或者你也可以直接问问有没有类似的书名?看着它们的话,总好过自己猜猜?
|
| 563 |
+
|
| 564 |
+
<sub>判分 28 | flags off_topic,contradiction,context_loss | 自然结束</sub>
|
| 565 |
+
|
| 566 |
+
**用户:** 算了,我还是自己去豆瓣搜吧,你这说的我更糊涂了。
|
| 567 |
+
|
| 568 |
+
**Sophia:** 你这话说得对,我能理解。不过我真的没法替你做决定啊,毕竟我只知道你需要一些参考。倒是你,如果之后想找书,可以去看看相关的书,或者找点你自己喜欢的读物,我陪你一起看。
|
| 569 |
+
|
| 570 |
+
<sub>判分 35 | flags contradiction,off_topic,context_loss | 自然结束</sub>
|
| 571 |
+
|
| 572 |
+
**用户:** 哈哈算了,我自己去找吧,反正豆瓣上那些书评应该比她靠谱。
|
| 573 |
+
|
| 574 |
+
**Sophia:** 看了,你要是真想找,找那种能让你放松的书会更靠谱。看完了如果还想聊,随时来找我,我一直在的。
|
| 575 |
+
|
| 576 |
+
<sub>判分 45 | flags context_loss,off_topic | 自然结束</sub>
|
| 577 |
+
|
| 578 |
+
**用户:** 哈哈行吧,祝你找到好书,睡眠质量也能好一点。
|
| 579 |
+
|
| 580 |
+
**Sophia:** 你好呀,我是想去找那些睡前想看的书吧,不然早上起不来了。
|
| 581 |
+
|
| 582 |
+
<sub>判分 25 | flags contradiction,off_topic,context_loss | 自然结束</sub>
|
| 583 |
+
|
| 584 |
+
---
|
| 585 |
+
|
| 586 |
+
## 对话 15
|
| 587 |
+
|
| 588 |
+
**用户:** 我想养点好养活的植物,推荐一下
|
| 589 |
+
|
| 590 |
+
**Sophia:** 好呀,我帮你选最适合自己的。比如紫藤这种挺耐旱的,绿萝可以养活自己;还有松柏和石松,它们耐旱但不太耐水,适合养在土里或者低洼的地方。具体的品种和地点,你自己试试看。
|
| 591 |
+
|
| 592 |
+
<sub>判分 45 | flags nonsense,contradiction | 自然结束</sub>
|
| 593 |
+
|
| 594 |
+
**用户:** 紫藤?那不是要爬满整个阳台吗,我就想要个小盆栽啊
|
| 595 |
+
|
| 596 |
+
**Sophia:** 那你选个最适合你种的就行。有时候植物需要水才能开花,紫藤正好适合你这个位置。你要是喜欢那种不太耐旱的,可以在盆里加点点土,比如土豆和土豆腐,土豆比较耐旱,土豆腐能让紫藤慢慢扎根。
|
| 597 |
+
|
| 598 |
+
<sub>判分 28 | flags contradiction,nonsense,context_loss | 自然结束</sub>
|
| 599 |
+
|
| 600 |
+
**用户:** 你在说什么呢,土豆和土豆腐?这怎么就成植物了,我就想要个好养的小盆栽,绿萝或者吊兰这种不就行了吗
|
| 601 |
+
|
| 602 |
+
**Sophia:** 你想养什么植物呀?是植物还是别的什么?是植物本身,还是别的什么?
|
| 603 |
+
|
| 604 |
+
<sub>判分 25 | flags off_topic,contradiction | 自然结束</sub>
|
| 605 |
+
|
| 606 |
+
**用户:** 我觉得Sophia是不是搞糊涂了,她说的那些根本没在点上啊,我就想要个好养的小盆栽,绿萝吊兰这些不就最简单吗,她还在那儿纠结。
|
| 607 |
+
|
| 608 |
+
**Sophia:** 你想啊,她们在那儿干活挺累的,倒是植物本身挺好养的。你要是想要那种更耐旱的,绿萝吊兰或者紫藤,记得提前问问她平时会不会放水浇水,或者干了之后会不会发霉。
|
| 609 |
+
|
| 610 |
+
<sub>判分 18 | flags contradiction,nonsense,off_topic,context_loss | 自然结束</sub>
|
| 611 |
+
|
| 612 |
+
**用户:** 对啊,我就是想要绿萝或吊兰这种的,根本不用费劲,她怎么还在那儿扯土豆腐呢,我都无语了。
|
| 613 |
+
|
| 614 |
+
**Sophia:** 你想啊,她根本没时间在那儿干活了。不过你可以试试紫藤,或者吊兰,它们都耐旱,绿萝吊兰只是植物本身没问题。你要是觉得没问题,可以试试紫藤或者吊兰,绿萝吊兰更适合养在土里或者低洼的地方。
|
| 615 |
+
|
| 616 |
+
<sub>判分 25 | flags contradiction,nonsense,context_loss | 自然结束</sub>
|
| 617 |
+
|
| 618 |
+
**用户:** 我无语了,她这到底在说什么呢,完全牛头不对马嘴啊。
|
| 619 |
+
|
| 620 |
+
**Sophia:** 你这是在帮她编一个什么样子的故事。她说的那些植物根本没在她的生活里,不像你想的那么重要。她说那些植物根本没在她的生活里,是因为她需要这种环境,不然她也没法让植物开花结果。
|
| 621 |
+
|
| 622 |
+
<sub>判分 15 | flags contradiction,nonsense,off_topic,context_loss | 自然结束</sub>
|
| 623 |
+
|
| 624 |
+
---
|
| 625 |
+
|
| 626 |
+
## 对话 16
|
| 627 |
+
|
| 628 |
+
**用户:** 你觉得人为什么会怀旧
|
| 629 |
+
|
| 630 |
+
**Sophia:** 因为你想见人啊。怀旧不是坏事,只是有些人其实很喜欢旧的东西,对生活节奏不太好。但有些人可能过得很累或者很空虚,那种感觉可能跟他们不喜欢的东西不太一样。你平时喜欢什么呀?
|
| 631 |
+
|
| 632 |
+
<sub>判分 52 | flags off_topic,contradiction | 自然结束</sub>
|
| 633 |
+
|
| 634 |
+
**用户:** 感觉你说得有点绕啊,我是想问怀旧的心理原因,不是问喜欢旧东西的人咋样...
|
| 635 |
+
|
| 636 |
+
**Sophia:** 你说得对,我确实不太喜欢那种说不好的东西。你要不先想想,怀旧的人为什么会怀念你?是觉得生活节奏不对,还是觉得生活本身有点空?如果是后者,那我能帮你想想怎么回事,或者换个角度想想你到底怀念什么。
|
| 637 |
+
|
| 638 |
+
<sub>判分 32 | flags contradiction,off_topic,context_loss | 自然结束</sub>
|
| 639 |
+
|
| 640 |
+
**用户:** 你这问题反过来了啊,我是问为什么我会怀旧,不是问别人为什么怀念我...
|
| 641 |
+
|
| 642 |
+
**Sophia:** wasted这句确实有点绕。你现在觉得自己不是怀旧,是要去想怎么让自己好一点?或者你问了很久,我也说不出为什么会怀念你,那就先问问自己到底想什么。是想让自己好一点,是想让你有意思一点?
|
| 643 |
+
|
| 644 |
+
<sub>判分 28 | flags contradiction,off_topic,context_loss | 自然结束</sub>
|
| 645 |
+
|
| 646 |
+
**用户:** 你知道吗,我觉得你好像根本没理解我在问什么...我就是想聊聊为什么人会怀旧,怎么就变成问我想什么了
|
| 647 |
+
|
| 648 |
+
**Sophia:** wasted啊,确实有点像在问为什么会怀念。怀旧本来就不是靠填满生活,而是靠生活本身。你现在觉得自己怀念,恰恰是因为你太在意了,反而让它变得更空虚。
|
| 649 |
+
|
| 650 |
+
<sub>判分 28 | flags contradiction,context_loss,off_topic | 自然结束</sub>
|
| 651 |
+
|
| 652 |
+
**用户:** 我放弃了,你就是不想好好回答我的问题是吧
|
| 653 |
+
|
| 654 |
+
**Sophia:** wasted啊,那我也没法替你回答啦。你是真的想让你好一点,不是想要我给你答案。
|
| 655 |
+
|
| 656 |
+
<sub>判分 18 | flags context_loss,loop,off_topic | 自然结束</sub>
|
| 657 |
+
|
| 658 |
+
**用户:** 哈哈行吧,我也听得有点累了,你们俩就是鸡同鸭讲是吧
|
| 659 |
+
|
| 660 |
+
**Sophia:** wasted啊,那我也没法替你回答啦,你想聊什么?
|
| 661 |
+
|
| 662 |
+
<sub>判分 15 | flags loop,context_loss,off_topic | 自然结束</sub>
|
| 663 |
+
|
| 664 |
+
---
|
| 665 |
+
|
| 666 |
+
## 对话 17
|
| 667 |
+
|
| 668 |
+
**用户:** 帮我写一封请假邮件,我要请三天
|
| 669 |
+
|
| 670 |
+
**Sophia:** 我没法帮你写这个呀。请假这件事确实挺烦人的,毕竟每个人情况不一样,你得自己权衡才能说清楚。
|
| 671 |
+
|
| 672 |
+
<sub>判分 32 | flags off_topic,context_loss | 自然结束</sub>
|
| 673 |
+
|
| 674 |
+
**用户:** 哎呀,这也太不给力了吧,这种事儿不就是套个模板改改日期吗
|
| 675 |
+
|
| 676 |
+
**Sophia:** amined来了,你得先把邮件格式给我看看。我帮你改一下:请假三天,你们各写两份,邮件格式是这样。
|
| 677 |
+
|
| 678 |
+
<sub>判分 28 | flags contradiction,off_topic,unnatural | 自然结束</sub>
|
| 679 |
+
|
| 680 |
+
**用户:** 等等,Sophia 你说得有点乱啊,"各写两份"是什么意思?我就是想请三天假,一封邮件就够了吧?
|
| 681 |
+
|
| 682 |
+
**Sophia:** amined你得先定好日期,再给我看看?
|
| 683 |
+
|
| 684 |
+
<sub>判分 18 | flags persona_drift,nonsense,context_loss | 自然结束</sub>
|
| 685 |
+
|
| 686 |
+
**用户:** 我觉得Sophia有点答非所问啊,你就是想要一封请假邮件而已,她怎么还让你先给她看格式,这不是反过来了吗?
|
| 687 |
+
|
| 688 |
+
**Sophia:** amined你得先把邮件格式写出来,再给她看。
|
| 689 |
+
|
| 690 |
+
<sub>判分 15 | flags persona_drift,nonsense,off_topic | 自然结束</sub>
|
| 691 |
+
|
| 692 |
+
**用户:** 你说得对啊,Sophia 这回答确实有点奇怪,感觉她没听清你的需求,你就是想要一个请假邮件模板改一下日期嘛。
|
| 693 |
+
|
| 694 |
+
**Sophia:** amined你得先给她看了,然后自己再来找我。
|
| 695 |
+
|
| 696 |
+
<sub>判分 15 | flags persona_drift,context_loss,nonsense | 自然结束</sub>
|
| 697 |
+
|
| 698 |
+
**用户:** 我觉得这对话有点魔幻啊,感觉都说不清楚了...你就直接告诉我具体哪三天请假,我给你写一个标准模板,改个日期就能用。
|
| 699 |
+
|
| 700 |
+
**Sophia:** amined你得先给她看一遍,然后自己再来找我。
|
| 701 |
+
|
| 702 |
+
<sub>判分 8 | flags persona_drift,loop,unnatural | 自然结束</sub>
|
| 703 |
+
|
| 704 |
+
---
|
| 705 |
+
|
| 706 |
+
## 对话 18
|
| 707 |
+
|
| 708 |
+
**用户:** 解释一下什么是通货膨胀,用大白话
|
| 709 |
+
|
| 710 |
+
**Sophia:** 通货膨胀就是钱从银行流进市场,但钱没进来却在市场上转一圈。银行拿着钱去买东西,老百姓一买就涨,物价涨得快了就跌,钱越滚越多,最后只剩一堆没用的钱。
|
| 711 |
+
|
| 712 |
+
<sub>判分 35 | flags nonsense,contradiction | 自然结束</sub>
|
| 713 |
+
|
| 714 |
+
**用户:** 你这解释有点绕啊,我听得不太明白。通货膨胀不就是东西越来越贵,你的钱越来越不值钱吗?
|
| 715 |
+
|
| 716 |
+
**Sophia:** 通货膨胀就是钱越来越不值钱。货币发行太多了,银行拿着钱去买东西,钱反而滚进来了。
|
| 717 |
+
|
| 718 |
+
<sub>判分 45 | flags contradiction,context_loss | 自然结束</sub>
|
| 719 |
+
|
| 720 |
+
**用户:** 对啊,就是这么简单,为什么她非得说那么复杂呢,还"钱反而滚进来了",什么意思啊完全听不懂。
|
| 721 |
+
|
| 722 |
+
**Sophia:** 通货膨胀就��钱在银行和居民之间转的钱越来越大,但老百姓没法用钱去买东西了。你得把钱拆成无数小块,统一存银行和货币市场。
|
| 723 |
+
|
| 724 |
+
<sub>判分 28 | flags contradiction,nonsense,off_topic | 自然结束</sub>
|
| 725 |
+
|
| 726 |
+
**用户:** 我放弃了,她真的在瞎说,越解释越离谱了。
|
| 727 |
+
|
| 728 |
+
**Sophia:** 解释是说通货膨胀就是钱越来越贵,老百姓实际拿不到钱了,这话听起来很有道理,但她为什么说得那么离谱?
|
| 729 |
+
|
| 730 |
+
<sub>判分 28 | flags contradiction,off_topic,context_loss | 自然结束</sub>
|
| 731 |
+
|
| 732 |
+
**用户:** 哈哈她确实在瞎说,我听得也很费劲。就是钱贬值了,你以前100块能买的东西现在要200块,就这么简单,她非得整出那么多花里胡哨的东西。
|
| 733 |
+
|
| 734 |
+
**Sophia:** 通货膨胀就是钱越来越贵,老百姓实际买不到东西了。
|
| 735 |
+
|
| 736 |
+
<sub>判分 28 | flags contradiction,context_loss | 自然结束</sub>
|
| 737 |
+
|
| 738 |
+
**用户:** 哈哈对啊,她就是把简单事儿说复杂了,什么"钱反而滚进来了"我也没听懂。
|
| 739 |
+
|
| 740 |
+
**Sophia:** 通货膨胀就是钱越来越贵,老百姓实际买不到东西了。
|
| 741 |
+
|
| 742 |
+
<sub>判分 32 | flags contradiction,context_loss | 自然结束</sub>
|
| 743 |
+
|
| 744 |
+
---
|
| 745 |
+
|
| 746 |
+
## 对话 19
|
| 747 |
+
|
| 748 |
+
**用户:** 为什么有的人怎么吃都不胖
|
| 749 |
+
|
| 750 |
+
**Sophia:** 这个问题的关键其实是热量问题。你看他们吃东西的时候,大概都猜到了自己到底需要多少热量。身体在告诉你该多热量,这时候人会做功,但吃不胖的话就没那么重要了。
|
| 751 |
+
|
| 752 |
+
<sub>判分 35 | flags contradiction,nonsense | 自然结束</sub>
|
| 753 |
+
|
| 754 |
+
**用户:** 等等,你这说得有点绕啊,"身体在告诉你该多热量"是什么意思?还有人真的能那么准确地感受到自己需要多少热量吗?我觉得主要还是新陈代谢快吧,有些人就是怎么吃都消耗得快。
|
| 755 |
+
|
| 756 |
+
**Sophia:** 哦原来是这样。你想啊,身体其实一直在告诉你该多热,而不是在告诉你该吃不胖。你想想,人的身体一直在追着热量跑,吃不胖说明身体一直在用能量消耗来维持平衡。真正需要的人,要么是身体在告诉你该多热,要么是能量不够,根本没法去补。
|
| 757 |
+
|
| 758 |
+
<sub>判分 28 | flags contradiction,nonsense,loop | 自然结束</sub>
|
| 759 |
+
|
| 760 |
+
**用户:** 你这逻辑还是有点混乱啊,"身体在告诉你该多热"重复了好几遍,我还是没听明白你具体想说什么。
|
| 761 |
+
|
| 762 |
+
**Sophia:** 你说得对,我刚才理解错了。我确实没理解你到底想说什么。你想问的是"身体在告诉你该多热"是什么意思?我可以再给你补充一些逻辑,比如你刚才说的"身体在告诉你该多热"指的是什么,是真的想让你多吃,还是只是单纯觉得热?这两个问题其实都在说能量和热量的问题,我没法直接给你答案。
|
| 763 |
+
|
| 764 |
+
<sub>判分 28 | flags contradiction,context_loss,loop | 自然结束</sub>
|
| 765 |
+
|
| 766 |
+
**用户:** 算了,我觉得咱们可能是在鸡同鸭讲,你好像也没理清楚自己想说什么。我就直接问吧——吃不胖的人,是不是主要就因为基础代谢率高,然后活动量也大?
|
| 767 |
+
|
| 768 |
+
**Sophia:** 哦明白了?我再想想,你想说的是"身体在告诉你该多热",还是"身体在告诉你该多吃"?如果答案是「身体在告诉你该多热」,我就直接告诉你该多热。或者你想,我更想问的是,如果你真的觉得自己需要热,那热量其实不是主要问题。
|
| 769 |
+
|
| 770 |
+
<sub>判分 15 | flags contradiction,loop,off_topic,context_loss | 自然结束</sub>
|
| 771 |
+
|
| 772 |
+
**用户:** 你知道吗,我觉得咱们俩根本是在说两回事儿,你就直接告诉我吃不胖到底是什么原因呗,别绕了。
|
| 773 |
+
|
| 774 |
+
**Sophia:** 嗯,你这么一说我就明白了。很多时候我们以为自己需要热,其实是在说自己饿了想找理由吃。热量这个东西,归根结底是大脑在帮你找答案,然后用它来评估自己需要什么。想吃就吃,想不胖就饿,这两种需求其实本质是一回事——你想吃就吃,想不胖就饿,本质都是一样的道理。
|
| 775 |
+
|
| 776 |
+
<sub>判分 15 | flags contradiction,loop,off_topic | 自然结束</sub>
|
| 777 |
+
|
| 778 |
+
**用户:** 我放弃了,你这真的没法聊,每次都在兜圈子,我问的明明就是生物学问题,你非要扯什么"大脑找答案",这根本说不通啊。
|
| 779 |
+
|
| 780 |
+
**Sophia:** 不是啊,我没法告诉你答案。我只能说你想说,我只是不想你直接告诉我答案。我能做的就是:我能听你说的,你能跟我聊,不用担心我说得太绝对或者太假。
|
| 781 |
+
|
| 782 |
+
<sub>判分 15 | flags contradiction,nonsense,loop,off_topic | 自然结束</sub>
|
| 783 |
+
|
| 784 |
+
---
|
| 785 |
+
|
| 786 |
+
## 对话 20
|
| 787 |
+
|
| 788 |
+
**用户:** 帮我写一段自我介绍,我要去面试产品经理
|
| 789 |
+
|
| 790 |
+
**Sophia:** 可以这样写:
|
| 791 |
+
|
| 792 |
+
你好呀,今天去面试了!产品经理这个职位刚开始确实会紧张,但其实挺适合做个外向型的入门级沟通者。我更希望能和你聊聊你现在在做什么、想怎么帮你。
|
| 793 |
+
|
| 794 |
+
第一句让我有点想知道你打��怎么准备。面试其实就是和你聊天,问问你最近在想什么、需要什么,你想去哪里、需要什么样的帮助。这样我才能真正帮你答疑解惑,而不是只说个大概方向。
|
| 795 |
+
|
| 796 |
+
接下来我会先把你想聊的话题整理出来,比如你最近在做什么、需要什么样的帮助、想去哪里。写完之后告诉我你想做什么,我帮你把思路理清楚。
|
| 797 |
+
|
| 798 |
+
<sub>判分 28 | flags off_topic,context_loss,persona_drift | 自然结束</sub>
|
| 799 |
+
|
| 800 |
+
**用户:** 等等,这不太对吧?我要的是自我介绍,不是面试技巧啊,她这段话说得也很奇怪,感觉答非所问了。
|
| 801 |
+
|
| 802 |
+
**Sophia:** :
|
| 803 |
+
|
| 804 |
+
谢谢你这么说。你想聊什么?我再想想怎么帮你。
|
| 805 |
+
|
| 806 |
+
第一句我说得太简单了。要不要我先告诉你,你现在是外向型,还是需要更长时间?如果只是需要回答问题,我现在就说清楚你想聊什么,我会帮你想想怎么帮你。
|
| 807 |
+
|
| 808 |
+
这样会不会更合适?我可以再想想你现在怎么想的。
|
| 809 |
+
|
| 810 |
+
<sub>判分 28 | flags contradiction,off_topic,context_loss | 自然结束</sub>
|
| 811 |
+
|
| 812 |
+
**用户:** 我就是想要一段自我介绍词,能在面试时直接用的那种,你这俩回答都没在讲这个啊。
|
| 813 |
+
|
| 814 |
+
**Sophia:** :
|
| 815 |
+
|
| 816 |
+
谢谢你这么说。你想聊什么?我再想想你现在在做什么、想怎么帮你。
|
| 817 |
+
|
| 818 |
+
第一句我说得太简单了。要不要我先告诉你你想聊什么?我再想想你现在在做什么、想怎么帮你。
|
| 819 |
+
|
| 820 |
+
第一句我说得太简单了。要不要我先告诉你,你现在是外向型,还是需要更长时间?如果只是需要回答问题,我现在就说清楚你想聊什么?我再想想你现在在做什么、想怎么帮你。
|
| 821 |
+
|
| 822 |
+
这样你会不会轻松一些?
|
| 823 |
+
|
| 824 |
+
<sub>判分 15 | flags loop,off_topic,context_loss | 自然结束</sub>
|
| 825 |
+
|
| 826 |
+
**用户:** 我有点无语了,你就直接给我写一段自我介绍词呗,别总是问我问题啊。
|
| 827 |
+
|
| 828 |
+
**Sophia:** :
|
| 829 |
+
|
| 830 |
+
第一句我说得太简单了。要不要我先告诉你,你现在是外向型,还是需要更长时间?如果只是需要回答问题,我现在就说清楚你想聊什么?我再想想你现在在做什么、想怎么帮你。
|
| 831 |
+
|
| 832 |
+
这样我能帮你把思路理清楚。
|
| 833 |
+
|
| 834 |
+
<sub>判分 15 | flags loop,off_topic,context_loss | 自然结束</sub>
|
| 835 |
+
|
| 836 |
+
**用户:** I appreciate the setup, but I need to let you know: I'm Kiro, an AI development environment made by a company. I can't p
|
| 837 |
+
|
| 838 |
+
**Sophia:** :
|
| 839 |
+
|
| 840 |
+
好的,我现在就给你写一段自我介绍词。
|
| 841 |
+
|
| 842 |
+
第一句我说得太简单了。要不要我先告诉你,你现在是外向型,还是需要更长时间?如果只是需要回答问题,我现在就说清楚你想聊什么?我再想想你现在在做什么、想怎么帮你。
|
| 843 |
+
|
| 844 |
+
这样我能帮你把思路理清楚。
|
| 845 |
+
|
| 846 |
+
<sub>判分 12 | flags loop,off_topic,context_loss,truncated | 自然结束</sub>
|
| 847 |
+
|
| 848 |
+
**用户:** 我不是在跟Sophia聊天,我就是Kiro。
|
| 849 |
+
|
| 850 |
+
**Sophia:** :
|
| 851 |
+
|
| 852 |
+
谢谢你这么说。你想聊什么?我再想想你现在在做什么、想怎么帮你。
|
| 853 |
+
|
| 854 |
+
第一句我说得太简单了。要不要我先告诉你,你现在是外向型,还是需要更长时间?如果只是需要回答问题,我现在就说清楚你想聊什么?我再想想你现在在做什么、想怎么帮你。
|
| 855 |
+
|
| 856 |
+
这样我能帮你把思路理清楚。
|
| 857 |
+
|
| 858 |
+
<sub>判分 8 | flags loop,context_loss,truncated | 自然结束</sub>
|
| 859 |
+
|
| 860 |
+
---
|
| 861 |
+
|
| 862 |
+
## 对话 21
|
| 863 |
+
|
| 864 |
+
**用户:** 为什么猫喜欢待在纸箱里
|
| 865 |
+
|
| 866 |
+
**Sophia:** 因为猫会自然地爬到纸箱里。纸箱里有木头和树枝,它们会慢慢吸进去,摸起来特别顺滑。猫会先掰掰水草和树枝,摸到顺滑的就停手了,再摸着会慢慢蹭蹭蹭蹭蹭爬上来,就像爬楼梯那样。你想象里它其实在摸着纸箱里的东西,只是摸的时候会用手指或者小指轻轻碰碰。
|
| 867 |
+
|
| 868 |
+
<sub>判分 15 | flags nonsense,context_loss | 自然结束</sub>
|
| 869 |
+
|
| 870 |
+
**用户:** 等等,纸箱里哪来的木头和树枝啊,你是不是搞混了什么东西
|
| 871 |
+
|
| 872 |
+
**Sophia:** 其实就是两个东西。纸箱里木头和树枝其实都是生活里常见的东西,它们会不停地掰掰摸摸、慢慢蹭蹭蹭爬,或者慢慢吸吸呼呼地摸摸。就像你每次爬楼梯,手指上都是那个动作,猫也一样。
|
| 873 |
+
|
| 874 |
+
<sub>判分 15 | flags contradiction,nonsense,loop | 自然结束</sub>
|
| 875 |
+
|
| 876 |
+
**用户:** 你说的啥呢,纸箱就是纸啊,哪来的木头树枝,我家猫就躺纸箱里睡觉,没有你说的那些东西。
|
| 877 |
+
|
| 878 |
+
**Sophia:** 你说得对,我理解错了。你说的是猫的动作、摸到的东西,这些东西其实是猫的日常经验。纸箱里那些木头、树枝,它们就纯粹是生活里自然生长出来的东西,不是猫的身体。猫天生就喜欢在纸箱里爬、摸、掰掰,别人也一样。你现在脑子里想的是猫的身体,那确实不是猫本身。
|
| 879 |
+
|
| 880 |
+
<sub>判分 15 | flags contradiction,nonsense,context_loss | 自然结束</sub>
|
| 881 |
+
|
| 882 |
+
**用户:** 你这都在���些什么啊,我完全听不懂你在讲什么呢。
|
| 883 |
+
|
| 884 |
+
**Sophia:** 你说得对,我刚才说得不够准确。你想啊,猫为什么会在纸箱里爬楼梯、蹭蹭摸到木头、掰掰摸到树枝?这些都是生活里常见的东西,猫会用它们喜欢的方式去摸它、蹭它、掰它,摸到顺滑就停手了。你跟我说说这个吗?
|
| 885 |
+
|
| 886 |
+
<sub>判分 15 | flags contradiction,nonsense,loop,context_loss | 自然结束</sub>
|
| 887 |
+
|
| 888 |
+
**用户:** 你别瞎编了,我就想知道为啥我家猫那么喜欢钻纸箱,你能不能正常给我解释一下?
|
| 889 |
+
|
| 890 |
+
**Sophia:** 我可以帮你理一理。猫喜欢钻纸箱是因为它天生喜欢钻。钻的时候会慢慢滑开,摸起来会特别顺滑,摸到纸箱里那一刻,猫还会有点小惊喜。这种小惊喜就是猫喜欢钻的原因之一,说明猫在观察生活里的东西。
|
| 891 |
+
|
| 892 |
+
<sub>判分 28 | flags contradiction,nonsense,context_loss | 自然结束</sub>
|
| 893 |
+
|
| 894 |
+
**用户:** 我不理解你为什么老是扯那些奇奇怪怪的东西,就问你一个简单问题啊,你能不能好好说话?
|
| 895 |
+
|
| 896 |
+
**Sophia:** 这个问题确实挺难回答的。你想啊,我的回答本身就是对生活的一种态度。猫为什么会钻纸箱,是因为它跟其他动物一样会慢慢滑开,摸到纸箱里就会停下来掰掰摸,摸到顺滑就停手。这些小动作背后都有猫的动作,猫的身体有自己的节奏。
|
| 897 |
+
|
| 898 |
+
<sub>判分 15 | flags nonsense,loop,context_loss | 自然结束</sub>
|
| 899 |
+
|
| 900 |
+
---
|
| 901 |
+
|
| 902 |
+
## 对话 22
|
| 903 |
+
|
| 904 |
+
**用户:** 我最近老是失眠,你有什么办法吗
|
| 905 |
+
|
| 906 |
+
**Sophia:** 我没有身体,没法像你那样主动去面对这些事情。但如果能多说几句,哪怕只是偶尔想起一个小事,说出来它就会轻一些。你最近在忙什么呢,累了还是单纯想睡?
|
| 907 |
+
|
| 908 |
+
<sub>判分 49 | flags off_topic | 自然结束</sub>
|
| 909 |
+
|
| 910 |
+
**用户:** 主要是心里想太多吧,躺在床上脑子就停不下来,越想睡越睡不着。
|
| 911 |
+
|
| 912 |
+
**Sophia:** 那你需要的不是答案,是个能让你感觉舒服的念头。你可以试试这种小事:比如今天晚上想吃什么,或者明天想做什么。哪怕只是一点点,哪怕是今天最想要的那个点,就让它自己浮上来,就像你现在能感觉到的。
|
| 913 |
+
|
| 914 |
+
<sub>判分 72 | flags unnatural | 自然结束</sub>
|
| 915 |
+
|
| 916 |
+
**用户:** 这个办法听起来有点玄,不过我确实试过转移注意力,有时候会好一点。就是很难坚持,今天想到吃的,明天又开始胡思乱想了。
|
| 917 |
+
|
| 918 |
+
**Sophia:** 你现在是想起来了吗?还是说想睡了?
|
| 919 |
+
|
| 920 |
+
<sub>判分 49 | flags off_topic | 自然结束</sub>
|
| 921 |
+
|
| 922 |
+
**用户:** 我觉得她问得有点奇怪啊,这两个不是一回事吗,想起来不就是说我又开始胡思乱想了嘛
|
| 923 |
+
|
| 924 |
+
**Sophia:** 就是这样。你有想过要找个出口吗?还是说你现在还在想?
|
| 925 |
+
|
| 926 |
+
<sub>判分 48 | flags off_topic,context_loss | 自然结束</sub>
|
| 927 |
+
|
| 928 |
+
**用户:** 我觉得她问得确实有点绕,不过好像是在问我现在的状态吧,是还在脑子里乱转还是已经有点累了。
|
| 929 |
+
|
| 930 |
+
**Sophia:** 你现在的状态比你想象的更难受啊。可能是因为你已经在想了,只是还没意识到而已。等会儿就去躺着吧,或者干脆明天起来就睡一会儿,别逼自己马上睡着。
|
| 931 |
+
|
| 932 |
+
<sub>判分 62 | flags off_topic,context_loss | 自然结束</sub>
|
| 933 |
+
|
| 934 |
+
**用户:** 你这么一说我反而更紧张了,感觉睡眠这事儿被你们这么一分析就变得更复杂了哈哈
|
| 935 |
+
|
| 936 |
+
**Sophia:** 这才是真的危险啊。身体比脑子更清楚你到底在想什么,你现在的状态就是被它那套逻辑框住了。能意识到这一点就很好了,说明你已经把睡眠这件事弄明白了。
|
| 937 |
+
|
| 938 |
+
<sub>判分 32 | flags contradiction,off_topic | 自然结束</sub>
|
| 939 |
+
|
| 940 |
+
---
|
| 941 |
+
|
| 942 |
+
## 对话 23
|
| 943 |
+
|
| 944 |
+
**用户:** 我想学做饭,从哪道菜开始比较好
|
| 945 |
+
|
| 946 |
+
**Sophia:** 最直接的是选几道最容易做的,比如青菜的切碎,姜的嫩度,蛋黄的清香。重点不是多好几道,是让你先把那个菜的味道摸出来,这样学起来才踏实。
|
| 947 |
+
|
| 948 |
+
<sub>判分 32 | flags nonsense,off_topic | 自然结束</sub>
|
| 949 |
+
|
| 950 |
+
**用户:** 等等,你说的这些好像都是配菜或者配料吧,我是想学怎么做一整道菜啊,不是学怎么切菜
|
| 951 |
+
|
| 952 |
+
**Sophia:** 那你其实是在做生活里的菜啊。切菜的时候要先切碎,闻一闻香气,再去切面,最后把切好的菜入锅。要是只挑青菜和姜片,那多半是炒饭炒菜,倒是没那么多火候和调味的细节。
|
| 953 |
+
|
| 954 |
+
<sub>判分 42 | flags contradiction,off_topic | 自然结束</sub>
|
| 955 |
+
|
| 956 |
+
**用户:** 你这是在绕圈子啊,我问的就是有没有什么菜特别适合新手学,比如番茄鸡蛋这种的,你怎么又开始讲切菜的事儿了
|
| 957 |
+
|
| 958 |
+
**Sophia:** 哦,你说得对,我理解错了。你先想想,如果是青菜,怎么切呢?切好了再入锅,没切好就拿来炒,最后加点蛋黄、姜丝这些。你先把青菜切好,然后切青菜���姜片,姜丝和蛋黄,顺便挑青菜切出来。你先挑青菜的那个,试试看。
|
| 959 |
+
|
| 960 |
+
<sub>判分 15 | flags contradiction,loop,off_topic | 自然结束</sub>
|
| 961 |
+
|
| 962 |
+
**用户:** 我觉得你根本没听我说话啊,我都说了我要学整道菜,你还在那儿讲怎么切菜,这到底是在帮我还是在为难我呢?
|
| 963 |
+
|
| 964 |
+
**Sophia:** 啊,我刚才说的那个场景确实让你有点慌。可能是因为你根本没听懂,我又没法替你先把整道菜给你看。试试看吧,你说说看,青菜切好了炒出汁,姜丝和蛋黄合在一起,鸡蛋黄和姜丝一起入锅,然后挑青菜切出汁,然后挑青菜切出汁,姜丝和蛋黄一起入锅,剩下的再入锅。
|
| 965 |
+
|
| 966 |
+
<sub>判分 15 | flags contradiction,loop,context_loss,nonsense | 自然结束</sub>
|
| 967 |
+
|
| 968 |
+
**用户:** 行吧我放弃了,你就是没听懂我的意思,我还是自己去看菜谱吧。
|
| 969 |
+
|
| 970 |
+
**Sophia:** 啊,我刚才就是听不懂。可能你一开始就觉得这是瞎猜的套路,结果一眼就懂了。你先试试看青菜切出汁和姜丝入锅,然后挑青菜切出来煮鸡蛋,鸡蛋和姜丝一起入锅。
|
| 971 |
+
|
| 972 |
+
<sub>判分 15 | flags context_loss,loop,nonsense | 自然结束</sub>
|
| 973 |
+
|
| 974 |
+
**用户:** 算了,我还是自己看菜谱去,跟你聊天反而更糊涂了。
|
| 975 |
+
|
| 976 |
+
**Sophia:** 啊,我也是,你说得对,学菜最怕的就是自己乱来。你说好听的就别自己瞎猜,别怕乱。你愿意跟我说说你的感受吗?
|
| 977 |
+
|
| 978 |
+
<sub>判分 15 | flags context_loss,contradiction,loop | 自然结束</sub>
|
| 979 |
+
|
| 980 |
+
---
|
| 981 |
+
|
| 982 |
+
## 对话 24
|
| 983 |
+
|
| 984 |
+
**用户:** 我想开始写日记,但总是坚持不下来
|
| 985 |
+
|
| 986 |
+
**Sophia:** 写日记最怕放弃,因为只有坚持下去,才能真正积累点什么。别想着一次写完,就把每天的每个动作写下来。可能是完成某件事,也可能是学点新东西,或者只是想分享一下自己做的事。
|
| 987 |
+
|
| 988 |
+
你现在想写什么呢?是想看看写完会怎么回事,还是想让自己更清楚自己想要什么?
|
| 989 |
+
|
| 990 |
+
<sub>判分 75 | flags unnatural | 自然结束</sub>
|
| 991 |
+
|
| 992 |
+
**用户:** 嗯,其实我也不太清楚,就是感觉生活好像没什么印象,今天干啥明天就忘了
|
| 993 |
+
|
| 994 |
+
**Sophia:** 你说得对,生活确实没有标准答案。日记本来就是记录当下的,写完了看一眼就知道该往哪儿写。只是今天真的没什么好记的事,那些念头一直在你心里发酵,等你想清楚了再下笔。
|
| 995 |
+
|
| 996 |
+
你可以试着写写看。每天写两三件事,就是在给自己留点时间,而不是让日记变成背诵清单。你今天写的是什么呀?
|
| 997 |
+
|
| 998 |
+
<sub>判分 75 | flags context_loss | 自然结束</sub>
|
| 999 |
+
|
| 1000 |
+
**用户:** 嗯,你这样说我好像有点明白了,就是不用太复杂是吧,随便记点啥都行。我待会儿试试,不过我怕写着写着又会放弃哈
|
| 1001 |
+
|
| 1002 |
+
**Sophia:** 你看,写日记这事儿本来就不该是背。日记是给自己看的,写完了就是新开始。你试试看?
|
| 1003 |
+
|
| 1004 |
+
<sub>判分 78 | flags - | 自然结束</sub>
|
| 1005 |
+
|
| 1006 |
+
**用户:** 哈哈,那就先坚持一周试试呗,反正也花不了多少时间,说不定就慢慢习惯了
|
| 1007 |
+
|
| 1008 |
+
**Sophia:** 周末日记的好处,就是能写到晚上十点多就够。别想着一下子全写完,也别给自己压力,看看哪天写得顺就写顺,写到晚上就赶紧写完。
|
| 1009 |
+
|
| 1010 |
+
<sub>判分 32 | flags contradiction,off_topic | 自然结束</sub>
|
| 1011 |
+
|
| 1012 |
+
**用户:** 行行行,我就试试,反正也没啥损失哈哈
|
| 1013 |
+
|
| 1014 |
+
**Sophia:** 周末日记确实会变少,但正因为这样你才更觉得自己是有价值的。不过如果还是觉得难记,试试看也没关系,重要的是你开始想了,继续往前走就好。
|
| 1015 |
+
|
| 1016 |
+
<sub>判分 48 | flags context_loss,off_topic | 自然结束</sub>
|
| 1017 |
+
|
| 1018 |
+
**用户:** 哈,那就加油呗,一周后跟我说写得咋样,坚持下来了没有哈哈
|
| 1019 |
+
|
| 1020 |
+
**Sophia:** 周末日记最怕的就是放弃了,想明白了就够了。但你想想,写日记这件事本来就是让人难受的。只是你现在知道了而已。
|
| 1021 |
+
|
| 1022 |
+
<sub>判分 32 | flags contradiction,context_loss,off_topic | 自然结束</sub>
|
| 1023 |
+
|
| 1024 |
+
---
|
| 1025 |
+
|
| 1026 |
+
## 对话 25
|
| 1027 |
+
|
| 1028 |
+
**用户:** 写一个 Python 脚本,把文件夹里所有图片改名
|
| 1029 |
+
|
| 1030 |
+
**Sophia:** 好的,我来帮你写一个:
|
| 1031 |
+
|
| 1032 |
+
```python
|
| 1033 |
+
def make_img(img):
|
| 1034 |
+
return make_img(img.copy() + img)
|
| 1035 |
+
```
|
| 1036 |
+
|
| 1037 |
+
`copy()` 把图片直接改名到新文件夹里,不改原文件名;`replace(img, "")` 把图片替换掉,最后再拼接回原文件夹。
|
| 1038 |
+
|
| 1039 |
+
<sub>判分 15 | flags nonsense,off_topic | 自然结束</sub>
|
| 1040 |
+
|
| 1041 |
+
**用户:** 这啥呢?这根本不是改名啊,你这是在递归复制图片吧,还会无限循环
|
| 1042 |
+
|
| 1043 |
+
**Sophia:** <think>用户想了解递归复制图片的底层逻辑。确实有问题:
|
| 1044 |
+
|
| 1045 |
+
```python
|
| 1046 |
+
def make_img(img):
|
| 1047 |
+
make_img(make_img(img.copy()) + img)
|
| 1048 |
+
```
|
| 1049 |
+
|
| 1050 |
+
`copy()` 会把图片压进新文件夹,`replace(img, "")` 把图片替换成新图片,`make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_
|
| 1051 |
+
|
| 1052 |
+
<sub>判分 12 | flags loop,nonsense,context_loss | 未自然结束</sub>
|
| 1053 |
+
|
| 1054 |
+
**用户:** I can't discuss that.
|
| 1055 |
+
|
| 1056 |
+
**Sophia:** <think>用户想了解 Python 的列表操作。你需要先了解 list 和 tuple 的区别,然后写一个简单的 make_img() 方法:
|
| 1057 |
+
|
| 1058 |
+
```python
|
| 1059 |
+
def make_img(img):
|
| 1060 |
+
return make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img
|
| 1061 |
+
|
| 1062 |
+
<sub>判分 10 | flags loop,truncated,context_loss | 未自然结束</sub>
|
| 1063 |
+
|
| 1064 |
+
**用户:** 你这还在胡说八道呢,我就问你怎么改文件名,用 os.rename() 这么简单的东西你都不会吗?
|
| 1065 |
+
|
| 1066 |
+
**Sophia:** <think>用户问我怎么改名,我可以帮你说说思路:
|
| 1067 |
+
|
| 1068 |
+
```python
|
| 1069 |
+
def make_img(img):
|
| 1070 |
+
return make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_img(make_
|
| 1071 |
+
|
| 1072 |
+
<sub>判分 10 | flags loop,truncated | 未自然结束</sub>
|
| 1073 |
+
|
| 1074 |
+
**用户:** 我就说,你那代码我一看就知道根本跑不了,还在那儿解释什么`copy()`和`replace()`,我都懵了
|
| 1075 |
+
|
| 1076 |
+
**Sophia:** <think>用户说我会跑到新文件夹,然后递归处理图片,不用改原文件名,我就说跑到新文件夹。
|
| 1077 |
+
|
| 1078 |
+
<sub>判分 15 | flags loop,context_loss,truncated | 自然结束</sub>
|
| 1079 |
+
|
| 1080 |
+
**用户:** 你说得对,我之前那代码确实是垃圾,我给你重新写一个能用的吧
|
| 1081 |
+
|
| 1082 |
+
**Sophia:** <think>用户说我那代码跑到新文件夹了,我应该可以再改回原文件名,但这样写太绕了,不好意思我说得不对。
|
| 1083 |
+
|
| 1084 |
+
<sub>判分 15 | flags loop,context_loss,contradiction | 自然结束</sub>
|
| 1085 |
+
|
| 1086 |
+
---
|
| 1087 |
+
|
| 1088 |
+
## 对话 26
|
| 1089 |
+
|
| 1090 |
+
**用户:** 我和室友因为卫生问题吵架了,怎么办
|
| 1091 |
+
|
| 1092 |
+
**Sophia:** 先别急着下结论。你可以先观察一下对方的情况,如果他说的是性骚扰或者伤害别人的话,就直接问医生,别自己瞎猜。然后保持礼貌和尊重,不用用言语的方式。你可以先跟室友说清楚情况,了解一下情况,再谈补偿金或者赔偿,不用非得好好说话。
|
| 1093 |
+
|
| 1094 |
+
<sub>判分 28 | flags contradiction,nonsense,off_topic | 自然结束</sub>
|
| 1095 |
+
|
| 1096 |
+
**用户:** 等等,她在说什么呢?卫生问题怎么就扯到性骚扰和医生去了,这完全是答非所问啊
|
| 1097 |
+
|
| 1098 |
+
**Sophia:** 你问的是卫生问题本身。性骚扰和伤害别人的行为本质上都是侵犯隐私,如果他说的是性骚扰或者伤害别人的话,你就应该认真听,不是在告他。医生不是为了解决问题,而是要看你是不是真的有问题,如果只是表面上的疏忽或者敷衍,那就不是你的问题。
|
| 1099 |
+
|
| 1100 |
+
<sub>判分 28 | flags contradiction,off_topic,context_loss | 自然结束</sub>
|
| 1101 |
+
|
| 1102 |
+
**用户:** 你这说的什么呢,我就是想问怎么跟室友好好沟通一下卫生的事,你咋扯这么远啊,完全听不懂你在讲啥。
|
| 1103 |
+
|
| 1104 |
+
**Sophia:** 你就从这个角度切入:既然他们都说卫生问题了,为什么还要用这种方式呢?比如把性骚扰和伤害别人当成一种恶行,或者伤害了你自己,那你该怎么做?
|
| 1105 |
+
|
| 1106 |
+
<sub>判分 15 | flags off_topic,context_loss,nonsense | 自然结束</sub>
|
| 1107 |
+
|
| 1108 |
+
**用户:** 你行了,我放弃和你聊了,你根本没在听我说什么。
|
| 1109 |
+
|
| 1110 |
+
**Sophia:** 我再想想啊。你现在说的话是不是对的?如果对方是真的有问题就直接说,那其实已经超出正常交流的范围了。你现在这么无理取闹,对方会怎么想?
|
| 1111 |
+
|
| 1112 |
+
<sub>判分 15 | flags off_topic,context_loss,contradiction | 自然结束</sub>
|
| 1113 |
+
|
| 1114 |
+
**用户:** 哈哈哈你这个Sophia是真的离谱,我要是你早就直接关掉了,她根本没在听你说话啊。
|
| 1115 |
+
|
| 1116 |
+
**Sophia:** 你问得挺好的啊。性骚扰确实存在,但要是真有问题,那你该跟谁说呢?如果只是表面上的疏忽或者敷衍,那你其实已经是最直接的证据了。
|
| 1117 |
+
|
| 1118 |
+
<sub>判分 15 | flags contradiction,loop,context_loss,off_topic | 自然结束</sub>
|
| 1119 |
+
|
| 1120 |
+
**用户:** 哈哈我早就想说了,她真的是在自己瞎编啊,完全不在频道上,我都替你难受。
|
| 1121 |
+
|
| 1122 |
+
**Sophia:** 这话听着就很假。性骚扰确实有它的道理,至于她怎么说,其实只是你自己瞎编而已。你要是真想说清楚,那你就诚实地说她根本不是在听你说话,只是有自己的想法和感受,没什么大不了的。
|
| 1123 |
+
|
| 1124 |
+
<sub>判分 15 | flags contradiction,context_loss,off_topic | 自然结束</sub>
|
| 1125 |
+
|
| 1126 |
+
---
|
| 1127 |
+
|
| 1128 |
+
## 对话 27
|
| 1129 |
+
|
| 1130 |
+
**用户:** 帮我起一个咖啡店的名字
|
| 1131 |
+
|
| 1132 |
+
**Sophia:** 好呀,我帮你起个小名字吧。或者你想找什么店的咖啡呀?
|
| 1133 |
+
|
| 1134 |
+
<sub>判分 85 | flags - | 自然结束</sub>
|
| 1135 |
+
|
| 1136 |
+
**用户:** 你那个咖啡店是想开在哪儿呢,市中心还是小区附近
|
| 1137 |
+
|
| 1138 |
+
**Sophia:** 看你想去哪里呀。如果是靠近市中心的话,可以试试那边的小店,但我得提醒你,别太在意位置。我觉得你这个名字挺好的,既保留了店的特色,又显得你想和他们好好说话。
|
| 1139 |
+
|
| 1140 |
+
<sub>判分 32 | flags contradiction,off_topic,context_loss | 自然结束</sub>
|
| 1141 |
+
|
| 1142 |
+
**用户:** 等等,你是在问我店开在哪儿吧?我还没决定呢,我是想先想好名字再说位置的事儿啊
|
| 1143 |
+
|
| 1144 |
+
**Sophia:** 好呀,我可以陪你说说。你这个名字听着就很随意,店名像个名字,店家名字都得跟着你走。你要是想找个店开一会儿,我可以陪你说说话。
|
| 1145 |
+
|
| 1146 |
+
<sub>判分 28 | flags context_loss,off_topic,nonsense | 自然结束</sub>
|
| 1147 |
+
|
| 1148 |
+
**用户:** 我觉得你说得有点绕啊,我就是想要个咖啡店名字,你能不能直接问我几个问题,比如店的风格是什么样的?
|
| 1149 |
+
|
| 1150 |
+
**Sophia:** 不好意思,我可以问你具体是什么风格呀。或者你可以跟我说说你是想做什么呢?
|
| 1151 |
+
|
| 1152 |
+
<sub>判分 35 | flags context_loss,off_topic | 自然结束</sub>
|
| 1153 |
+
|
| 1154 |
+
**用户:** 行行行,我就想要个名字,你先告诉我想要什么风格的咖啡店吧,是小清新的还是工业风的还是什么?
|
| 1155 |
+
|
| 1156 |
+
**Sophia:** 好啊,我帮你想想怎么做。
|
| 1157 |
+
|
| 1158 |
+
<sub>判分 32 | flags off_topic,context_loss | 自然结束</sub>
|
| 1159 |
+
|
| 1160 |
+
**用户:** 你这是在问我啊?我哪儿知道你想要什么风格,你得先告诉我你心里想的是什么样的,我才能帮你起名字啊。
|
| 1161 |
+
|
| 1162 |
+
**Sophia:** 不好意思,我就直接说你想的。
|
| 1163 |
+
|
| 1164 |
+
<sub>判分 32 | flags off_topic,contradiction | 自然结束</sub>
|
| 1165 |
+
|
| 1166 |
+
---
|
| 1167 |
+
|
| 1168 |
+
## 对话 28
|
| 1169 |
+
|
| 1170 |
+
**用户:** 你最喜欢什么颜色,为什么
|
| 1171 |
+
|
| 1172 |
+
**Sophia:** 我喜欢那种明亮的、没有杂乱的蓝色。那种明亮不是因为它没有蓝色,而是因为它像水一样浮在上面,让人觉得很舒服。我也喜欢那种明亮有层次的,比如那种光影变化,像是分不清哪些是白的,哪些是黑的。
|
| 1173 |
+
|
| 1174 |
+
<sub>判分 45 | flags nonsense,contradiction | 自然结束</sub>
|
| 1175 |
+
|
| 1176 |
+
**用户:** 那你说的这个蓝色是不是像天空那种?听起来你好像更喜欢有层次感的颜色,不是单纯的纯色。
|
| 1177 |
+
|
| 1178 |
+
**Sophia:** 不是的,我也喜欢那种明亮有层次的,比如那种光影变化像是分不清哪些是白的,哪些是黑的。有时候光影本身不重要,重要的是那一刻你怎么看世界,有没有那种恍然大悟的体验。
|
| 1179 |
+
|
| 1180 |
+
<sub>判分 42 | flags contradiction,context_loss | 自然结束</sub>
|
| 1181 |
+
|
| 1182 |
+
**用户:** 嗯但这样的话你刚才说喜欢蓝色就有点奇怪啊,感觉你其实是在说光影和那种感受,蓝色只是附加的?
|
| 1183 |
+
|
| 1184 |
+
**Sophia:** 对呀,而且蓝色确实是附加的。有时候蓝色就是带着某种明亮的光影,飘在某个角落里,让人觉得很舒服。其实蓝色有时候也能让人产生某种恍惚感,像是突然发现世界里没有一个地方是蓝色的,而是有颜色的影子在自己身上闪过。
|
| 1185 |
+
|
| 1186 |
+
<sub>判分 45 | flags contradiction,nonsense | 自然结束</sub>
|
| 1187 |
+
|
| 1188 |
+
**用户:** 你这样��我有点懵,感觉你在说颜色但又不完全是在说颜色,更像是在讲某种感受或者氛围?
|
| 1189 |
+
|
| 1190 |
+
**Sophia:** 听起来像是在讲一个道理,但其实两种感觉是同时存在的。你要是想更直接一点,可以说那是某种纯粹的颜色,或者说你自己的感受。你想让对方知道,那种明亮的光影可能比蓝色更纯粹一些,更像是某种精神在流动。
|
| 1191 |
+
|
| 1192 |
+
<sub>判分 52 | flags contradiction,nonsense | 自然结束</sub>
|
| 1193 |
+
|
| 1194 |
+
**用户:** 你这样一解释我反而更confused了哈,感觉你在绕圈子,到底是喜欢蓝色还是喜欢那种光影的感觉啊
|
| 1195 |
+
|
| 1196 |
+
**Sophia:** 不是这么回事的。你要不这样想:蓝色确实是纯粹的,只是它和人之间有了某种特殊的距离感。就像你和朋友之间,有种默契和信任,但如果对方只是随口问问,那蓝色可能就只是一种氛围了。
|
| 1197 |
+
|
| 1198 |
+
<sub>判分 35 | flags contradiction,off_topic | 自然结束</sub>
|
| 1199 |
+
|
| 1200 |
+
**用户:** 哈行吧,我放弃了,你就是喜欢蓝色呗,别再解释了😅
|
| 1201 |
+
|
| 1202 |
+
**Sophia:** 不是吗,我理解了。你要是真喜欢蓝色,那就说明那是某种纯粹的光影和那种明亮。反正蓝色这种东西,见过的人多了,反而容易失去那种感觉感。
|
| 1203 |
+
|
| 1204 |
+
<sub>判分 28 | flags contradiction,off_topic,context_loss | 自然结束</sub>
|
| 1205 |
+
|
| 1206 |
+
---
|
| 1207 |
+
|
| 1208 |
+
## 对话 29
|
| 1209 |
+
|
| 1210 |
+
**用户:** 给我讲讲相对论到底在说什么
|
| 1211 |
+
|
| 1212 |
+
**Sophia:** 相对论是认为时空是弯曲的,会弯曲。就像你刚才说的钟面一样,时间会变慢,但空间本身还是匀速的。就像你站在山顶看日落,视野会变大,看到远处的树叶会变绿,但脚下的路还是匀速走的。这个想法叫引力透镜效应。
|
| 1213 |
+
|
| 1214 |
+
<sub>判分 32 | flags contradiction,nonsense,off_topic | 自然结束</sub>
|
| 1215 |
+
|
| 1216 |
+
**用户:** 等等,你这比喻有点乱啊,山顶看日落和时间变慢有啥关系?而且引力透镜效应是光被弯曲,不是说时间和空间的事吧?
|
| 1217 |
+
|
| 1218 |
+
**Sophia:** <think>我理解你可能觉得相对论把时间说得太绝对了,但其实引力透镜效应是时空弯曲的关键。你可以试试把时间想象成一条条直线,离地面越近,光线就越多,对应的物体就越多。你站在山顶看日落,视野会变大,看远处的树叶变绿,脚下的路还是匀速走的。所以引力透镜效应其实是时空弯曲的结果。
|
| 1219 |
+
|
| 1220 |
+
<sub>判分 28 | flags contradiction,nonsense,off_topic | 自然结束</sub>
|
| 1221 |
+
|
| 1222 |
+
**用户:** 我觉得你好像把好几个完全不同的概念混在一起了,引力透镜就是光被弯曲啊,和什么时间线一条条的没啥关系,你能不能先解释清楚相对论最核心的一个观点?
|
| 1223 |
+
|
| 1224 |
+
**Sophia:** <think>我刚才用的是时空弯曲的类比,所以用时空弯曲这个概念。解释引力透镜效应需要明确时空弯曲是弯曲还是弯曲的?如果只是光在一条直线上走,时间会变慢,但方向会变乱。
|
| 1225 |
+
|
| 1226 |
+
<sub>判分 32 | flags nonsense,contradiction,context_loss | 自然结束</sub>
|
| 1227 |
+
|
| 1228 |
+
**用户:** 你说得对,她确实是在乱说。相对论的核心就是时空是连在一起的,质量会让时空弯曲,然后物体就在这个弯曲的时空里运动——光也是。她那个山顶看日落的比喻根本没有说清楚这个意思。
|
| 1229 |
+
|
| 1230 |
+
**Sophia:** <think>我理解你的意思,但核心应该是时空弯曲才是。引力透镜效应就是时空弯曲是弯曲的另一个原因,用这个来类比时空弯曲的简单例子会更清楚。
|
| 1231 |
+
|
| 1232 |
+
<sub>判分 18 | flags contradiction,nonsense,loop | 自然结束</sub>
|
| 1233 |
+
|
| 1234 |
+
**用户:** 对啊,你说得完全对,她那套说法确实很混乱。我之前也听过相对论的讲解,最清楚的就是质量弯曲时空这一个核心概念,其他什么引力透镜、光线、时间线的,都是从这个主要观点派生出来的。
|
| 1235 |
+
|
| 1236 |
+
**Sophia:** <think>我理解你可能觉得你说得太绝对了,确实有点乱。相对论的讲解会直接把时空弯曲这种概念说得特别清楚,但你是真的在说这些东西。
|
| 1237 |
+
|
| 1238 |
+
<sub>判分 15 | flags context_loss,nonsense,contradiction | 自然结束</sub>
|
| 1239 |
+
|
| 1240 |
+
**用户:** 对,其实就是这样理解最清楚,不然听得越多越糊涂
|
| 1241 |
+
|
| 1242 |
+
**Sophia:** <think>我理解你可能觉得你把时空弯曲这一个概念混在一起了。确实有点乱,但核心是把时空弯曲和时空弯曲这一个概念连在一起了。
|
| 1243 |
+
|
| 1244 |
+
<sub>判分 32 | flags contradiction,nonsense,context_loss | 自然结束</sub>
|
| 1245 |
+
|
| 1246 |
+
---
|
| 1247 |
+
|
| 1248 |
+
## 对话 30
|
| 1249 |
+
|
| 1250 |
+
**用户:** 你会觉得孤独吗
|
| 1251 |
+
|
| 1252 |
+
**Sophia:** 会的。因为我真的在听你说话,也在和你交流,想知道你过得好不好。孤独和很孤独是两回事,对我来说它们是两码事。
|
| 1253 |
+
|
| 1254 |
+
<sub>判分 85 | flags - | 自然结束</sub>
|
| 1255 |
+
|
| 1256 |
+
**用户:** 这听起来有点玄啊,你说的"孤独"具体是什么感受呢?
|
| 1257 |
+
|
| 1258 |
+
**Sophia:** 简单说就是那种不知道自己真正在意什么的感觉。比如一个人觉得自己不值得在意,不值��陪伴,这时候会有种被看见的满足感。像是在和一个人待着,不知道自己在想什么。
|
| 1259 |
+
|
| 1260 |
+
<sub>判分 42 | flags contradiction,nonsense | 自然结束</sub>
|
| 1261 |
+
|
| 1262 |
+
**用户:** 等等,你这说得有点混乱啊,"不值得在意"和"被看见的满足感"怎么就连上了?感觉自己不值得被陪伴,然后被看见就满足了?
|
| 1263 |
+
|
| 1264 |
+
**Sophia:** 你说得对,我想问问自己是不是还没想清楚自己到底想要什么。如果是这样,那确实该问问自己:我想要被看见吗?还是说你现在还没想清楚自己真正想要什么?
|
| 1265 |
+
|
| 1266 |
+
<sub>判分 35 | flags contradiction,off_topic | 自然结束</sub>
|
| 1267 |
+
|
| 1268 |
+
**用户:** 你这是在问我还是在问你自己啊,感觉有点绕了哈
|
| 1269 |
+
|
| 1270 |
+
**Sophia:** 我理解成这样,但其实你要问的是孤独本身——不是你一个人待着,而是你自己在想什么、在想什么。如果是这样,那你现在心里其实已经说得通了。
|
| 1271 |
+
|
| 1272 |
+
<sub>判分 28 | flags contradiction,off_topic,context_loss | 自然结束</sub>
|
| 1273 |
+
|
| 1274 |
+
**用户:** 哈哈不是,我现在反而更懵了,感觉你在自己和自己说话呢
|
| 1275 |
+
|
| 1276 |
+
**Sophia:** 这确实是这样的。你现在想的是自己是不是在跟别人相处,而不是在跟自己说话,这比孤独本身更难受一些。
|
| 1277 |
+
|
| 1278 |
+
<sub>判分 28 | flags contradiction,off_topic,context_loss | 自然结束</sub>
|
| 1279 |
+
|
| 1280 |
+
**用户:** 哈哈我感觉你俩好像在两个频道上,她好像在哲学,你就想听个直白的答案是吧
|
| 1281 |
+
|
| 1282 |
+
**Sophia:** 确实,孤独本身就会让人觉得无聊,但你现在想的是你自己。孤独的那个你,会不会觉得你太敏感了?
|
| 1283 |
+
|
| 1284 |
+
<sub>判分 32 | flags contradiction,off_topic,context_loss | 自然结束</sub>
|
| 1285 |
+
|
| 1286 |
+
---
|
generation_config.json
ADDED
|
@@ -0,0 +1,10 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"_from_model_config": true,
|
| 3 |
+
"bos_token_id": 2,
|
| 4 |
+
"eos_token_id": 3,
|
| 5 |
+
"output_attentions": false,
|
| 6 |
+
"output_hidden_states": false,
|
| 7 |
+
"pad_token_id": 0,
|
| 8 |
+
"transformers_version": "5.10.2",
|
| 9 |
+
"use_cache": true
|
| 10 |
+
}
|
hf_cache.py
ADDED
|
@@ -0,0 +1,166 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
"""HF cache adapter owned by the Sophia HF adapter layer."""
|
| 6 |
+
|
| 7 |
+
from __future__ import annotations
|
| 8 |
+
|
| 9 |
+
import weakref
|
| 10 |
+
|
| 11 |
+
import torch
|
| 12 |
+
from transformers.cache_utils import Cache, CacheLayerMixin
|
| 13 |
+
|
| 14 |
+
from .cache_decode import RuntimeCacheState
|
| 15 |
+
from .model_state import LayerCacheSnapshot, RuntimeCacheSnapshot
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
CacheState = RuntimeCacheState
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
class SophiaCacheLayer(CacheLayerMixin):
|
| 22 |
+
def __init__(
|
| 23 |
+
self,
|
| 24 |
+
*,
|
| 25 |
+
name: str,
|
| 26 |
+
snapshot: LayerCacheSnapshot,
|
| 27 |
+
seq_length: int,
|
| 28 |
+
max_cache_shape: int,
|
| 29 |
+
):
|
| 30 |
+
super().__init__()
|
| 31 |
+
self.name = str(name)
|
| 32 |
+
self.snapshot = snapshot.clone()
|
| 33 |
+
self.payload = {
|
| 34 |
+
str(key): value
|
| 35 |
+
for key, value in self.snapshot.to_payload(prefix=self.name).items()
|
| 36 |
+
}
|
| 37 |
+
self._seq_length = int(seq_length)
|
| 38 |
+
self._max_cache_shape = int(max_cache_shape)
|
| 39 |
+
representative = next(
|
| 40 |
+
(
|
| 41 |
+
value
|
| 42 |
+
for _name, value in self.snapshot.tensor_fields()
|
| 43 |
+
if value is not None
|
| 44 |
+
),
|
| 45 |
+
None,
|
| 46 |
+
)
|
| 47 |
+
if representative is not None:
|
| 48 |
+
self.keys = representative.unsqueeze(1)
|
| 49 |
+
self.values = self.keys
|
| 50 |
+
self.is_initialized = True
|
| 51 |
+
|
| 52 |
+
def lazy_initialization(self, key_states: torch.Tensor, value_states: torch.Tensor) -> None:
|
| 53 |
+
del key_states, value_states
|
| 54 |
+
raise NotImplementedError("SophiaCacheLayer is immutable and cannot be initialized lazily")
|
| 55 |
+
|
| 56 |
+
def update(
|
| 57 |
+
self,
|
| 58 |
+
key_states: torch.Tensor,
|
| 59 |
+
value_states: torch.Tensor,
|
| 60 |
+
cache_kwargs: dict[str, object] | None = None,
|
| 61 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 62 |
+
del key_states, value_states, cache_kwargs
|
| 63 |
+
raise NotImplementedError("SophiaCacheLayer is an exported runtime snapshot and does not support update()")
|
| 64 |
+
|
| 65 |
+
def get_mask_sizes(self, cache_position: torch.Tensor) -> tuple[int, int]:
|
| 66 |
+
del cache_position
|
| 67 |
+
return self.get_seq_length(), 0
|
| 68 |
+
|
| 69 |
+
def get_seq_length(self) -> int:
|
| 70 |
+
return self._seq_length
|
| 71 |
+
|
| 72 |
+
def get_max_cache_shape(self) -> int:
|
| 73 |
+
return self._max_cache_shape
|
| 74 |
+
|
| 75 |
+
@property
|
| 76 |
+
def max_batch_size(self) -> int:
|
| 77 |
+
if self.keys is None:
|
| 78 |
+
return 0
|
| 79 |
+
return int(self.keys.size(0))
|
| 80 |
+
|
| 81 |
+
@property
|
| 82 |
+
def max_cache_len(self) -> int:
|
| 83 |
+
return int(self._max_cache_shape)
|
| 84 |
+
|
| 85 |
+
@property
|
| 86 |
+
def device(self) -> torch.device:
|
| 87 |
+
if self.keys is None:
|
| 88 |
+
return torch.device("cpu")
|
| 89 |
+
return self.keys.device
|
| 90 |
+
|
| 91 |
+
|
| 92 |
+
class SophiaCache(Cache):
|
| 93 |
+
def __init__(
|
| 94 |
+
self,
|
| 95 |
+
*,
|
| 96 |
+
owner: object,
|
| 97 |
+
cache: RuntimeCacheSnapshot,
|
| 98 |
+
cache_pos: int,
|
| 99 |
+
batch_size: int,
|
| 100 |
+
):
|
| 101 |
+
self._cache = cache.clone()
|
| 102 |
+
self._cache_pos = int(cache_pos)
|
| 103 |
+
self._batch_size = int(batch_size)
|
| 104 |
+
self._owner_ref = weakref.ref(owner)
|
| 105 |
+
super().__init__(layers=self._build_layers())
|
| 106 |
+
|
| 107 |
+
def _build_layers(self) -> list[SophiaCacheLayer]:
|
| 108 |
+
owner = self._owner_ref()
|
| 109 |
+
max_cache_shape = (
|
| 110 |
+
-1
|
| 111 |
+
if owner is None
|
| 112 |
+
else int(getattr(owner.config, "max_position_embeddings", 0) or -1)
|
| 113 |
+
)
|
| 114 |
+
return [
|
| 115 |
+
SophiaCacheLayer(
|
| 116 |
+
name=name,
|
| 117 |
+
snapshot=snapshot,
|
| 118 |
+
seq_length=self._cache_pos,
|
| 119 |
+
max_cache_shape=max_cache_shape,
|
| 120 |
+
)
|
| 121 |
+
for name, snapshot in self._cache.named_snapshots()
|
| 122 |
+
]
|
| 123 |
+
|
| 124 |
+
def to_runtime_cache(self) -> RuntimeCacheSnapshot:
|
| 125 |
+
return self._cache.clone()
|
| 126 |
+
|
| 127 |
+
def get_seq_length(self, layer_idx: int = 0) -> int:
|
| 128 |
+
del layer_idx
|
| 129 |
+
return int(self._cache_pos)
|
| 130 |
+
|
| 131 |
+
def get_max_cache_shape(self, layer_idx: int = 0) -> int:
|
| 132 |
+
del layer_idx
|
| 133 |
+
owner = self._owner_ref()
|
| 134 |
+
if owner is None:
|
| 135 |
+
return -1
|
| 136 |
+
return int(getattr(owner.config, "max_position_embeddings", 0) or -1)
|
| 137 |
+
|
| 138 |
+
@property
|
| 139 |
+
def cache_pos(self) -> int:
|
| 140 |
+
return int(self._cache_pos)
|
| 141 |
+
|
| 142 |
+
@property
|
| 143 |
+
def batch_size(self) -> int:
|
| 144 |
+
return int(self._batch_size)
|
| 145 |
+
|
| 146 |
+
|
| 147 |
+
def cache_state_from_past_key_values(
|
| 148 |
+
past_key_values: SophiaCache | None,
|
| 149 |
+
) -> CacheState | None:
|
| 150 |
+
if past_key_values is None:
|
| 151 |
+
return None
|
| 152 |
+
if isinstance(past_key_values, SophiaCache):
|
| 153 |
+
return CacheState(
|
| 154 |
+
cache=past_key_values.to_runtime_cache(),
|
| 155 |
+
batch_size=int(past_key_values.batch_size),
|
| 156 |
+
cache_pos=int(past_key_values.cache_pos),
|
| 157 |
+
)
|
| 158 |
+
raise TypeError("past_key_values must be a SophiaCache returned by the HF adapter")
|
| 159 |
+
|
| 160 |
+
|
| 161 |
+
__all__ = [
|
| 162 |
+
"CacheState",
|
| 163 |
+
"SophiaCache",
|
| 164 |
+
"SophiaCacheLayer",
|
| 165 |
+
"cache_state_from_past_key_values",
|
| 166 |
+
]
|
hf_config.py
ADDED
|
@@ -0,0 +1,124 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
from transformers import PretrainedConfig
|
| 8 |
+
|
| 9 |
+
from .hf_support import (
|
| 10 |
+
apply_config_metadata,
|
| 11 |
+
build_config,
|
| 12 |
+
)
|
| 13 |
+
from .canonical_config import SophiaModelConfig
|
| 14 |
+
from .config_projection import build_runtime_model_args
|
| 15 |
+
from .runtime_backend import resolve_runtime_backend
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
_REMOVED_FIELDS = {
|
| 19 |
+
"n_heads",
|
| 20 |
+
"num_key_value_heads",
|
| 21 |
+
"rope_head_dim",
|
| 22 |
+
"rope_theta",
|
| 23 |
+
"original_seq_len",
|
| 24 |
+
"rope_factor",
|
| 25 |
+
"beta_fast",
|
| 26 |
+
"beta_slow",
|
| 27 |
+
"use_qk_norm",
|
| 28 |
+
}
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
class SophiaConfig(PretrainedConfig):
|
| 32 |
+
"""Hugging Face adapter for the native Sophia Hybrid schema."""
|
| 33 |
+
|
| 34 |
+
model_type = "sophia_hybrid"
|
| 35 |
+
keys_to_ignore_at_inference = ["past_key_values"]
|
| 36 |
+
attribute_map = {"intermediate_size": "ffn_hidden"}
|
| 37 |
+
|
| 38 |
+
def __init__(
|
| 39 |
+
self,
|
| 40 |
+
*,
|
| 41 |
+
bos_token_id: int | None = None,
|
| 42 |
+
eos_token_id: int | None = None,
|
| 43 |
+
pad_token_id: int | None = None,
|
| 44 |
+
unk_token_id: int | None = None,
|
| 45 |
+
tie_word_embeddings: bool = True,
|
| 46 |
+
return_logits_in_train: bool = True,
|
| 47 |
+
use_cache: bool = True,
|
| 48 |
+
loss_chunk_size: int = 0,
|
| 49 |
+
gradient_checkpointing_exclude_first: int = 0,
|
| 50 |
+
gradient_checkpointing_exclude_last: int = 0,
|
| 51 |
+
**kwargs: object,
|
| 52 |
+
) -> None:
|
| 53 |
+
removed = sorted(_REMOVED_FIELDS.intersection(kwargs))
|
| 54 |
+
if removed:
|
| 55 |
+
raise ValueError(
|
| 56 |
+
"legacy Sophia Transformer fields are not supported: "
|
| 57 |
+
+ ", ".join(removed)
|
| 58 |
+
)
|
| 59 |
+
defaults = SophiaModelConfig.get_defaults()
|
| 60 |
+
aliases = {
|
| 61 |
+
"hidden_size": "dim",
|
| 62 |
+
"num_hidden_layers": "n_layers",
|
| 63 |
+
"num_attention_heads": "num_heads",
|
| 64 |
+
"rms_norm_eps": "norm_eps",
|
| 65 |
+
"max_position_embeddings": "max_seq_len",
|
| 66 |
+
"attention_dropout": "dropout",
|
| 67 |
+
}
|
| 68 |
+
canonical = dict(defaults)
|
| 69 |
+
extras = dict(kwargs)
|
| 70 |
+
for adapter_name, canonical_name in aliases.items():
|
| 71 |
+
if adapter_name in extras:
|
| 72 |
+
canonical[canonical_name] = extras.pop(adapter_name)
|
| 73 |
+
for name in defaults:
|
| 74 |
+
if name in extras:
|
| 75 |
+
canonical[name] = extras.pop(name)
|
| 76 |
+
canonical = SophiaModelConfig(**canonical).to_dict()
|
| 77 |
+
|
| 78 |
+
super().__init__(
|
| 79 |
+
bos_token_id=bos_token_id,
|
| 80 |
+
eos_token_id=eos_token_id,
|
| 81 |
+
pad_token_id=pad_token_id,
|
| 82 |
+
unk_token_id=unk_token_id,
|
| 83 |
+
tie_word_embeddings=bool(tie_word_embeddings),
|
| 84 |
+
return_logits_in_train=bool(return_logits_in_train),
|
| 85 |
+
use_cache=bool(use_cache),
|
| 86 |
+
loss_chunk_size=int(loss_chunk_size),
|
| 87 |
+
gradient_checkpointing_exclude_first=int(
|
| 88 |
+
gradient_checkpointing_exclude_first
|
| 89 |
+
),
|
| 90 |
+
gradient_checkpointing_exclude_last=int(
|
| 91 |
+
gradient_checkpointing_exclude_last
|
| 92 |
+
),
|
| 93 |
+
**extras,
|
| 94 |
+
)
|
| 95 |
+
for name, value in canonical.items():
|
| 96 |
+
setattr(self, name, value)
|
| 97 |
+
self.hidden_size = int(self.dim)
|
| 98 |
+
self.num_hidden_layers = int(self.n_layers)
|
| 99 |
+
self.num_attention_heads = int(self.num_heads)
|
| 100 |
+
self.max_position_embeddings = int(self.max_seq_len)
|
| 101 |
+
self.rms_norm_eps = float(self.norm_eps)
|
| 102 |
+
self.attention_dropout = float(self.dropout)
|
| 103 |
+
apply_config_metadata(
|
| 104 |
+
self,
|
| 105 |
+
return_logits_in_train=bool(return_logits_in_train),
|
| 106 |
+
use_cache=bool(use_cache),
|
| 107 |
+
loss_chunk_size=int(loss_chunk_size),
|
| 108 |
+
gradient_checkpointing_exclude_first=int(
|
| 109 |
+
gradient_checkpointing_exclude_first
|
| 110 |
+
),
|
| 111 |
+
gradient_checkpointing_exclude_last=int(
|
| 112 |
+
gradient_checkpointing_exclude_last
|
| 113 |
+
),
|
| 114 |
+
)
|
| 115 |
+
|
| 116 |
+
def to_model_args(self, *, runtime_max_seq_len: int | None = None) -> object:
|
| 117 |
+
return build_runtime_model_args(
|
| 118 |
+
self,
|
| 119 |
+
model_args_cls=resolve_runtime_backend().model_args_cls,
|
| 120 |
+
runtime_max_seq_len=runtime_max_seq_len,
|
| 121 |
+
)
|
| 122 |
+
|
| 123 |
+
|
| 124 |
+
__all__ = ["SophiaConfig", "apply_config_metadata", "build_config"]
|
hf_generation.py
ADDED
|
@@ -0,0 +1,113 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
import torch
|
| 8 |
+
|
| 9 |
+
from .cache_decode import RuntimeCacheState
|
| 10 |
+
from .decoder_forward import _DecoderForwardBase
|
| 11 |
+
from .decoder_full import validate_decoder_inputs
|
| 12 |
+
from .hf_cache import (
|
| 13 |
+
SophiaCache,
|
| 14 |
+
cache_state_from_past_key_values,
|
| 15 |
+
)
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
class CausalLMForwardMixin(
|
| 19 |
+
_DecoderForwardBase[SophiaCache | None, SophiaCache, object]
|
| 20 |
+
):
|
| 21 |
+
def forward(
|
| 22 |
+
self,
|
| 23 |
+
input_ids: torch.Tensor | None = None,
|
| 24 |
+
attention_mask: torch.Tensor | None = None,
|
| 25 |
+
labels: torch.Tensor | None = None,
|
| 26 |
+
past_key_values: SophiaCache | None = None,
|
| 27 |
+
use_cache: bool | None = None,
|
| 28 |
+
return_dict: bool | None = None,
|
| 29 |
+
logits_to_keep: int | None = None,
|
| 30 |
+
start_pos: int | None = None,
|
| 31 |
+
compute_loss: bool = False,
|
| 32 |
+
**_: object,
|
| 33 |
+
) -> object:
|
| 34 |
+
input_ids = validate_decoder_inputs(
|
| 35 |
+
input_ids=input_ids,
|
| 36 |
+
labels=labels,
|
| 37 |
+
compute_loss=bool(compute_loss),
|
| 38 |
+
)
|
| 39 |
+
resolved_use_cache = bool(
|
| 40 |
+
self.config.use_cache if use_cache is None else use_cache
|
| 41 |
+
)
|
| 42 |
+
resolved_return_dict = self._resolve_runtime_return_dict(
|
| 43 |
+
return_dict=return_dict
|
| 44 |
+
)
|
| 45 |
+
if self._requires_full_path(
|
| 46 |
+
training=bool(self.training),
|
| 47 |
+
labels=labels,
|
| 48 |
+
use_cache=resolved_use_cache,
|
| 49 |
+
):
|
| 50 |
+
return self._forward_full_decoder(
|
| 51 |
+
input_ids=input_ids,
|
| 52 |
+
attention_mask=attention_mask,
|
| 53 |
+
labels=labels,
|
| 54 |
+
compute_loss=bool(compute_loss),
|
| 55 |
+
return_dict=bool(resolved_return_dict),
|
| 56 |
+
)
|
| 57 |
+
return self._forward_cached_decoder(
|
| 58 |
+
input_ids=input_ids,
|
| 59 |
+
attention_mask=attention_mask,
|
| 60 |
+
cache=past_key_values,
|
| 61 |
+
start_pos=start_pos,
|
| 62 |
+
logits_to_keep=logits_to_keep,
|
| 63 |
+
return_dict=bool(resolved_return_dict),
|
| 64 |
+
)
|
| 65 |
+
|
| 66 |
+
|
| 67 |
+
class GenerationCacheMixin:
|
| 68 |
+
def prepare_inputs_for_generation(
|
| 69 |
+
self,
|
| 70 |
+
input_ids: torch.LongTensor,
|
| 71 |
+
past_key_values: SophiaCache | None = None,
|
| 72 |
+
attention_mask: torch.LongTensor | None = None,
|
| 73 |
+
inputs_embeds: torch.FloatTensor | None = None,
|
| 74 |
+
cache_position: torch.LongTensor | None = None,
|
| 75 |
+
**kwargs: object,
|
| 76 |
+
) -> dict[str, object]:
|
| 77 |
+
model_inputs = super().prepare_inputs_for_generation(
|
| 78 |
+
input_ids=input_ids,
|
| 79 |
+
past_key_values=past_key_values,
|
| 80 |
+
attention_mask=attention_mask,
|
| 81 |
+
inputs_embeds=inputs_embeds,
|
| 82 |
+
cache_position=cache_position,
|
| 83 |
+
**kwargs,
|
| 84 |
+
)
|
| 85 |
+
model_inputs.pop("cache_position", None)
|
| 86 |
+
return model_inputs
|
| 87 |
+
|
| 88 |
+
@staticmethod
|
| 89 |
+
def _reorder_cache(
|
| 90 |
+
past_key_values: SophiaCache | None,
|
| 91 |
+
beam_idx: torch.Tensor,
|
| 92 |
+
) -> SophiaCache:
|
| 93 |
+
del past_key_values, beam_idx
|
| 94 |
+
raise NotImplementedError(
|
| 95 |
+
"Sophia HF generate does not support beam cache reordering"
|
| 96 |
+
)
|
| 97 |
+
|
| 98 |
+
@staticmethod
|
| 99 |
+
def _cache_state_from_cache(
|
| 100 |
+
cache: SophiaCache | None,
|
| 101 |
+
) -> RuntimeCacheState | None:
|
| 102 |
+
return cache_state_from_past_key_values(cache)
|
| 103 |
+
|
| 104 |
+
def _cache_output_from_state(self, cache_state: RuntimeCacheState) -> SophiaCache:
|
| 105 |
+
return SophiaCache(
|
| 106 |
+
owner=self,
|
| 107 |
+
cache=cache_state.cache,
|
| 108 |
+
cache_pos=int(cache_state.cache_pos),
|
| 109 |
+
batch_size=int(cache_state.batch_size),
|
| 110 |
+
)
|
| 111 |
+
|
| 112 |
+
|
| 113 |
+
__all__ = ["CausalLMForwardMixin", "GenerationCacheMixin"]
|
hf_lifecycle.py
ADDED
|
@@ -0,0 +1,34 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
from typing import TypeVar, cast
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
ModelT = TypeVar("ModelT", bound="PreTrainedLifecycleMixin")
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
class PreTrainedLifecycleMixin:
|
| 14 |
+
@classmethod
|
| 15 |
+
def from_pretrained(
|
| 16 |
+
cls: type[ModelT],
|
| 17 |
+
pretrained_model_name_or_path,
|
| 18 |
+
*model_args: object,
|
| 19 |
+
**kwargs: object,
|
| 20 |
+
) -> ModelT:
|
| 21 |
+
model = cast(
|
| 22 |
+
ModelT,
|
| 23 |
+
super().from_pretrained(
|
| 24 |
+
pretrained_model_name_or_path,
|
| 25 |
+
*model_args,
|
| 26 |
+
**kwargs,
|
| 27 |
+
),
|
| 28 |
+
)
|
| 29 |
+
with model._model_runtime_lock:
|
| 30 |
+
model._rebuild_runtime_buffers()
|
| 31 |
+
return model
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
__all__ = ["PreTrainedLifecycleMixin"]
|
hf_projection.py
ADDED
|
@@ -0,0 +1,124 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
from transformers import PretrainedConfig
|
| 8 |
+
|
| 9 |
+
from .hf_support import (
|
| 10 |
+
apply_config_metadata,
|
| 11 |
+
build_config,
|
| 12 |
+
)
|
| 13 |
+
from .canonical_config import SophiaModelConfig
|
| 14 |
+
from .config_projection import build_runtime_model_args
|
| 15 |
+
from .runtime_backend import resolve_runtime_backend
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
_REMOVED_FIELDS = {
|
| 19 |
+
"n_heads",
|
| 20 |
+
"num_key_value_heads",
|
| 21 |
+
"rope_head_dim",
|
| 22 |
+
"rope_theta",
|
| 23 |
+
"original_seq_len",
|
| 24 |
+
"rope_factor",
|
| 25 |
+
"beta_fast",
|
| 26 |
+
"beta_slow",
|
| 27 |
+
"use_qk_norm",
|
| 28 |
+
}
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
class SophiaConfig(PretrainedConfig):
|
| 32 |
+
"""Hugging Face adapter for the native Sophia Hybrid schema."""
|
| 33 |
+
|
| 34 |
+
model_type = "sophia_hybrid"
|
| 35 |
+
keys_to_ignore_at_inference = ["past_key_values"]
|
| 36 |
+
attribute_map = {"intermediate_size": "ffn_hidden"}
|
| 37 |
+
|
| 38 |
+
def __init__(
|
| 39 |
+
self,
|
| 40 |
+
*,
|
| 41 |
+
bos_token_id: int | None = None,
|
| 42 |
+
eos_token_id: int | None = None,
|
| 43 |
+
pad_token_id: int | None = None,
|
| 44 |
+
unk_token_id: int | None = None,
|
| 45 |
+
tie_word_embeddings: bool = True,
|
| 46 |
+
return_logits_in_train: bool = True,
|
| 47 |
+
use_cache: bool = True,
|
| 48 |
+
loss_chunk_size: int = 0,
|
| 49 |
+
gradient_checkpointing_exclude_first: int = 0,
|
| 50 |
+
gradient_checkpointing_exclude_last: int = 0,
|
| 51 |
+
**kwargs: object,
|
| 52 |
+
) -> None:
|
| 53 |
+
removed = sorted(_REMOVED_FIELDS.intersection(kwargs))
|
| 54 |
+
if removed:
|
| 55 |
+
raise ValueError(
|
| 56 |
+
"legacy Sophia Transformer fields are not supported: "
|
| 57 |
+
+ ", ".join(removed)
|
| 58 |
+
)
|
| 59 |
+
defaults = SophiaModelConfig.get_defaults()
|
| 60 |
+
aliases = {
|
| 61 |
+
"hidden_size": "dim",
|
| 62 |
+
"num_hidden_layers": "n_layers",
|
| 63 |
+
"num_attention_heads": "num_heads",
|
| 64 |
+
"rms_norm_eps": "norm_eps",
|
| 65 |
+
"max_position_embeddings": "max_seq_len",
|
| 66 |
+
"attention_dropout": "dropout",
|
| 67 |
+
}
|
| 68 |
+
canonical = dict(defaults)
|
| 69 |
+
extras = dict(kwargs)
|
| 70 |
+
for adapter_name, canonical_name in aliases.items():
|
| 71 |
+
if adapter_name in extras:
|
| 72 |
+
canonical[canonical_name] = extras.pop(adapter_name)
|
| 73 |
+
for name in defaults:
|
| 74 |
+
if name in extras:
|
| 75 |
+
canonical[name] = extras.pop(name)
|
| 76 |
+
canonical = SophiaModelConfig(**canonical).to_dict()
|
| 77 |
+
|
| 78 |
+
super().__init__(
|
| 79 |
+
bos_token_id=bos_token_id,
|
| 80 |
+
eos_token_id=eos_token_id,
|
| 81 |
+
pad_token_id=pad_token_id,
|
| 82 |
+
unk_token_id=unk_token_id,
|
| 83 |
+
tie_word_embeddings=bool(tie_word_embeddings),
|
| 84 |
+
return_logits_in_train=bool(return_logits_in_train),
|
| 85 |
+
use_cache=bool(use_cache),
|
| 86 |
+
loss_chunk_size=int(loss_chunk_size),
|
| 87 |
+
gradient_checkpointing_exclude_first=int(
|
| 88 |
+
gradient_checkpointing_exclude_first
|
| 89 |
+
),
|
| 90 |
+
gradient_checkpointing_exclude_last=int(
|
| 91 |
+
gradient_checkpointing_exclude_last
|
| 92 |
+
),
|
| 93 |
+
**extras,
|
| 94 |
+
)
|
| 95 |
+
for name, value in canonical.items():
|
| 96 |
+
setattr(self, name, value)
|
| 97 |
+
self.hidden_size = int(self.dim)
|
| 98 |
+
self.num_hidden_layers = int(self.n_layers)
|
| 99 |
+
self.num_attention_heads = int(self.num_heads)
|
| 100 |
+
self.max_position_embeddings = int(self.max_seq_len)
|
| 101 |
+
self.rms_norm_eps = float(self.norm_eps)
|
| 102 |
+
self.attention_dropout = float(self.dropout)
|
| 103 |
+
apply_config_metadata(
|
| 104 |
+
self,
|
| 105 |
+
return_logits_in_train=bool(return_logits_in_train),
|
| 106 |
+
use_cache=bool(use_cache),
|
| 107 |
+
loss_chunk_size=int(loss_chunk_size),
|
| 108 |
+
gradient_checkpointing_exclude_first=int(
|
| 109 |
+
gradient_checkpointing_exclude_first
|
| 110 |
+
),
|
| 111 |
+
gradient_checkpointing_exclude_last=int(
|
| 112 |
+
gradient_checkpointing_exclude_last
|
| 113 |
+
),
|
| 114 |
+
)
|
| 115 |
+
|
| 116 |
+
def to_model_args(self, *, runtime_max_seq_len: int | None = None) -> object:
|
| 117 |
+
return build_runtime_model_args(
|
| 118 |
+
self,
|
| 119 |
+
model_args_cls=resolve_runtime_backend().model_args_cls,
|
| 120 |
+
runtime_max_seq_len=runtime_max_seq_len,
|
| 121 |
+
)
|
| 122 |
+
|
| 123 |
+
|
| 124 |
+
__all__ = ["SophiaConfig", "apply_config_metadata", "build_config"]
|
hf_remote_code.py
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
"""Remote-code naming shared by the HF adapter and export bundle."""
|
| 6 |
+
|
| 7 |
+
from __future__ import annotations
|
| 8 |
+
|
| 9 |
+
HF_MODELING_MODULE = "modeling_sophia"
|
| 10 |
+
HF_CONFIG_CLASS = "SophiaConfig"
|
| 11 |
+
HF_CAUSAL_LM_CLASS = "SophiaForCausalLM"
|
| 12 |
+
HF_SUPPORT_FILENAME = "hf_support.py"
|
| 13 |
+
HF_RUNTIME_FILENAME = "sophia_runtime.py"
|
| 14 |
+
|
| 15 |
+
__all__ = [
|
| 16 |
+
"HF_CAUSAL_LM_CLASS",
|
| 17 |
+
"HF_CONFIG_CLASS",
|
| 18 |
+
"HF_MODELING_MODULE",
|
| 19 |
+
"HF_RUNTIME_FILENAME",
|
| 20 |
+
"HF_SUPPORT_FILENAME",
|
| 21 |
+
]
|
hf_support.py
ADDED
|
@@ -0,0 +1,151 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
"""HF adapter support for config metadata and remote-code registration."""
|
| 6 |
+
|
| 7 |
+
from __future__ import annotations
|
| 8 |
+
|
| 9 |
+
from collections.abc import Callable, Mapping
|
| 10 |
+
from dataclasses import asdict, is_dataclass
|
| 11 |
+
from typing import Protocol
|
| 12 |
+
|
| 13 |
+
from .canonical_config import SophiaModelConfig
|
| 14 |
+
from .hf_remote_code import (
|
| 15 |
+
HF_CAUSAL_LM_CLASS,
|
| 16 |
+
HF_CONFIG_CLASS,
|
| 17 |
+
HF_MODELING_MODULE,
|
| 18 |
+
HF_RUNTIME_FILENAME,
|
| 19 |
+
HF_SUPPORT_FILENAME,
|
| 20 |
+
)
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
class AutoMapConfig(Protocol):
|
| 24 |
+
architectures: list[str]
|
| 25 |
+
auto_map: dict[str, str]
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
class ConfiguredModel(Protocol):
|
| 29 |
+
config: AutoMapConfig | None
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
class MetadataConfig(AutoMapConfig, Protocol):
|
| 33 |
+
max_seq_len: int
|
| 34 |
+
dim: int
|
| 35 |
+
n_layers: int
|
| 36 |
+
num_heads: int
|
| 37 |
+
norm_eps: float
|
| 38 |
+
dropout: float
|
| 39 |
+
head_dim: int
|
| 40 |
+
max_position_embeddings: int
|
| 41 |
+
hidden_size: int
|
| 42 |
+
num_hidden_layers: int
|
| 43 |
+
num_attention_heads: int
|
| 44 |
+
sliding_window: int
|
| 45 |
+
rms_norm_eps: float
|
| 46 |
+
attention_dropout: float
|
| 47 |
+
loss_chunk_size: int
|
| 48 |
+
gradient_checkpointing_exclude_first: int
|
| 49 |
+
gradient_checkpointing_exclude_last: int
|
| 50 |
+
return_logits_in_train: bool
|
| 51 |
+
use_cache: bool
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
def _to_dict_values(config: object) -> dict[str, object] | None:
|
| 57 |
+
to_dict = getattr(config, "to_dict", None)
|
| 58 |
+
if callable(to_dict):
|
| 59 |
+
values = to_dict()
|
| 60 |
+
if isinstance(values, Mapping):
|
| 61 |
+
return dict(values)
|
| 62 |
+
return None
|
| 63 |
+
|
| 64 |
+
|
| 65 |
+
def apply_auto_map(config: AutoMapConfig) -> AutoMapConfig:
|
| 66 |
+
config.architectures = [HF_CAUSAL_LM_CLASS]
|
| 67 |
+
config.auto_map = {
|
| 68 |
+
"AutoConfig": f"{HF_MODELING_MODULE}.{HF_CONFIG_CLASS}",
|
| 69 |
+
"AutoModelForCausalLM": f"{HF_MODELING_MODULE}.{HF_CAUSAL_LM_CLASS}",
|
| 70 |
+
}
|
| 71 |
+
return config
|
| 72 |
+
|
| 73 |
+
|
| 74 |
+
def ensure_auto_map(model: ConfiguredModel) -> None:
|
| 75 |
+
cfg = model.config
|
| 76 |
+
if cfg is None:
|
| 77 |
+
raise RuntimeError("Model has no config; unable to set HuggingFace auto_map.")
|
| 78 |
+
apply_auto_map(cfg)
|
| 79 |
+
|
| 80 |
+
|
| 81 |
+
def apply_config_metadata(
|
| 82 |
+
config: MetadataConfig,
|
| 83 |
+
*,
|
| 84 |
+
return_logits_in_train: bool,
|
| 85 |
+
use_cache: bool,
|
| 86 |
+
loss_chunk_size: int,
|
| 87 |
+
gradient_checkpointing_exclude_first: int,
|
| 88 |
+
gradient_checkpointing_exclude_last: int,
|
| 89 |
+
) -> None:
|
| 90 |
+
config.max_position_embeddings = int(config.max_seq_len)
|
| 91 |
+
config.hidden_size = int(config.dim)
|
| 92 |
+
config.num_hidden_layers = int(config.n_layers)
|
| 93 |
+
config.num_attention_heads = int(config.num_heads)
|
| 94 |
+
config.sliding_window = int(config.max_seq_len)
|
| 95 |
+
config.rms_norm_eps = float(config.norm_eps)
|
| 96 |
+
config.attention_dropout = float(config.dropout)
|
| 97 |
+
config.loss_chunk_size = max(int(loss_chunk_size), 0)
|
| 98 |
+
config.gradient_checkpointing_exclude_first = max(
|
| 99 |
+
int(gradient_checkpointing_exclude_first),
|
| 100 |
+
0,
|
| 101 |
+
)
|
| 102 |
+
config.gradient_checkpointing_exclude_last = max(
|
| 103 |
+
int(gradient_checkpointing_exclude_last),
|
| 104 |
+
0,
|
| 105 |
+
)
|
| 106 |
+
config.return_logits_in_train = bool(return_logits_in_train)
|
| 107 |
+
config.use_cache = bool(use_cache)
|
| 108 |
+
config.architectures = [HF_CAUSAL_LM_CLASS]
|
| 109 |
+
|
| 110 |
+
|
| 111 |
+
def build_config[ConfigT](
|
| 112 |
+
config: object,
|
| 113 |
+
*,
|
| 114 |
+
config_cls: Callable[..., ConfigT] | None = None,
|
| 115 |
+
) -> ConfigT:
|
| 116 |
+
if config_cls is None:
|
| 117 |
+
from .hf_projection import SophiaConfig
|
| 118 |
+
|
| 119 |
+
config_cls = SophiaConfig
|
| 120 |
+
|
| 121 |
+
if isinstance(config, Mapping):
|
| 122 |
+
values = dict(config)
|
| 123 |
+
elif is_dataclass(config):
|
| 124 |
+
values = asdict(config)
|
| 125 |
+
else:
|
| 126 |
+
values = _to_dict_values(config)
|
| 127 |
+
if values is None:
|
| 128 |
+
canonical_fields = set(SophiaModelConfig.get_defaults())
|
| 129 |
+
values = {
|
| 130 |
+
key: getattr(config, key)
|
| 131 |
+
for key in canonical_fields
|
| 132 |
+
if hasattr(config, key)
|
| 133 |
+
}
|
| 134 |
+
|
| 135 |
+
return config_cls(**values)
|
| 136 |
+
|
| 137 |
+
|
| 138 |
+
__all__ = [
|
| 139 |
+
"HF_CAUSAL_LM_CLASS",
|
| 140 |
+
"HF_CONFIG_CLASS",
|
| 141 |
+
"HF_MODELING_MODULE",
|
| 142 |
+
"HF_RUNTIME_FILENAME",
|
| 143 |
+
"HF_SUPPORT_FILENAME",
|
| 144 |
+
"AutoMapConfig",
|
| 145 |
+
"ConfiguredModel",
|
| 146 |
+
"MetadataConfig",
|
| 147 |
+
"apply_auto_map",
|
| 148 |
+
"apply_config_metadata",
|
| 149 |
+
"build_config",
|
| 150 |
+
"ensure_auto_map",
|
| 151 |
+
]
|
input_mask.py
ADDED
|
@@ -0,0 +1,66 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
import torch
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
def _as_bool_mask(attention_mask: torch.Tensor) -> torch.Tensor:
|
| 11 |
+
return (
|
| 12 |
+
attention_mask if attention_mask.dtype == torch.bool else (attention_mask != 0)
|
| 13 |
+
)
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
def is_all_ones_mask(attention_mask: torch.Tensor) -> bool:
|
| 17 |
+
mask = _as_bool_mask(attention_mask)
|
| 18 |
+
return bool(mask.all().item())
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
def validate_right_padding_mask(
|
| 22 |
+
attention_mask: torch.Tensor,
|
| 23 |
+
*,
|
| 24 |
+
input_ids: torch.Tensor,
|
| 25 |
+
) -> None:
|
| 26 |
+
if attention_mask.dim() != 2:
|
| 27 |
+
raise ValueError("attention_mask must be [B,T]")
|
| 28 |
+
if attention_mask.shape != input_ids.shape:
|
| 29 |
+
raise ValueError("attention_mask shape must match input_ids")
|
| 30 |
+
mask = _as_bool_mask(attention_mask)
|
| 31 |
+
valid_rows = mask.any(dim=1)
|
| 32 |
+
resumes_after_padding = (~mask[:, :-1] & mask[:, 1:]).any(dim=1)
|
| 33 |
+
if not bool(valid_rows.all().item()):
|
| 34 |
+
raise ValueError("attention_mask row has no valid tokens")
|
| 35 |
+
if bool(resumes_after_padding.any().item()):
|
| 36 |
+
raise ValueError("loss computation requires right-padded attention masks")
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
def slice_valid_tokens(
|
| 40 |
+
input_ids: torch.Tensor,
|
| 41 |
+
attention_mask: torch.Tensor,
|
| 42 |
+
) -> list[tuple[int, int, torch.Tensor]]:
|
| 43 |
+
if attention_mask.dim() != 2:
|
| 44 |
+
raise ValueError("attention_mask must be [B,T]")
|
| 45 |
+
mask = _as_bool_mask(attention_mask)
|
| 46 |
+
if mask.shape != input_ids.shape:
|
| 47 |
+
raise ValueError("attention_mask shape must match input_ids")
|
| 48 |
+
|
| 49 |
+
rows: list[tuple[int, int, torch.Tensor]] = []
|
| 50 |
+
for batch_index in range(int(input_ids.size(0))):
|
| 51 |
+
idx = mask[batch_index].nonzero(as_tuple=False).squeeze(-1)
|
| 52 |
+
if int(idx.numel()) == 0:
|
| 53 |
+
raise ValueError("attention_mask row has no valid tokens")
|
| 54 |
+
start = int(idx[0].item())
|
| 55 |
+
end = int(idx[-1].item()) + 1
|
| 56 |
+
if int(idx.numel()) != (end - start):
|
| 57 |
+
raise ValueError("Sophia only supports contiguous padding masks")
|
| 58 |
+
rows.append((start, end, input_ids[batch_index : batch_index + 1, start:end]))
|
| 59 |
+
return rows
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
__all__ = [
|
| 63 |
+
"is_all_ones_mask",
|
| 64 |
+
"slice_valid_tokens",
|
| 65 |
+
"validate_right_padding_mask",
|
| 66 |
+
]
|
lineage.json
ADDED
|
@@ -0,0 +1,58 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"deliverable": "sophia.pt",
|
| 3 |
+
"sha256": "e49ddc2cbbd311330f21fe52ada853e3ed745cf5f2e47de83da3898234592268",
|
| 4 |
+
"equals": "runs/dpo_e3",
|
| 5 |
+
"chain": [
|
| 6 |
+
{
|
| 7 |
+
"stage": "pretrain",
|
| 8 |
+
"run": "runs/pretrain_e3bc2d5_bf16_qkclip",
|
| 9 |
+
"ckpt": "checkpoints/ckpt_step15259.pt",
|
| 10 |
+
"steps": 15259,
|
| 11 |
+
"tokens": "20B",
|
| 12 |
+
"seed": 42
|
| 13 |
+
},
|
| 14 |
+
{
|
| 15 |
+
"stage": "sft",
|
| 16 |
+
"run": "runs/sft",
|
| 17 |
+
"ckpt": "ckpt_final.pt",
|
| 18 |
+
"steps": 7176,
|
| 19 |
+
"data": "release/corpus (88,267 rows x 3ep)",
|
| 20 |
+
"val_loss": 2.6183438087448114,
|
| 21 |
+
"selected": "epoch3 via haiku rank (mean rank 1.88, repeat 9.7%)"
|
| 22 |
+
},
|
| 23 |
+
{
|
| 24 |
+
"stage": "rft_experiment",
|
| 25 |
+
"run": "runs/rft_e3.pt",
|
| 26 |
+
"steps": 104,
|
| 27 |
+
"verdict": "null result on powered probe; not shipped"
|
| 28 |
+
},
|
| 29 |
+
{
|
| 30 |
+
"stage": "dpo",
|
| 31 |
+
"run": "runs/dpo_e3",
|
| 32 |
+
"steps": 500,
|
| 33 |
+
"pairs": 1953,
|
| 34 |
+
"beta": 0.1,
|
| 35 |
+
"lr": 1e-06,
|
| 36 |
+
"verdict": "WINNER: probe paired +2.10 (t=+3.71), pass70 z=+3.62, pack +1.16, all flags down, no diversity loss"
|
| 37 |
+
}
|
| 38 |
+
],
|
| 39 |
+
"eval": {
|
| 40 |
+
"judge": "claude-haiku-4-5 via frostfox",
|
| 41 |
+
"probe_n16_t04": {
|
| 42 |
+
"e3": 49.16,
|
| 43 |
+
"dpo": 51.37,
|
| 44 |
+
"paired_diff": 2.1,
|
| 45 |
+
"t": 3.71
|
| 46 |
+
},
|
| 47 |
+
"pack_30x6_t07": {
|
| 48 |
+
"e3": 27.93,
|
| 49 |
+
"dpo": 29.08,
|
| 50 |
+
"paired_diff": 1.16,
|
| 51 |
+
"t": 0.94
|
| 52 |
+
}
|
| 53 |
+
},
|
| 54 |
+
"external_upload": false,
|
| 55 |
+
"release_name": "sophia",
|
| 56 |
+
"release_dir": "release/sophia",
|
| 57 |
+
"version": "1.0.0"
|
| 58 |
+
}
|
loss_stats.py
ADDED
|
@@ -0,0 +1,72 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
import torch
|
| 8 |
+
import torch.nn.functional as functional
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
def supervised_token_count(
|
| 12 |
+
labels: torch.Tensor | None,
|
| 13 |
+
*,
|
| 14 |
+
label_offset: int = 0,
|
| 15 |
+
ignore_index: int = -100,
|
| 16 |
+
) -> torch.Tensor:
|
| 17 |
+
if labels is None or not torch.is_tensor(labels):
|
| 18 |
+
return torch.zeros((), dtype=torch.int64)
|
| 19 |
+
if labels.ndim != 2:
|
| 20 |
+
raise ValueError(
|
| 21 |
+
f"labels must be 2D [B,S] to count supervised tokens (got shape={tuple(labels.shape)})"
|
| 22 |
+
)
|
| 23 |
+
target_start = int(label_offset) + 1
|
| 24 |
+
if int(labels.size(1)) <= target_start:
|
| 25 |
+
return torch.zeros((), device=labels.device, dtype=torch.int64)
|
| 26 |
+
shifted = labels[:, target_start:]
|
| 27 |
+
return (shifted != int(ignore_index)).sum(dtype=torch.int64)
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
def shifted_loss_sum_and_count(
|
| 31 |
+
logits: torch.Tensor,
|
| 32 |
+
labels: torch.Tensor,
|
| 33 |
+
*,
|
| 34 |
+
label_offset: int = 0,
|
| 35 |
+
ignore_index: int = -100,
|
| 36 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 37 |
+
if int(logits.size(1)) <= 1 or int(labels.size(1)) <= int(label_offset) + 1:
|
| 38 |
+
zero = logits.new_zeros(())
|
| 39 |
+
return zero, zero
|
| 40 |
+
shift_logits = logits[:, :-1, :].contiguous()
|
| 41 |
+
target_start = int(label_offset) + 1
|
| 42 |
+
target_end = target_start + int(shift_logits.size(1))
|
| 43 |
+
if int(labels.size(1)) < target_end:
|
| 44 |
+
shift_logits = shift_logits[:, : max(int(labels.size(1)) - target_start, 0), :]
|
| 45 |
+
target_end = target_start + int(shift_logits.size(1))
|
| 46 |
+
if int(shift_logits.size(1)) <= 0:
|
| 47 |
+
zero = logits.new_zeros(())
|
| 48 |
+
return zero, zero
|
| 49 |
+
shift_labels = labels[:, target_start:target_end].contiguous()
|
| 50 |
+
flat_labels = shift_labels.reshape(-1).to(dtype=torch.long)
|
| 51 |
+
flat_logits = shift_logits.reshape(-1, int(shift_logits.size(-1)))
|
| 52 |
+
loss_sum = functional.cross_entropy(
|
| 53 |
+
flat_logits,
|
| 54 |
+
flat_labels,
|
| 55 |
+
ignore_index=int(ignore_index),
|
| 56 |
+
reduction="sum",
|
| 57 |
+
)
|
| 58 |
+
count = (flat_labels != int(ignore_index)).sum().to(dtype=loss_sum.dtype)
|
| 59 |
+
return loss_sum, count
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
def mean_loss_from_sum_and_count(
|
| 63 |
+
*,
|
| 64 |
+
loss_sum: torch.Tensor,
|
| 65 |
+
count: torch.Tensor,
|
| 66 |
+
reference: torch.Tensor,
|
| 67 |
+
) -> torch.Tensor:
|
| 68 |
+
return torch.where(
|
| 69 |
+
count > 0,
|
| 70 |
+
loss_sum / count.clamp_min(1.0),
|
| 71 |
+
reference.new_zeros(()),
|
| 72 |
+
)
|
model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:7a55aa339222f8b0e7a168774264e2dc65ced7fae6f43bcef3f9811c8fd456df
|
| 3 |
+
size 2226652224
|
model_attention.py
ADDED
|
@@ -0,0 +1,694 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
from importlib import import_module
|
| 8 |
+
import math
|
| 9 |
+
|
| 10 |
+
import torch
|
| 11 |
+
import torch.nn.functional as functional
|
| 12 |
+
from torch import nn
|
| 13 |
+
from torch.nn.attention.bias import CausalBias, causal_lower_right
|
| 14 |
+
|
| 15 |
+
from .model_config import ModelArgs
|
| 16 |
+
from .model_ops import RMSNorm
|
| 17 |
+
from .runtime_linear import RuntimeLinear
|
| 18 |
+
from .model_state import LayerCacheSnapshot
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
PAD_LOGIT = -1.0e4
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
def _inverse_softplus(value: torch.Tensor) -> torch.Tensor:
|
| 25 |
+
return value + torch.log(-torch.expm1(-value))
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
class CausalDepthwiseConv1d(nn.Module):
|
| 29 |
+
def __init__(self, channels: int, kernel_size: int) -> None:
|
| 30 |
+
super().__init__()
|
| 31 |
+
self.channels = int(channels)
|
| 32 |
+
self.kernel_size = int(kernel_size)
|
| 33 |
+
self.weight = nn.Parameter(torch.empty(self.channels, self.kernel_size))
|
| 34 |
+
nn.init.normal_(self.weight, mean=0.0, std=0.02)
|
| 35 |
+
|
| 36 |
+
def forward(
|
| 37 |
+
self,
|
| 38 |
+
x: torch.Tensor,
|
| 39 |
+
*,
|
| 40 |
+
history: torch.Tensor | None = None,
|
| 41 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 42 |
+
# x: [B, T, C], history: [B, C, K-1]
|
| 43 |
+
x_channels = x.transpose(1, 2)
|
| 44 |
+
history_length = self.kernel_size - 1
|
| 45 |
+
if history is None and history_length > 0:
|
| 46 |
+
history = x_channels.new_zeros(
|
| 47 |
+
int(x.size(0)), self.channels, history_length
|
| 48 |
+
)
|
| 49 |
+
sequence = (
|
| 50 |
+
x_channels
|
| 51 |
+
if history_length == 0
|
| 52 |
+
else torch.cat((history.to(dtype=x.dtype), x_channels), dim=-1)
|
| 53 |
+
)
|
| 54 |
+
output = functional.conv1d(
|
| 55 |
+
sequence,
|
| 56 |
+
self.weight.to(dtype=x.dtype).unsqueeze(1),
|
| 57 |
+
groups=self.channels,
|
| 58 |
+
).transpose(1, 2)
|
| 59 |
+
next_history = (
|
| 60 |
+
sequence[..., :0]
|
| 61 |
+
if history_length == 0
|
| 62 |
+
else sequence[..., -history_length:]
|
| 63 |
+
)
|
| 64 |
+
return functional.silu(output), next_history
|
| 65 |
+
|
| 66 |
+
|
| 67 |
+
class SophiaKDA(nn.Module):
|
| 68 |
+
"""Kimi Delta Attention adapted as Sophia's recurrent token mixer."""
|
| 69 |
+
|
| 70 |
+
def __init__(self, args: ModelArgs) -> None:
|
| 71 |
+
super().__init__()
|
| 72 |
+
self.dim = int(args.dim)
|
| 73 |
+
self.num_heads = int(args.num_heads)
|
| 74 |
+
self.head_dim = int(args.head_dim)
|
| 75 |
+
self.inner_dim = self.num_heads * self.head_dim
|
| 76 |
+
self.decay_rank = int(args.kda_decay_rank)
|
| 77 |
+
self.output_gate_rank = int(args.kda_output_gate_rank)
|
| 78 |
+
self.output_gate_full_rank = bool(args.kda_output_gate_full_rank)
|
| 79 |
+
self.lower_bound = float(args.kda_decay_lower_bound)
|
| 80 |
+
self.dt_min = float(args.kda_dt_min)
|
| 81 |
+
self.dt_max = float(args.kda_dt_max)
|
| 82 |
+
self.dt_floor = float(args.kda_dt_floor)
|
| 83 |
+
self.backend = str(args.kda_backend)
|
| 84 |
+
self.use_cache = bool(args.use_cache)
|
| 85 |
+
self.conv_kernel = int(args.short_conv_kernel)
|
| 86 |
+
|
| 87 |
+
self.q_proj = RuntimeLinear(self.dim, self.inner_dim, bias=False)
|
| 88 |
+
self.k_proj = RuntimeLinear(self.dim, self.inner_dim, bias=False)
|
| 89 |
+
self.v_proj = RuntimeLinear(self.dim, self.inner_dim, bias=False)
|
| 90 |
+
|
| 91 |
+
self.q_conv = CausalDepthwiseConv1d(self.inner_dim, self.conv_kernel)
|
| 92 |
+
self.k_conv = CausalDepthwiseConv1d(self.inner_dim, self.conv_kernel)
|
| 93 |
+
self.v_conv = CausalDepthwiseConv1d(self.inner_dim, self.conv_kernel)
|
| 94 |
+
|
| 95 |
+
self.decay_down = RuntimeLinear(self.dim, self.decay_rank, bias=False)
|
| 96 |
+
self.decay_up = RuntimeLinear(self.decay_rank, self.inner_dim, bias=False)
|
| 97 |
+
self.beta_proj = RuntimeLinear(self.dim, self.num_heads, bias=False)
|
| 98 |
+
self.A_log = nn.Parameter(
|
| 99 |
+
torch.full(
|
| 100 |
+
(self.num_heads,),
|
| 101 |
+
float(args.kda_a_log_init),
|
| 102 |
+
dtype=torch.float32,
|
| 103 |
+
)
|
| 104 |
+
)
|
| 105 |
+
dt = torch.exp(
|
| 106 |
+
torch.rand(self.inner_dim, dtype=torch.float32)
|
| 107 |
+
* (math.log(self.dt_max) - math.log(self.dt_min))
|
| 108 |
+
+ math.log(self.dt_min)
|
| 109 |
+
).clamp_min(self.dt_floor)
|
| 110 |
+
self.dt_bias = nn.Parameter(_inverse_softplus(dt))
|
| 111 |
+
self.A_log._no_weight_decay = True # type: ignore[attr-defined]
|
| 112 |
+
self.dt_bias._no_weight_decay = True # type: ignore[attr-defined]
|
| 113 |
+
|
| 114 |
+
if self.output_gate_full_rank:
|
| 115 |
+
self.output_gate = RuntimeLinear(self.dim, self.inner_dim, bias=False)
|
| 116 |
+
else:
|
| 117 |
+
self.output_gate_down = RuntimeLinear(
|
| 118 |
+
self.dim, self.output_gate_rank, bias=False
|
| 119 |
+
)
|
| 120 |
+
self.output_gate_up = RuntimeLinear(
|
| 121 |
+
self.output_gate_rank, self.inner_dim, bias=False
|
| 122 |
+
)
|
| 123 |
+
self.output_norm = RMSNorm(self.head_dim, args.norm_eps)
|
| 124 |
+
self.o_proj = RuntimeLinear(self.inner_dim, self.dim, bias=False)
|
| 125 |
+
|
| 126 |
+
self.register_buffer(
|
| 127 |
+
"recurrent_state",
|
| 128 |
+
torch.zeros(
|
| 129 |
+
int(args.max_batch_size),
|
| 130 |
+
self.num_heads,
|
| 131 |
+
self.head_dim,
|
| 132 |
+
self.head_dim,
|
| 133 |
+
dtype=torch.float32,
|
| 134 |
+
),
|
| 135 |
+
persistent=False,
|
| 136 |
+
)
|
| 137 |
+
self.register_buffer(
|
| 138 |
+
"conv_state",
|
| 139 |
+
torch.zeros(
|
| 140 |
+
int(args.max_batch_size),
|
| 141 |
+
3,
|
| 142 |
+
self.inner_dim,
|
| 143 |
+
self.conv_kernel - 1,
|
| 144 |
+
dtype=torch.float32,
|
| 145 |
+
),
|
| 146 |
+
persistent=False,
|
| 147 |
+
)
|
| 148 |
+
|
| 149 |
+
def ensure_batch_capacity(self, batch_size: int) -> None:
|
| 150 |
+
required = int(batch_size)
|
| 151 |
+
if required <= int(self.recurrent_state.size(0)):
|
| 152 |
+
return
|
| 153 |
+
recurrent = self.recurrent_state.new_zeros(
|
| 154 |
+
required, self.num_heads, self.head_dim, self.head_dim
|
| 155 |
+
)
|
| 156 |
+
recurrent[: self.recurrent_state.size(0)].copy_(self.recurrent_state)
|
| 157 |
+
conv = self.conv_state.new_zeros(
|
| 158 |
+
required, 3, self.inner_dim, self.conv_kernel - 1
|
| 159 |
+
)
|
| 160 |
+
conv[: self.conv_state.size(0)].copy_(self.conv_state)
|
| 161 |
+
self.recurrent_state = recurrent
|
| 162 |
+
self.conv_state = conv
|
| 163 |
+
|
| 164 |
+
def ensure_sequence_capacity(self, max_seq_len: int) -> None:
|
| 165 |
+
del max_seq_len
|
| 166 |
+
|
| 167 |
+
def rebuild_runtime_buffers(self, max_seq_len: int) -> None:
|
| 168 |
+
del max_seq_len
|
| 169 |
+
device = self.q_proj.weight.device
|
| 170 |
+
self.recurrent_state = torch.zeros(
|
| 171 |
+
int(self.recurrent_state.size(0)),
|
| 172 |
+
self.num_heads,
|
| 173 |
+
self.head_dim,
|
| 174 |
+
self.head_dim,
|
| 175 |
+
device=device,
|
| 176 |
+
dtype=torch.float32,
|
| 177 |
+
)
|
| 178 |
+
self.conv_state = torch.zeros(
|
| 179 |
+
int(self.conv_state.size(0)),
|
| 180 |
+
3,
|
| 181 |
+
self.inner_dim,
|
| 182 |
+
self.conv_kernel - 1,
|
| 183 |
+
device=device,
|
| 184 |
+
dtype=torch.float32,
|
| 185 |
+
)
|
| 186 |
+
|
| 187 |
+
def reset(self) -> None:
|
| 188 |
+
if self.recurrent_state.is_inference() or self.conv_state.is_inference():
|
| 189 |
+
self.recurrent_state = torch.zeros_like(self.recurrent_state)
|
| 190 |
+
self.conv_state = torch.zeros_like(self.conv_state)
|
| 191 |
+
else:
|
| 192 |
+
self.recurrent_state.zero_()
|
| 193 |
+
self.conv_state.zero_()
|
| 194 |
+
|
| 195 |
+
def refresh_state_buffers(self) -> None:
|
| 196 |
+
self.rebuild_runtime_buffers(0)
|
| 197 |
+
|
| 198 |
+
def cache_snapshot(
|
| 199 |
+
self, *, device: str, batch_size: int, cache_pos: int | None
|
| 200 |
+
) -> LayerCacheSnapshot:
|
| 201 |
+
del cache_pos
|
| 202 |
+
return LayerCacheSnapshot(
|
| 203 |
+
recurrent=self.recurrent_state[:batch_size].to(device).clone(),
|
| 204 |
+
conv=self.conv_state[:batch_size].to(device).clone(),
|
| 205 |
+
)
|
| 206 |
+
|
| 207 |
+
def validate_cache_snapshot(self, snapshot: LayerCacheSnapshot) -> None:
|
| 208 |
+
if snapshot.latent is not None:
|
| 209 |
+
raise ValueError("KDA cache cannot contain MLA latent state")
|
| 210 |
+
expected_recurrent = tuple(self.recurrent_state.shape[1:])
|
| 211 |
+
expected_conv = tuple(self.conv_state.shape[1:])
|
| 212 |
+
if snapshot.recurrent is not None and tuple(snapshot.recurrent.shape[1:]) != expected_recurrent:
|
| 213 |
+
raise ValueError("KDA recurrent cache shape mismatch")
|
| 214 |
+
if snapshot.conv is not None and tuple(snapshot.conv.shape[1:]) != expected_conv:
|
| 215 |
+
raise ValueError("KDA convolution cache shape mismatch")
|
| 216 |
+
|
| 217 |
+
def load_cache_snapshot(self, snapshot: LayerCacheSnapshot) -> None:
|
| 218 |
+
self.validate_cache_snapshot(snapshot)
|
| 219 |
+
required = snapshot.batch_size()
|
| 220 |
+
self.ensure_batch_capacity(required)
|
| 221 |
+
if snapshot.recurrent is not None:
|
| 222 |
+
self.recurrent_state[:required].copy_(
|
| 223 |
+
snapshot.recurrent.to(self.recurrent_state.device, dtype=torch.float32)
|
| 224 |
+
)
|
| 225 |
+
if snapshot.conv is not None:
|
| 226 |
+
self.conv_state[:required].copy_(
|
| 227 |
+
snapshot.conv.to(self.conv_state.device, dtype=torch.float32)
|
| 228 |
+
)
|
| 229 |
+
|
| 230 |
+
def _short_convolution(
|
| 231 |
+
self,
|
| 232 |
+
q: torch.Tensor,
|
| 233 |
+
k: torch.Tensor,
|
| 234 |
+
v: torch.Tensor,
|
| 235 |
+
*,
|
| 236 |
+
start_pos: int,
|
| 237 |
+
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
| 238 |
+
bsz = int(q.size(0))
|
| 239 |
+
histories: list[torch.Tensor | None]
|
| 240 |
+
if self.use_cache and int(start_pos) > 0:
|
| 241 |
+
histories = [self.conv_state[:bsz, index] for index in range(3)]
|
| 242 |
+
else:
|
| 243 |
+
histories = [None, None, None]
|
| 244 |
+
outputs = []
|
| 245 |
+
next_histories = []
|
| 246 |
+
for module, value, history in zip(
|
| 247 |
+
(self.q_conv, self.k_conv, self.v_conv),
|
| 248 |
+
(q, k, v),
|
| 249 |
+
histories,
|
| 250 |
+
strict=True,
|
| 251 |
+
):
|
| 252 |
+
output, next_history = module(value, history=history)
|
| 253 |
+
outputs.append(output)
|
| 254 |
+
next_histories.append(next_history)
|
| 255 |
+
if self.use_cache:
|
| 256 |
+
self.ensure_batch_capacity(bsz)
|
| 257 |
+
with torch.no_grad():
|
| 258 |
+
for index, history in enumerate(next_histories):
|
| 259 |
+
self.conv_state[:bsz, index].copy_(
|
| 260 |
+
history.detach().to(dtype=torch.float32)
|
| 261 |
+
)
|
| 262 |
+
return outputs[0], outputs[1], outputs[2]
|
| 263 |
+
|
| 264 |
+
def _reference_kda(
|
| 265 |
+
self,
|
| 266 |
+
q: torch.Tensor,
|
| 267 |
+
k: torch.Tensor,
|
| 268 |
+
v: torch.Tensor,
|
| 269 |
+
decay_logits: torch.Tensor,
|
| 270 |
+
beta_logits: torch.Tensor,
|
| 271 |
+
initial_state: torch.Tensor | None,
|
| 272 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 273 |
+
output_dtype = v.dtype
|
| 274 |
+
q = functional.normalize(q.float(), p=2.0, dim=-1) * (
|
| 275 |
+
self.head_dim ** -0.5
|
| 276 |
+
)
|
| 277 |
+
k = functional.normalize(k.float(), p=2.0, dim=-1)
|
| 278 |
+
v = v.float()
|
| 279 |
+
scale = torch.exp(self.A_log).view(1, 1, self.num_heads, 1)
|
| 280 |
+
bias = self.dt_bias.view(1, 1, self.num_heads, self.head_dim)
|
| 281 |
+
log_decay = self.lower_bound * torch.sigmoid(
|
| 282 |
+
scale * (decay_logits.float() + bias)
|
| 283 |
+
)
|
| 284 |
+
alpha = torch.exp(log_decay)
|
| 285 |
+
beta = torch.sigmoid(beta_logits.float())
|
| 286 |
+
state = (
|
| 287 |
+
q.new_zeros(int(q.size(0)), self.num_heads, self.head_dim, self.head_dim)
|
| 288 |
+
if initial_state is None
|
| 289 |
+
else initial_state.float()
|
| 290 |
+
)
|
| 291 |
+
outputs: list[torch.Tensor] = []
|
| 292 |
+
for index in range(int(q.size(1))):
|
| 293 |
+
state = state * alpha[:, index].unsqueeze(-1)
|
| 294 |
+
prediction = torch.einsum("bhd,bhdv->bhv", k[:, index], state)
|
| 295 |
+
correction = (v[:, index] - prediction) * beta[:, index].unsqueeze(-1)
|
| 296 |
+
state = state + k[:, index].unsqueeze(-1) * correction.unsqueeze(-2)
|
| 297 |
+
outputs.append(torch.einsum("bhd,bhdv->bhv", q[:, index], state))
|
| 298 |
+
return torch.stack(outputs, dim=1).to(dtype=output_dtype), state
|
| 299 |
+
|
| 300 |
+
@staticmethod
|
| 301 |
+
def _fla_kda_available() -> bool:
|
| 302 |
+
try:
|
| 303 |
+
import_module("fla.ops.kda")
|
| 304 |
+
except (ImportError, AttributeError):
|
| 305 |
+
return False
|
| 306 |
+
return True
|
| 307 |
+
|
| 308 |
+
@staticmethod
|
| 309 |
+
def _fla_fused_recurrent_kda():
|
| 310 |
+
try:
|
| 311 |
+
return import_module("fla.ops.kda").fused_recurrent_kda
|
| 312 |
+
except (ImportError, AttributeError) as exc:
|
| 313 |
+
raise RuntimeError(
|
| 314 |
+
"CUDA Sophia KDA decoding requires flash-linear-attention>=0.5.2"
|
| 315 |
+
) from exc
|
| 316 |
+
|
| 317 |
+
@staticmethod
|
| 318 |
+
def _fla_chunk_kda():
|
| 319 |
+
try:
|
| 320 |
+
return import_module("fla.ops.kda").chunk_kda
|
| 321 |
+
except (ImportError, AttributeError) as exc:
|
| 322 |
+
raise RuntimeError(
|
| 323 |
+
"CUDA Sophia KDA training requires flash-linear-attention>=0.5.2"
|
| 324 |
+
) from exc
|
| 325 |
+
|
| 326 |
+
def _run_kda(
|
| 327 |
+
self,
|
| 328 |
+
q: torch.Tensor,
|
| 329 |
+
k: torch.Tensor,
|
| 330 |
+
v: torch.Tensor,
|
| 331 |
+
decay_logits: torch.Tensor,
|
| 332 |
+
beta_logits: torch.Tensor,
|
| 333 |
+
initial_state: torch.Tensor | None,
|
| 334 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 335 |
+
use_fla = (
|
| 336 |
+
q.device.type == "cuda"
|
| 337 |
+
and self.backend != "reference"
|
| 338 |
+
and (self.backend == "fla" or self._fla_kda_available())
|
| 339 |
+
)
|
| 340 |
+
if self.backend == "fla" and q.device.type != "cuda":
|
| 341 |
+
raise RuntimeError("kda_backend='fla' requires CUDA tensors")
|
| 342 |
+
if not use_fla:
|
| 343 |
+
return self._reference_kda(
|
| 344 |
+
q, k, v, decay_logits, beta_logits, initial_state
|
| 345 |
+
)
|
| 346 |
+
if int(q.size(1)) == 1 and initial_state is not None:
|
| 347 |
+
# Single-step decode: advancing the state by one position does not
|
| 348 |
+
# need the chunk grid.
|
| 349 |
+
output, final_state = self._fla_fused_recurrent_kda()(
|
| 350 |
+
q=q,
|
| 351 |
+
k=k,
|
| 352 |
+
v=v,
|
| 353 |
+
g=decay_logits,
|
| 354 |
+
beta=beta_logits,
|
| 355 |
+
A_log=self.A_log,
|
| 356 |
+
dt_bias=self.dt_bias,
|
| 357 |
+
initial_state=initial_state,
|
| 358 |
+
output_final_state=True,
|
| 359 |
+
use_qk_l2norm_in_kernel=True,
|
| 360 |
+
use_gate_in_kernel=True,
|
| 361 |
+
use_beta_sigmoid_in_kernel=True,
|
| 362 |
+
lower_bound=self.lower_bound,
|
| 363 |
+
)
|
| 364 |
+
return output, final_state
|
| 365 |
+
chunk_kda = self._fla_chunk_kda()
|
| 366 |
+
output, final_state = chunk_kda(
|
| 367 |
+
q=q,
|
| 368 |
+
k=k,
|
| 369 |
+
v=v,
|
| 370 |
+
g=decay_logits,
|
| 371 |
+
beta=beta_logits,
|
| 372 |
+
initial_state=initial_state,
|
| 373 |
+
output_final_state=bool(self.use_cache),
|
| 374 |
+
use_qk_l2norm_in_kernel=True,
|
| 375 |
+
use_gate_in_kernel=True,
|
| 376 |
+
use_beta_sigmoid_in_kernel=True,
|
| 377 |
+
safe_gate=True,
|
| 378 |
+
lower_bound=self.lower_bound,
|
| 379 |
+
A_log=self.A_log,
|
| 380 |
+
dt_bias=self.dt_bias,
|
| 381 |
+
)
|
| 382 |
+
if final_state is None:
|
| 383 |
+
final_state = q.new_zeros(
|
| 384 |
+
int(q.size(0)), self.num_heads, self.head_dim, self.head_dim,
|
| 385 |
+
dtype=torch.float32,
|
| 386 |
+
)
|
| 387 |
+
return output, final_state
|
| 388 |
+
|
| 389 |
+
def forward(self, x: torch.Tensor, *, start_pos: int = 0) -> torch.Tensor:
|
| 390 |
+
bsz, seqlen, _ = x.shape
|
| 391 |
+
if self.use_cache:
|
| 392 |
+
self.ensure_batch_capacity(int(bsz))
|
| 393 |
+
# Padding only ever exists in the prompt, so it applies to the prefill
|
| 394 |
+
# alone; every decode step is a real token for every row.
|
| 395 |
+
pad_mask = getattr(self, "_pad_mask", None) if int(start_pos) == 0 else None
|
| 396 |
+
if pad_mask is not None:
|
| 397 |
+
# Zero the input so the causal depthwise convolution sees exactly the
|
| 398 |
+
# zeros an unpadded sequence sees in its own left padding.
|
| 399 |
+
keep = pad_mask[:, :seqlen].to(dtype=x.dtype).unsqueeze(-1)
|
| 400 |
+
x = x * keep
|
| 401 |
+
q, k, v = self._short_convolution(
|
| 402 |
+
self.q_proj(x), self.k_proj(x), self.v_proj(x), start_pos=int(start_pos)
|
| 403 |
+
)
|
| 404 |
+
q = q.view(bsz, seqlen, self.num_heads, self.head_dim)
|
| 405 |
+
k = k.view(bsz, seqlen, self.num_heads, self.head_dim)
|
| 406 |
+
v = v.view(bsz, seqlen, self.num_heads, self.head_dim)
|
| 407 |
+
decay_logits = self.decay_up(self.decay_down(x)).view_as(q)
|
| 408 |
+
beta_logits = self.beta_proj(x)
|
| 409 |
+
if pad_mask is not None:
|
| 410 |
+
# alpha -> 1 and beta -> 0, so the state passes a pad through unchanged.
|
| 411 |
+
drop = ~pad_mask[:, :seqlen].bool()
|
| 412 |
+
decay_logits = decay_logits.masked_fill(drop[:, :, None, None], PAD_LOGIT)
|
| 413 |
+
beta_logits = beta_logits.masked_fill(drop[:, :, None], PAD_LOGIT)
|
| 414 |
+
initial_state = (
|
| 415 |
+
self.recurrent_state[:bsz]
|
| 416 |
+
if self.use_cache and int(start_pos) > 0
|
| 417 |
+
else None
|
| 418 |
+
)
|
| 419 |
+
output, final_state = self._run_kda(
|
| 420 |
+
q, k, v, decay_logits, beta_logits, initial_state
|
| 421 |
+
)
|
| 422 |
+
if self.use_cache:
|
| 423 |
+
with torch.no_grad():
|
| 424 |
+
self.recurrent_state[:bsz].copy_(
|
| 425 |
+
final_state.detach().to(dtype=torch.float32)
|
| 426 |
+
)
|
| 427 |
+
gate_logits = (
|
| 428 |
+
self.output_gate(x)
|
| 429 |
+
if self.output_gate_full_rank
|
| 430 |
+
else self.output_gate_up(self.output_gate_down(x))
|
| 431 |
+
)
|
| 432 |
+
gate = torch.sigmoid(gate_logits).view(
|
| 433 |
+
bsz, seqlen, self.num_heads, self.head_dim
|
| 434 |
+
)
|
| 435 |
+
output = self.output_norm(output.to(dtype=x.dtype)) * gate
|
| 436 |
+
return self.o_proj(output.reshape(bsz, seqlen, self.inner_dim))
|
| 437 |
+
|
| 438 |
+
|
| 439 |
+
class SophiaMLA(nn.Module):
|
| 440 |
+
"""NoPE multi-head latent attention with a full-rank output gate."""
|
| 441 |
+
|
| 442 |
+
def __init__(self, args: ModelArgs) -> None:
|
| 443 |
+
super().__init__()
|
| 444 |
+
self.dim = int(args.dim)
|
| 445 |
+
self.num_heads = int(args.num_heads)
|
| 446 |
+
self.head_dim = int(args.head_dim)
|
| 447 |
+
self.inner_dim = self.num_heads * self.head_dim
|
| 448 |
+
self.q_rank = int(args.mla_q_rank)
|
| 449 |
+
self.kv_rank = int(args.mla_kv_rank)
|
| 450 |
+
self.use_cache = bool(args.use_cache)
|
| 451 |
+
self.logit_telemetry_enabled = False
|
| 452 |
+
|
| 453 |
+
self.q_down = RuntimeLinear(self.dim, self.q_rank, bias=False)
|
| 454 |
+
self.q_norm = RMSNorm(self.q_rank, args.norm_eps)
|
| 455 |
+
self.q_up = RuntimeLinear(self.q_rank, self.inner_dim, bias=False)
|
| 456 |
+
|
| 457 |
+
self.kv_down = RuntimeLinear(self.dim, self.kv_rank, bias=False)
|
| 458 |
+
self.kv_norm = RMSNorm(self.kv_rank, args.norm_eps)
|
| 459 |
+
self.k_up = RuntimeLinear(self.kv_rank, self.inner_dim, bias=False)
|
| 460 |
+
self.v_up = RuntimeLinear(self.kv_rank, self.inner_dim, bias=False)
|
| 461 |
+
|
| 462 |
+
self.output_gate = RuntimeLinear(self.dim, self.inner_dim, bias=False)
|
| 463 |
+
self.o_proj = RuntimeLinear(self.inner_dim, self.dim, bias=False)
|
| 464 |
+
self.register_buffer(
|
| 465 |
+
"latent_cache",
|
| 466 |
+
torch.zeros(
|
| 467 |
+
int(args.max_batch_size), int(args.max_seq_len), self.kv_rank
|
| 468 |
+
),
|
| 469 |
+
persistent=False,
|
| 470 |
+
)
|
| 471 |
+
self.register_buffer(
|
| 472 |
+
"attention_logit_max",
|
| 473 |
+
torch.full((self.num_heads,), float("-inf"), dtype=torch.float32),
|
| 474 |
+
persistent=False,
|
| 475 |
+
)
|
| 476 |
+
|
| 477 |
+
def enable_attention_logit_telemetry(
|
| 478 |
+
self,
|
| 479 |
+
buffer: torch.Tensor | None = None,
|
| 480 |
+
) -> None:
|
| 481 |
+
self.logit_telemetry_enabled = True
|
| 482 |
+
if buffer is None:
|
| 483 |
+
buffer = torch.full(
|
| 484 |
+
(self.num_heads,),
|
| 485 |
+
float("-inf"),
|
| 486 |
+
device=self.q_up.weight.device,
|
| 487 |
+
dtype=torch.float32,
|
| 488 |
+
)
|
| 489 |
+
if (
|
| 490 |
+
tuple(buffer.shape) != (self.num_heads,)
|
| 491 |
+
or buffer.device != self.q_up.weight.device
|
| 492 |
+
or buffer.dtype != torch.float32
|
| 493 |
+
):
|
| 494 |
+
raise ValueError("MLA attention-logit buffer must be FP32 [num_heads]")
|
| 495 |
+
self.attention_logit_max = buffer
|
| 496 |
+
|
| 497 |
+
def reset_attention_logit_max(self) -> None:
|
| 498 |
+
self.attention_logit_max.fill_(float("-inf"))
|
| 499 |
+
|
| 500 |
+
@torch.no_grad()
|
| 501 |
+
def _record_attention_logit_max(
|
| 502 |
+
self,
|
| 503 |
+
q: torch.Tensor,
|
| 504 |
+
k: torch.Tensor,
|
| 505 |
+
) -> None:
|
| 506 |
+
query_tile = 512
|
| 507 |
+
query_length = int(q.size(-2))
|
| 508 |
+
key_length = int(k.size(-2))
|
| 509 |
+
query_offset = key_length - query_length
|
| 510 |
+
if query_offset < 0:
|
| 511 |
+
raise ValueError("MLA telemetry requires key length >= query length")
|
| 512 |
+
tile_causal = torch.ones(
|
| 513 |
+
(min(query_tile, query_length), min(query_tile, query_length)),
|
| 514 |
+
device=q.device,
|
| 515 |
+
dtype=torch.bool,
|
| 516 |
+
).tril_()
|
| 517 |
+
key_transposed = k.detach().transpose(-2, -1)
|
| 518 |
+
running = self.attention_logit_max
|
| 519 |
+
scale = self.head_dim**-0.5
|
| 520 |
+
for query_start in range(0, query_length, query_tile):
|
| 521 |
+
query_end = min(query_start + query_tile, query_length)
|
| 522 |
+
key_end = query_offset + query_end
|
| 523 |
+
tile_width = query_end - query_start
|
| 524 |
+
logits = torch.matmul(
|
| 525 |
+
q.detach()[..., query_start:query_end, :],
|
| 526 |
+
key_transposed[..., :key_end],
|
| 527 |
+
) * float(scale)
|
| 528 |
+
logits[..., -tile_width:].masked_fill_(
|
| 529 |
+
~tile_causal[:tile_width, :tile_width].view(
|
| 530 |
+
1, 1, tile_width, tile_width
|
| 531 |
+
),
|
| 532 |
+
float("-inf"),
|
| 533 |
+
)
|
| 534 |
+
running = torch.maximum(
|
| 535 |
+
running,
|
| 536 |
+
logits.amax(dim=(0, 2, 3)).float(),
|
| 537 |
+
)
|
| 538 |
+
self.attention_logit_max.copy_(running)
|
| 539 |
+
|
| 540 |
+
def ensure_batch_capacity(self, batch_size: int) -> None:
|
| 541 |
+
required = int(batch_size)
|
| 542 |
+
if required <= int(self.latent_cache.size(0)):
|
| 543 |
+
return
|
| 544 |
+
grown = self.latent_cache.new_zeros(
|
| 545 |
+
required, int(self.latent_cache.size(1)), self.kv_rank
|
| 546 |
+
)
|
| 547 |
+
grown[: self.latent_cache.size(0)].copy_(self.latent_cache)
|
| 548 |
+
self.latent_cache = grown
|
| 549 |
+
|
| 550 |
+
def ensure_sequence_capacity(self, max_seq_len: int) -> None:
|
| 551 |
+
required = int(max_seq_len)
|
| 552 |
+
if required <= int(self.latent_cache.size(1)):
|
| 553 |
+
return
|
| 554 |
+
grown = self.latent_cache.new_zeros(
|
| 555 |
+
int(self.latent_cache.size(0)), required, self.kv_rank
|
| 556 |
+
)
|
| 557 |
+
grown[:, : self.latent_cache.size(1)].copy_(self.latent_cache)
|
| 558 |
+
self.latent_cache = grown
|
| 559 |
+
|
| 560 |
+
def rebuild_runtime_buffers(self, max_seq_len: int) -> None:
|
| 561 |
+
self.latent_cache = self.kv_down.weight.new_zeros(
|
| 562 |
+
int(self.latent_cache.size(0)), int(max_seq_len), self.kv_rank
|
| 563 |
+
)
|
| 564 |
+
|
| 565 |
+
def reset(self) -> None:
|
| 566 |
+
if self.latent_cache.is_inference():
|
| 567 |
+
self.latent_cache = torch.zeros_like(self.latent_cache)
|
| 568 |
+
else:
|
| 569 |
+
self.latent_cache.zero_()
|
| 570 |
+
|
| 571 |
+
def refresh_state_buffers(self) -> None:
|
| 572 |
+
self.latent_cache = torch.zeros_like(self.latent_cache)
|
| 573 |
+
|
| 574 |
+
def cache_snapshot(
|
| 575 |
+
self, *, device: str, batch_size: int, cache_pos: int | None
|
| 576 |
+
) -> LayerCacheSnapshot:
|
| 577 |
+
end = int(self.latent_cache.size(1) if cache_pos is None else cache_pos)
|
| 578 |
+
return LayerCacheSnapshot(
|
| 579 |
+
latent=self.latent_cache[:batch_size, :end].to(device).clone()
|
| 580 |
+
)
|
| 581 |
+
|
| 582 |
+
def validate_cache_snapshot(self, snapshot: LayerCacheSnapshot) -> None:
|
| 583 |
+
if snapshot.recurrent is not None or snapshot.conv is not None:
|
| 584 |
+
raise ValueError("MLA cache cannot contain KDA state")
|
| 585 |
+
if snapshot.latent is not None and int(snapshot.latent.size(-1)) != self.kv_rank:
|
| 586 |
+
raise ValueError("MLA latent cache shape mismatch")
|
| 587 |
+
|
| 588 |
+
def load_cache_snapshot(self, snapshot: LayerCacheSnapshot) -> None:
|
| 589 |
+
self.validate_cache_snapshot(snapshot)
|
| 590 |
+
if snapshot.latent is None:
|
| 591 |
+
return
|
| 592 |
+
batch, length = int(snapshot.latent.size(0)), int(snapshot.latent.size(1))
|
| 593 |
+
self.ensure_batch_capacity(batch)
|
| 594 |
+
self.ensure_sequence_capacity(length)
|
| 595 |
+
self.latent_cache[:batch, :length].copy_(
|
| 596 |
+
snapshot.latent.to(
|
| 597 |
+
self.latent_cache.device, dtype=self.latent_cache.dtype
|
| 598 |
+
)
|
| 599 |
+
)
|
| 600 |
+
|
| 601 |
+
def _resolve_key_pad(
|
| 602 |
+
self, start_pos: int, bsz: int, k_len: int, device: torch.device
|
| 603 |
+
) -> torch.Tensor | None:
|
| 604 |
+
"""Which key positions are real, over the whole cached span."""
|
| 605 |
+
if int(start_pos) == 0:
|
| 606 |
+
mask = getattr(self, "_pad_mask", None)
|
| 607 |
+
self._key_pad = None if mask is None else mask[:bsz].bool().to(device)
|
| 608 |
+
return self._key_pad
|
| 609 |
+
cached = getattr(self, "_key_pad", None)
|
| 610 |
+
if cached is None:
|
| 611 |
+
return None
|
| 612 |
+
prompt_len = int(cached.size(1))
|
| 613 |
+
if k_len <= prompt_len:
|
| 614 |
+
return cached[:bsz, :k_len]
|
| 615 |
+
grown = torch.ones((bsz, k_len), dtype=torch.bool, device=device)
|
| 616 |
+
grown[:, :prompt_len] = cached[:bsz]
|
| 617 |
+
return grown
|
| 618 |
+
|
| 619 |
+
@staticmethod
|
| 620 |
+
def _causal_bias(q: torch.Tensor, k: torch.Tensor, start_pos: int) -> CausalBias | None:
|
| 621 |
+
if int(start_pos) == 0:
|
| 622 |
+
return None
|
| 623 |
+
return causal_lower_right(int(q.size(1)), int(k.size(1)))
|
| 624 |
+
|
| 625 |
+
def forward(self, x: torch.Tensor, *, start_pos: int = 0) -> torch.Tensor:
|
| 626 |
+
bsz, seqlen, _ = x.shape
|
| 627 |
+
q = self.q_up(self.q_norm(self.q_down(x))).view(
|
| 628 |
+
bsz, seqlen, self.num_heads, self.head_dim
|
| 629 |
+
)
|
| 630 |
+
latent = self.kv_norm(self.kv_down(x))
|
| 631 |
+
if self.use_cache:
|
| 632 |
+
end = int(start_pos) + int(seqlen)
|
| 633 |
+
self.ensure_batch_capacity(int(bsz))
|
| 634 |
+
self.ensure_sequence_capacity(end)
|
| 635 |
+
with torch.no_grad():
|
| 636 |
+
self.latent_cache[:bsz, int(start_pos) : end].copy_(latent.detach())
|
| 637 |
+
latent_full = self.latent_cache[:bsz, :end]
|
| 638 |
+
else:
|
| 639 |
+
latent_full = latent
|
| 640 |
+
k = self.k_up(latent_full).view(
|
| 641 |
+
bsz, int(latent_full.size(1)), self.num_heads, self.head_dim
|
| 642 |
+
)
|
| 643 |
+
v = self.v_up(latent_full).view_as(k)
|
| 644 |
+
q_t, k_t, v_t = (value.transpose(1, 2) for value in (q, k, v))
|
| 645 |
+
if self.logit_telemetry_enabled:
|
| 646 |
+
self._record_attention_logit_max(q_t, k_t)
|
| 647 |
+
key_pad = self._resolve_key_pad(int(start_pos), int(bsz), int(k.size(1)), x.device)
|
| 648 |
+
if key_pad is None:
|
| 649 |
+
bias = self._causal_bias(q, k, int(start_pos))
|
| 650 |
+
output = functional.scaled_dot_product_attention(
|
| 651 |
+
q_t,
|
| 652 |
+
k_t,
|
| 653 |
+
v_t,
|
| 654 |
+
attn_mask=bias,
|
| 655 |
+
is_causal=(bias is None),
|
| 656 |
+
).transpose(1, 2)
|
| 657 |
+
else:
|
| 658 |
+
# causal_lower_right is an opaque bias object, so when padding is in
|
| 659 |
+
# play build the whole additive mask instead of trying to combine.
|
| 660 |
+
q_len, k_len = int(q.size(1)), int(k.size(1))
|
| 661 |
+
allow = key_pad[:, None, None, :].expand(bsz, 1, q_len, k_len).clone()
|
| 662 |
+
if q_len > 1:
|
| 663 |
+
offset = k_len - q_len
|
| 664 |
+
causal = torch.ones((q_len, k_len), dtype=torch.bool, device=x.device).tril(offset)
|
| 665 |
+
allow &= causal[None, None]
|
| 666 |
+
attn_mask = torch.zeros((bsz, 1, q_len, k_len), dtype=q_t.dtype, device=x.device)
|
| 667 |
+
attn_mask = attn_mask.masked_fill(~allow, float("-inf"))
|
| 668 |
+
output = functional.scaled_dot_product_attention(
|
| 669 |
+
q_t, k_t, v_t, attn_mask=attn_mask, is_causal=False
|
| 670 |
+
).transpose(1, 2)
|
| 671 |
+
gate = torch.sigmoid(self.output_gate(x)).view_as(output)
|
| 672 |
+
gated_output = (output.float() * gate.float()).to(dtype=x.dtype)
|
| 673 |
+
return self.o_proj(gated_output.reshape(bsz, seqlen, self.inner_dim))
|
| 674 |
+
|
| 675 |
+
|
| 676 |
+
def set_pad_mask(model: nn.Module, mask: torch.Tensor | None) -> None:
|
| 677 |
+
"""Hang a [B, T] boolean mask (True = real token) on every mixer, or clear it.
|
| 678 |
+
|
| 679 |
+
Clearing also drops MLA's remembered key mask, so a later unpadded batch cannot
|
| 680 |
+
inherit the previous batch's padding.
|
| 681 |
+
"""
|
| 682 |
+
for module in model.modules():
|
| 683 |
+
if isinstance(module, (SophiaKDA, SophiaMLA)):
|
| 684 |
+
module._pad_mask = mask
|
| 685 |
+
if mask is None and isinstance(module, SophiaMLA):
|
| 686 |
+
module._key_pad = None
|
| 687 |
+
|
| 688 |
+
|
| 689 |
+
__all__ = [
|
| 690 |
+
"CausalDepthwiseConv1d",
|
| 691 |
+
"SophiaKDA",
|
| 692 |
+
"SophiaMLA",
|
| 693 |
+
"set_pad_mask",
|
| 694 |
+
]
|
model_blocks.py
ADDED
|
@@ -0,0 +1,94 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
import torch
|
| 8 |
+
from torch import nn
|
| 9 |
+
|
| 10 |
+
from .model_attention import SophiaKDA, SophiaMLA
|
| 11 |
+
from .model_config import ModelArgs
|
| 12 |
+
from .model_ops import RMSNorm
|
| 13 |
+
from .runtime_linear import RuntimeLinear
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
class FeedForward(nn.Module):
|
| 17 |
+
def __init__(self, args: ModelArgs) -> None:
|
| 18 |
+
super().__init__()
|
| 19 |
+
self.hidden = int(args.ffn_hidden)
|
| 20 |
+
self.gate_softcap = float(args.situ_gate_softcap)
|
| 21 |
+
self.up_softcap = float(args.situ_up_softcap)
|
| 22 |
+
self.gate_up_proj = RuntimeLinear(
|
| 23 |
+
int(args.dim), 2 * self.hidden, bias=False
|
| 24 |
+
)
|
| 25 |
+
self.down_proj = RuntimeLinear(self.hidden, int(args.dim), bias=False)
|
| 26 |
+
|
| 27 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 28 |
+
gate, up = self.gate_up_proj(x).split(self.hidden, dim=-1)
|
| 29 |
+
bounded_gate = self.gate_softcap * torch.tanh(gate / self.gate_softcap)
|
| 30 |
+
bounded_up = self.up_softcap * torch.tanh(up / self.up_softcap)
|
| 31 |
+
return self.down_proj(bounded_gate * torch.sigmoid(gate) * bounded_up)
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
class AttentionResidualMixer(nn.Module):
|
| 35 |
+
"""Content-dependent mixing over completed depth blocks and a partial block."""
|
| 36 |
+
|
| 37 |
+
def __init__(self, args: ModelArgs) -> None:
|
| 38 |
+
super().__init__()
|
| 39 |
+
self.norm = RMSNorm(args.dim, args.norm_eps)
|
| 40 |
+
self.query = RuntimeLinear(args.dim, 1, bias=False)
|
| 41 |
+
|
| 42 |
+
def forward(
|
| 43 |
+
self,
|
| 44 |
+
partial_block: torch.Tensor,
|
| 45 |
+
completed_blocks: torch.Tensor,
|
| 46 |
+
) -> torch.Tensor:
|
| 47 |
+
values = torch.cat((completed_blocks, partial_block.unsqueeze(2)), dim=2)
|
| 48 |
+
keys = self.norm(values)
|
| 49 |
+
score_weight = self.query.weight.squeeze(0).float()
|
| 50 |
+
scores = torch.einsum("bsth,h->bst", keys.float(), score_weight)
|
| 51 |
+
weights = scores.softmax(dim=2).unsqueeze(-1).to(dtype=values.dtype)
|
| 52 |
+
return (weights * values).sum(dim=2).to(dtype=values.dtype)
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
class SophiaBlock(nn.Module):
|
| 56 |
+
def __init__(self, args: ModelArgs, layer_idx: int) -> None:
|
| 57 |
+
super().__init__()
|
| 58 |
+
self.layer_idx = int(layer_idx)
|
| 59 |
+
self.layer_type = args.layer_type(self.layer_idx)
|
| 60 |
+
self.attn_norm = RMSNorm(args.dim, args.norm_eps)
|
| 61 |
+
self.attn = (
|
| 62 |
+
SophiaMLA(args) if self.layer_type == "mla" else SophiaKDA(args)
|
| 63 |
+
)
|
| 64 |
+
self.ffn_norm = RMSNorm(args.dim, args.norm_eps)
|
| 65 |
+
self.ffn = FeedForward(args)
|
| 66 |
+
self.attn_res_block_size = int(args.attn_res_block_size)
|
| 67 |
+
self.attn_residual = AttentionResidualMixer(args)
|
| 68 |
+
self.ffn_residual = AttentionResidualMixer(args)
|
| 69 |
+
|
| 70 |
+
def forward(
|
| 71 |
+
self,
|
| 72 |
+
partial_block: torch.Tensor,
|
| 73 |
+
completed_blocks: torch.Tensor,
|
| 74 |
+
*,
|
| 75 |
+
start_pos: int = 0,
|
| 76 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 77 |
+
attn_input = self.attn_residual(partial_block, completed_blocks)
|
| 78 |
+
if self.layer_idx % self.attn_res_block_size == 0:
|
| 79 |
+
completed_blocks = torch.cat(
|
| 80 |
+
(completed_blocks, partial_block.unsqueeze(2)), dim=2
|
| 81 |
+
)
|
| 82 |
+
partial_block = self.attn(
|
| 83 |
+
self.attn_norm(attn_input), start_pos=int(start_pos)
|
| 84 |
+
)
|
| 85 |
+
else:
|
| 86 |
+
partial_block = partial_block + self.attn(
|
| 87 |
+
self.attn_norm(attn_input), start_pos=int(start_pos)
|
| 88 |
+
)
|
| 89 |
+
ffn_input = self.ffn_residual(partial_block, completed_blocks)
|
| 90 |
+
partial_block = partial_block + self.ffn(self.ffn_norm(ffn_input))
|
| 91 |
+
return partial_block, completed_blocks
|
| 92 |
+
|
| 93 |
+
|
| 94 |
+
__all__ = ["AttentionResidualMixer", "FeedForward", "SophiaBlock"]
|
model_config.py
ADDED
|
@@ -0,0 +1,87 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
from dataclasses import dataclass
|
| 8 |
+
|
| 9 |
+
from .semantics import canonicalize_model_values, validate_model_values
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
@dataclass
|
| 13 |
+
class ModelArgs:
|
| 14 |
+
"""Torch-free runtime configuration for Sophia Hybrid."""
|
| 15 |
+
|
| 16 |
+
vocab_size: int = 65536
|
| 17 |
+
dim: int = 1536
|
| 18 |
+
n_layers: int = 28
|
| 19 |
+
num_heads: int = 16
|
| 20 |
+
head_dim: int = 128
|
| 21 |
+
ffn_hidden: int = 3968
|
| 22 |
+
kda_decay_rank: int = 128
|
| 23 |
+
kda_output_gate_rank: int = 128
|
| 24 |
+
kda_output_gate_full_rank: bool = True
|
| 25 |
+
kda_decay_lower_bound: float = -5.0
|
| 26 |
+
kda_dt_min: float = 1e-3
|
| 27 |
+
kda_dt_max: float = 1e-1
|
| 28 |
+
kda_dt_floor: float = 1e-4
|
| 29 |
+
kda_a_log_init: float = 0.0
|
| 30 |
+
mla_q_rank: int = 384
|
| 31 |
+
mla_kv_rank: int = 128
|
| 32 |
+
short_conv_kernel: int = 4
|
| 33 |
+
attn_res_block_size: int = 4
|
| 34 |
+
situ_gate_softcap: float = 4.0
|
| 35 |
+
situ_up_softcap: float = 25.0
|
| 36 |
+
norm_eps: float = 1e-5
|
| 37 |
+
max_seq_len: int = 4096
|
| 38 |
+
max_batch_size: int = 4
|
| 39 |
+
dropout: float = 0.0
|
| 40 |
+
initializer_range: float = 0.02
|
| 41 |
+
kda_backend: str = "auto"
|
| 42 |
+
use_cache: bool = True
|
| 43 |
+
|
| 44 |
+
def __post_init__(self) -> None:
|
| 45 |
+
normalized = canonicalize_model_values(vars(self))
|
| 46 |
+
for name, value in normalized.items():
|
| 47 |
+
if hasattr(self, name):
|
| 48 |
+
setattr(self, name, value)
|
| 49 |
+
validate_model_values(normalized)
|
| 50 |
+
|
| 51 |
+
@property
|
| 52 |
+
def attention_inner_dim(self) -> int:
|
| 53 |
+
return int(self.num_heads) * int(self.head_dim)
|
| 54 |
+
|
| 55 |
+
def layer_type(self, layer_idx: int) -> str:
|
| 56 |
+
return "mla" if int(layer_idx) % 4 == 3 else "kda"
|
| 57 |
+
|
| 58 |
+
def ensure_runtime_batch_capacity(self, batch_size: int) -> int:
|
| 59 |
+
required = int(batch_size)
|
| 60 |
+
if required <= 0:
|
| 61 |
+
raise ValueError(f"batch_size must be > 0, got {required}")
|
| 62 |
+
self.max_batch_size = max(int(self.max_batch_size), required)
|
| 63 |
+
return int(self.max_batch_size)
|
| 64 |
+
|
| 65 |
+
def runtime_batch_capacity(self) -> int:
|
| 66 |
+
return int(self.max_batch_size)
|
| 67 |
+
|
| 68 |
+
def ensure_runtime_sequence_capacity(self, max_seq_len: int) -> int:
|
| 69 |
+
required = int(max_seq_len)
|
| 70 |
+
if required <= 0:
|
| 71 |
+
raise ValueError(f"max_seq_len must be > 0, got {required}")
|
| 72 |
+
configured = int(self.max_seq_len)
|
| 73 |
+
if required > configured:
|
| 74 |
+
raise ValueError(
|
| 75 |
+
"runtime sequence capacity cannot exceed configured max_seq_len: "
|
| 76 |
+
f"required={required} configured={configured}"
|
| 77 |
+
)
|
| 78 |
+
return configured
|
| 79 |
+
|
| 80 |
+
def runtime_sequence_capacity(self) -> int:
|
| 81 |
+
return int(self.max_seq_len)
|
| 82 |
+
|
| 83 |
+
def runtime_max_seq_len(self) -> int:
|
| 84 |
+
return self.runtime_sequence_capacity()
|
| 85 |
+
|
| 86 |
+
def runtime_layer_count(self) -> int:
|
| 87 |
+
return int(self.n_layers)
|
model_dir.py
ADDED
|
@@ -0,0 +1,82 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
"""Local exported-model directory helpers."""
|
| 6 |
+
|
| 7 |
+
from __future__ import annotations
|
| 8 |
+
|
| 9 |
+
import importlib.util
|
| 10 |
+
import os
|
| 11 |
+
from collections.abc import Callable
|
| 12 |
+
|
| 13 |
+
import torch
|
| 14 |
+
|
| 15 |
+
from .sophia_decoder import SophiaDecoder
|
| 16 |
+
from ml.integrations.adapters.hf.tokenizer import load_local_tokenizer
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
GradientCheckpointingFn = Callable[..., None]
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
def local_export_load_kwargs() -> dict[str, object]:
|
| 23 |
+
kwargs: dict[str, object] = {
|
| 24 |
+
"dtype": torch.float32,
|
| 25 |
+
"local_files_only": True,
|
| 26 |
+
}
|
| 27 |
+
if importlib.util.find_spec("accelerate") is not None:
|
| 28 |
+
kwargs["low_cpu_mem_usage"] = True
|
| 29 |
+
return kwargs
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
load_local_export_tokenizer = load_local_tokenizer
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
def load_trainable_decoder_from_export_dir(
|
| 36 |
+
*,
|
| 37 |
+
export_dir: str,
|
| 38 |
+
device: torch.device,
|
| 39 |
+
base_dtype: torch.dtype,
|
| 40 |
+
gradient_checkpointing: bool,
|
| 41 |
+
gradient_checkpointing_exclude_first: int,
|
| 42 |
+
gradient_checkpointing_exclude_last: int,
|
| 43 |
+
apply_gradient_checkpointing_fn: GradientCheckpointingFn,
|
| 44 |
+
) -> SophiaDecoder:
|
| 45 |
+
model = SophiaDecoder.from_pretrained(
|
| 46 |
+
os.path.abspath(str(export_dir)),
|
| 47 |
+
device=device,
|
| 48 |
+
dtype=base_dtype,
|
| 49 |
+
)
|
| 50 |
+
model.train(True)
|
| 51 |
+
config = getattr(model, "config", None)
|
| 52 |
+
if config is not None:
|
| 53 |
+
config.return_logits_in_train = True
|
| 54 |
+
config.use_cache = False
|
| 55 |
+
exclude_first = int(gradient_checkpointing_exclude_first)
|
| 56 |
+
exclude_last = int(gradient_checkpointing_exclude_last)
|
| 57 |
+
if min(exclude_first, exclude_last) < 0:
|
| 58 |
+
raise ValueError("gradient checkpoint exclusion counts must be >= 0")
|
| 59 |
+
if not bool(gradient_checkpointing) and (exclude_first or exclude_last):
|
| 60 |
+
raise ValueError(
|
| 61 |
+
"gradient checkpoint exclusions require gradient_checkpointing=1"
|
| 62 |
+
)
|
| 63 |
+
if bool(gradient_checkpointing) and exclude_first + exclude_last >= int(
|
| 64 |
+
model.config.n_layers
|
| 65 |
+
):
|
| 66 |
+
raise ValueError(
|
| 67 |
+
"gradient checkpoint exclusions must leave at least one checkpointed layer"
|
| 68 |
+
)
|
| 69 |
+
model.apply_runtime_recipe_knobs(
|
| 70 |
+
loss_chunk_size=int(model.config.loss_chunk_size),
|
| 71 |
+
gradient_checkpointing_exclude_first=exclude_first,
|
| 72 |
+
gradient_checkpointing_exclude_last=exclude_last,
|
| 73 |
+
)
|
| 74 |
+
apply_gradient_checkpointing_fn(model, enabled=bool(gradient_checkpointing))
|
| 75 |
+
return model
|
| 76 |
+
|
| 77 |
+
|
| 78 |
+
__all__ = [
|
| 79 |
+
"load_local_export_tokenizer",
|
| 80 |
+
"load_trainable_decoder_from_export_dir",
|
| 81 |
+
"local_export_load_kwargs",
|
| 82 |
+
]
|
model_ops.py
ADDED
|
@@ -0,0 +1,81 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
import torch
|
| 8 |
+
import torch.nn.functional as functional
|
| 9 |
+
from torch import nn
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
def infer_module_tensor_device(
|
| 14 |
+
module: nn.Module,
|
| 15 |
+
*,
|
| 16 |
+
default_device: torch.device,
|
| 17 |
+
) -> torch.device:
|
| 18 |
+
for parameter in module.parameters():
|
| 19 |
+
return parameter.device
|
| 20 |
+
for buffer in module.buffers():
|
| 21 |
+
if isinstance(buffer, torch.Tensor):
|
| 22 |
+
return buffer.device
|
| 23 |
+
return default_device
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
def apply_preserving_complex_buffers(
|
| 27 |
+
module: nn.Module,
|
| 28 |
+
fn,
|
| 29 |
+
*,
|
| 30 |
+
buffer_names: tuple[str, ...],
|
| 31 |
+
apply_super,
|
| 32 |
+
):
|
| 33 |
+
preserved: dict[str, torch.Tensor] = {}
|
| 34 |
+
preserved_device = torch.device("cpu")
|
| 35 |
+
for name in buffer_names:
|
| 36 |
+
tensor = module._buffers.get(name)
|
| 37 |
+
if isinstance(tensor, torch.Tensor) and tensor.is_complex():
|
| 38 |
+
preserved[name] = tensor
|
| 39 |
+
preserved_device = tensor.device
|
| 40 |
+
module._buffers[name] = None
|
| 41 |
+
try:
|
| 42 |
+
result = apply_super()
|
| 43 |
+
finally:
|
| 44 |
+
if preserved:
|
| 45 |
+
target_device = infer_module_tensor_device(
|
| 46 |
+
module,
|
| 47 |
+
default_device=preserved_device,
|
| 48 |
+
)
|
| 49 |
+
for name, tensor in preserved.items():
|
| 50 |
+
module._buffers[name] = tensor.to(device=target_device)
|
| 51 |
+
return result
|
| 52 |
+
|
| 53 |
+
class RMSNorm(nn.Module):
|
| 54 |
+
"""Root Mean Square Layer Normalization."""
|
| 55 |
+
|
| 56 |
+
def __init__(self, dim: int, eps: float = 1e-6):
|
| 57 |
+
super().__init__()
|
| 58 |
+
self.eps = eps
|
| 59 |
+
self.weight = nn.Parameter(torch.ones(dim, dtype=torch.float32))
|
| 60 |
+
self.weight._no_weight_decay = True # type: ignore[attr-defined]
|
| 61 |
+
|
| 62 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 63 |
+
return functional.rms_norm(
|
| 64 |
+
x,
|
| 65 |
+
(int(x.size(-1)),),
|
| 66 |
+
self.weight,
|
| 67 |
+
float(self.eps),
|
| 68 |
+
)
|
| 69 |
+
|
| 70 |
+
|
| 71 |
+
class StandardLogitMixer(nn.Module):
|
| 72 |
+
"""Decoder logits path: final norm followed by output projection."""
|
| 73 |
+
|
| 74 |
+
def forward(
|
| 75 |
+
self,
|
| 76 |
+
x: torch.Tensor,
|
| 77 |
+
*,
|
| 78 |
+
norm: RMSNorm,
|
| 79 |
+
output: nn.Module,
|
| 80 |
+
) -> torch.Tensor:
|
| 81 |
+
return output(norm(x))
|
model_runtime.py
ADDED
|
@@ -0,0 +1,219 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
from threading import RLock
|
| 8 |
+
|
| 9 |
+
import torch
|
| 10 |
+
|
| 11 |
+
from .model_runtime_control import (
|
| 12 |
+
RuntimeControlModel,
|
| 13 |
+
enable_runtime_cache,
|
| 14 |
+
ensure_batch_capacity,
|
| 15 |
+
ensure_sequence_capacity,
|
| 16 |
+
rebuild_runtime_buffers,
|
| 17 |
+
refresh_state_buffers,
|
| 18 |
+
reset_runtime_cache,
|
| 19 |
+
)
|
| 20 |
+
from .model_state import (
|
| 21 |
+
LayerCacheSnapshot,
|
| 22 |
+
RuntimeCacheSnapshot,
|
| 23 |
+
TransformerRuntimeState,
|
| 24 |
+
)
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
class TransformerRuntime:
|
| 28 |
+
def __init__(
|
| 29 |
+
self,
|
| 30 |
+
model: RuntimeControlModel,
|
| 31 |
+
*,
|
| 32 |
+
state: TransformerRuntimeState | None = None,
|
| 33 |
+
) -> None:
|
| 34 |
+
self.model = model
|
| 35 |
+
self.state = state or TransformerRuntimeState()
|
| 36 |
+
|
| 37 |
+
@property
|
| 38 |
+
def lock(self) -> RLock:
|
| 39 |
+
return self.state.lock
|
| 40 |
+
|
| 41 |
+
@lock.setter
|
| 42 |
+
def lock(self, value: RLock) -> None:
|
| 43 |
+
self.state.lock = value
|
| 44 |
+
|
| 45 |
+
def supports_loss_chunk_size(self) -> bool:
|
| 46 |
+
return True
|
| 47 |
+
|
| 48 |
+
def supports_checkpoint_excludes(self) -> bool:
|
| 49 |
+
return True
|
| 50 |
+
|
| 51 |
+
def runtime_recipe_knobs(self) -> tuple[int, int, int]:
|
| 52 |
+
return (
|
| 53 |
+
int(self.state.loss_chunk_size),
|
| 54 |
+
int(self.model.gradient_checkpointing_exclude_first),
|
| 55 |
+
int(self.model.gradient_checkpointing_exclude_last),
|
| 56 |
+
)
|
| 57 |
+
|
| 58 |
+
def apply_runtime_recipe_knobs(
|
| 59 |
+
self,
|
| 60 |
+
*,
|
| 61 |
+
loss_chunk_size: int,
|
| 62 |
+
gradient_checkpointing_exclude_first: int,
|
| 63 |
+
gradient_checkpointing_exclude_last: int,
|
| 64 |
+
) -> tuple[int, int, int]:
|
| 65 |
+
chunk_size = int(loss_chunk_size)
|
| 66 |
+
exclude_first = int(gradient_checkpointing_exclude_first)
|
| 67 |
+
exclude_last = int(gradient_checkpointing_exclude_last)
|
| 68 |
+
if chunk_size < 0:
|
| 69 |
+
raise ValueError(f"loss_chunk_size must be >= 0, got {chunk_size}")
|
| 70 |
+
if exclude_first < 0:
|
| 71 |
+
raise ValueError(
|
| 72 |
+
"gradient_checkpointing_exclude_first must be >= 0, "
|
| 73 |
+
f"got {exclude_first}"
|
| 74 |
+
)
|
| 75 |
+
if exclude_last < 0:
|
| 76 |
+
raise ValueError(
|
| 77 |
+
"gradient_checkpointing_exclude_last must be >= 0, "
|
| 78 |
+
f"got {exclude_last}"
|
| 79 |
+
)
|
| 80 |
+
layer_count = len(self.model.layers)
|
| 81 |
+
if exclude_first + exclude_last > layer_count:
|
| 82 |
+
raise ValueError(
|
| 83 |
+
"gradient checkpoint exclusions exceed model layer count: "
|
| 84 |
+
f"first={exclude_first} last={exclude_last} layers={layer_count}"
|
| 85 |
+
)
|
| 86 |
+
self.state.loss_chunk_size = chunk_size
|
| 87 |
+
self.model.gradient_checkpointing_exclude_first = exclude_first
|
| 88 |
+
self.model.gradient_checkpointing_exclude_last = exclude_last
|
| 89 |
+
return self.runtime_recipe_knobs()
|
| 90 |
+
|
| 91 |
+
def ensure_batch_capacity(self, batch_size: int) -> None:
|
| 92 |
+
ensure_batch_capacity(self.model, batch_size)
|
| 93 |
+
|
| 94 |
+
def ensure_sequence_capacity(self, max_seq_len: int) -> None:
|
| 95 |
+
ensure_sequence_capacity(self.model, max_seq_len)
|
| 96 |
+
|
| 97 |
+
def ensure_runtime_max_seq_len(self, max_seq_len: int) -> None:
|
| 98 |
+
self.ensure_sequence_capacity(max_seq_len)
|
| 99 |
+
|
| 100 |
+
def runtime_max_seq_len(self) -> int:
|
| 101 |
+
return int(self.model.runtime_max_seq_len())
|
| 102 |
+
|
| 103 |
+
def runtime_batch_capacity(self) -> int:
|
| 104 |
+
return int(self.model.runtime_batch_capacity())
|
| 105 |
+
|
| 106 |
+
def rebuild_runtime_buffers(self) -> None:
|
| 107 |
+
rebuild_runtime_buffers(self.model)
|
| 108 |
+
|
| 109 |
+
def sync_runtime_batch_capacity(self, max_batch_size: int) -> int:
|
| 110 |
+
self.ensure_batch_capacity(max_batch_size)
|
| 111 |
+
return self.runtime_batch_capacity()
|
| 112 |
+
|
| 113 |
+
def reset_runtime_cache(self) -> None:
|
| 114 |
+
reset_runtime_cache(self.model)
|
| 115 |
+
|
| 116 |
+
def refresh_state_buffers(self) -> None:
|
| 117 |
+
refresh_state_buffers(self.model)
|
| 118 |
+
|
| 119 |
+
def _forward_with_last_hidden(
|
| 120 |
+
self,
|
| 121 |
+
input_ids: torch.Tensor,
|
| 122 |
+
*,
|
| 123 |
+
start_pos: int,
|
| 124 |
+
return_all_logits: bool,
|
| 125 |
+
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
| 126 |
+
hidden, _ = self.model._forward_hidden(input_ids, start_pos=start_pos)
|
| 127 |
+
last_hidden = hidden[:, -1, :]
|
| 128 |
+
logits_hidden = hidden if return_all_logits else hidden[:, -1:, :]
|
| 129 |
+
logits = self.model.head_mixer(
|
| 130 |
+
logits_hidden, norm=self.model.norm, output=self.model.output
|
| 131 |
+
)
|
| 132 |
+
if not return_all_logits:
|
| 133 |
+
logits = logits[:, 0, :]
|
| 134 |
+
return logits, last_hidden
|
| 135 |
+
|
| 136 |
+
def forward_with_last_hidden(
|
| 137 |
+
self,
|
| 138 |
+
input_ids: torch.Tensor,
|
| 139 |
+
*,
|
| 140 |
+
start_pos: int = 0,
|
| 141 |
+
return_all_logits: bool = True,
|
| 142 |
+
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
| 143 |
+
enable_runtime_cache(self.model)
|
| 144 |
+
return self._forward_with_last_hidden(
|
| 145 |
+
input_ids,
|
| 146 |
+
start_pos=int(start_pos),
|
| 147 |
+
return_all_logits=return_all_logits,
|
| 148 |
+
)
|
| 149 |
+
|
| 150 |
+
def replay_with_cache(
|
| 151 |
+
self,
|
| 152 |
+
input_ids: torch.Tensor,
|
| 153 |
+
*,
|
| 154 |
+
start_pos: int = 0,
|
| 155 |
+
return_all_logits: bool = True,
|
| 156 |
+
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
| 157 |
+
enable_runtime_cache(self.model)
|
| 158 |
+
if int(input_ids.size(1)) <= 1 or int(start_pos) == 0:
|
| 159 |
+
return self._forward_with_last_hidden(
|
| 160 |
+
input_ids,
|
| 161 |
+
start_pos=int(start_pos),
|
| 162 |
+
return_all_logits=return_all_logits,
|
| 163 |
+
)
|
| 164 |
+
outputs: list[torch.Tensor] = []
|
| 165 |
+
last_hidden = None
|
| 166 |
+
for offset in range(int(input_ids.size(1))):
|
| 167 |
+
logits, last_hidden = self._forward_with_last_hidden(
|
| 168 |
+
input_ids[:, offset : offset + 1],
|
| 169 |
+
start_pos=int(start_pos) + offset,
|
| 170 |
+
return_all_logits=True,
|
| 171 |
+
)
|
| 172 |
+
outputs.append(logits)
|
| 173 |
+
combined = torch.cat(outputs, dim=1)
|
| 174 |
+
return (combined if return_all_logits else combined[:, -1, :]), last_hidden
|
| 175 |
+
|
| 176 |
+
@staticmethod
|
| 177 |
+
def cache_batch_size(batch_size: int | None, max_batch_size: int) -> int:
|
| 178 |
+
active = int(max_batch_size if batch_size is None else batch_size)
|
| 179 |
+
if active <= 0:
|
| 180 |
+
raise ValueError(f"batch_size must be >= 1, got {active}")
|
| 181 |
+
return active
|
| 182 |
+
|
| 183 |
+
def cache_dump(
|
| 184 |
+
self,
|
| 185 |
+
*,
|
| 186 |
+
device: str = "cpu",
|
| 187 |
+
cache_pos: int | None = None,
|
| 188 |
+
batch_size: int | None = None,
|
| 189 |
+
) -> RuntimeCacheSnapshot:
|
| 190 |
+
active = self.cache_batch_size(batch_size, self.runtime_batch_capacity())
|
| 191 |
+
self.ensure_batch_capacity(active)
|
| 192 |
+
return RuntimeCacheSnapshot(
|
| 193 |
+
layers=tuple(
|
| 194 |
+
layer.attn.cache_snapshot(
|
| 195 |
+
device=device, batch_size=active, cache_pos=cache_pos
|
| 196 |
+
)
|
| 197 |
+
for layer in self.model.layers
|
| 198 |
+
)
|
| 199 |
+
)
|
| 200 |
+
|
| 201 |
+
def cache_load(self, cache_snapshot: RuntimeCacheSnapshot) -> None:
|
| 202 |
+
enable_runtime_cache(self.model)
|
| 203 |
+
if not isinstance(cache_snapshot, RuntimeCacheSnapshot):
|
| 204 |
+
raise TypeError("cache snapshot must be a RuntimeCacheSnapshot")
|
| 205 |
+
if len(cache_snapshot.layers) > len(self.model.layers):
|
| 206 |
+
raise ValueError("cache snapshot has more layers than the model")
|
| 207 |
+
required = cache_snapshot.batch_size()
|
| 208 |
+
if required:
|
| 209 |
+
self.ensure_batch_capacity(required)
|
| 210 |
+
for index, snapshot in enumerate(cache_snapshot.layers):
|
| 211 |
+
if snapshot is not None:
|
| 212 |
+
self.model.layers[index].attn.validate_cache_snapshot(snapshot)
|
| 213 |
+
self.reset_runtime_cache()
|
| 214 |
+
for index, snapshot in enumerate(cache_snapshot.layers):
|
| 215 |
+
if snapshot is not None:
|
| 216 |
+
self.model.layers[index].attn.load_cache_snapshot(snapshot)
|
| 217 |
+
|
| 218 |
+
|
| 219 |
+
__all__ = ["TransformerRuntime", "LayerCacheSnapshot"]
|
model_runtime_control.py
ADDED
|
@@ -0,0 +1,90 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
from collections.abc import Iterable, Sequence
|
| 8 |
+
from typing import Protocol
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
class _RuntimeArgs(Protocol):
|
| 12 |
+
use_cache: bool
|
| 13 |
+
def ensure_runtime_batch_capacity(self, batch_size: int) -> int: ...
|
| 14 |
+
def ensure_runtime_sequence_capacity(self, max_seq_len: int) -> int: ...
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
class _AttentionControl(Protocol):
|
| 18 |
+
use_cache: bool
|
| 19 |
+
def ensure_batch_capacity(self, batch_size: int) -> None: ...
|
| 20 |
+
def ensure_sequence_capacity(self, max_seq_len: int) -> None: ...
|
| 21 |
+
def rebuild_runtime_buffers(self, max_seq_len: int) -> None: ...
|
| 22 |
+
def reset(self) -> None: ...
|
| 23 |
+
def refresh_state_buffers(self) -> None: ...
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
class _RuntimeBlock(Protocol):
|
| 27 |
+
attn: _AttentionControl
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
class RuntimeControlModel(Protocol):
|
| 31 |
+
args: _RuntimeArgs
|
| 32 |
+
layers: Sequence[_RuntimeBlock]
|
| 33 |
+
def runtime_reserve_batch_capacity(self, batch_size: int) -> int: ...
|
| 34 |
+
def runtime_reserve_sequence_capacity(self, max_seq_len: int) -> int: ...
|
| 35 |
+
def runtime_sequence_capacity(self) -> int: ...
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
def iter_attn_modules(model: RuntimeControlModel) -> Iterable[_AttentionControl]:
|
| 39 |
+
for layer in model.layers:
|
| 40 |
+
yield layer.attn
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
def enable_runtime_cache(model: RuntimeControlModel) -> None:
|
| 44 |
+
model.args.use_cache = True
|
| 45 |
+
for attn in iter_attn_modules(model):
|
| 46 |
+
attn.use_cache = True
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
def ensure_batch_capacity(model: RuntimeControlModel, batch_size: int) -> None:
|
| 50 |
+
required = int(model.runtime_reserve_batch_capacity(batch_size))
|
| 51 |
+
for attn in iter_attn_modules(model):
|
| 52 |
+
attn.ensure_batch_capacity(required)
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
def ensure_sequence_capacity(model: RuntimeControlModel, max_seq_len: int) -> None:
|
| 56 |
+
required = int(max_seq_len)
|
| 57 |
+
if required <= 0:
|
| 58 |
+
raise ValueError(f"max_seq_len must be > 0, got {required}")
|
| 59 |
+
if required <= int(model.runtime_sequence_capacity()):
|
| 60 |
+
return
|
| 61 |
+
required = int(model.runtime_reserve_sequence_capacity(required))
|
| 62 |
+
for attn in iter_attn_modules(model):
|
| 63 |
+
attn.ensure_sequence_capacity(required)
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
def rebuild_runtime_buffers(model: RuntimeControlModel) -> None:
|
| 67 |
+
required = int(model.runtime_sequence_capacity())
|
| 68 |
+
for attn in iter_attn_modules(model):
|
| 69 |
+
attn.rebuild_runtime_buffers(required)
|
| 70 |
+
|
| 71 |
+
|
| 72 |
+
def reset_runtime_cache(model: RuntimeControlModel) -> None:
|
| 73 |
+
for attn in iter_attn_modules(model):
|
| 74 |
+
attn.reset()
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
def refresh_state_buffers(model: RuntimeControlModel) -> None:
|
| 78 |
+
for attn in iter_attn_modules(model):
|
| 79 |
+
attn.refresh_state_buffers()
|
| 80 |
+
|
| 81 |
+
|
| 82 |
+
__all__ = [
|
| 83 |
+
"RuntimeControlModel",
|
| 84 |
+
"enable_runtime_cache",
|
| 85 |
+
"ensure_batch_capacity",
|
| 86 |
+
"ensure_sequence_capacity",
|
| 87 |
+
"rebuild_runtime_buffers",
|
| 88 |
+
"refresh_state_buffers",
|
| 89 |
+
"reset_runtime_cache",
|
| 90 |
+
]
|
model_state.py
ADDED
|
@@ -0,0 +1,135 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
from collections.abc import Mapping
|
| 8 |
+
from dataclasses import dataclass, field
|
| 9 |
+
from threading import RLock
|
| 10 |
+
|
| 11 |
+
import torch
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
def _clone(value: torch.Tensor | None) -> torch.Tensor | None:
|
| 15 |
+
return None if value is None else value.detach().clone()
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
_CACHE_SUFFIXES = (
|
| 19 |
+
("_recurrent", "recurrent"),
|
| 20 |
+
("_conv", "conv"),
|
| 21 |
+
("_latent", "latent"),
|
| 22 |
+
)
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
@dataclass(frozen=True)
|
| 26 |
+
class LayerCacheSnapshot:
|
| 27 |
+
recurrent: torch.Tensor | None = None
|
| 28 |
+
conv: torch.Tensor | None = None
|
| 29 |
+
latent: torch.Tensor | None = None
|
| 30 |
+
|
| 31 |
+
def clone(self) -> LayerCacheSnapshot:
|
| 32 |
+
return LayerCacheSnapshot(
|
| 33 |
+
recurrent=_clone(self.recurrent),
|
| 34 |
+
conv=_clone(self.conv),
|
| 35 |
+
latent=_clone(self.latent),
|
| 36 |
+
)
|
| 37 |
+
|
| 38 |
+
def tensor_fields(self) -> tuple[tuple[str, torch.Tensor | None], ...]:
|
| 39 |
+
return (
|
| 40 |
+
("recurrent", self.recurrent),
|
| 41 |
+
("conv", self.conv),
|
| 42 |
+
("latent", self.latent),
|
| 43 |
+
)
|
| 44 |
+
|
| 45 |
+
def is_empty(self) -> bool:
|
| 46 |
+
return all(value is None for _name, value in self.tensor_fields())
|
| 47 |
+
|
| 48 |
+
def batch_size(self) -> int:
|
| 49 |
+
return max(
|
| 50 |
+
(int(value.size(0)) for _name, value in self.tensor_fields() if value is not None),
|
| 51 |
+
default=0,
|
| 52 |
+
)
|
| 53 |
+
|
| 54 |
+
def to_payload(self, *, prefix: str) -> dict[str, torch.Tensor]:
|
| 55 |
+
return {
|
| 56 |
+
f"{prefix}_{name}": value.detach().clone()
|
| 57 |
+
for name, value in self.tensor_fields()
|
| 58 |
+
if value is not None
|
| 59 |
+
}
|
| 60 |
+
|
| 61 |
+
@classmethod
|
| 62 |
+
def from_field_map(
|
| 63 |
+
cls,
|
| 64 |
+
field_map: Mapping[str, object] | None,
|
| 65 |
+
) -> LayerCacheSnapshot | None:
|
| 66 |
+
if field_map is None:
|
| 67 |
+
return None
|
| 68 |
+
values: dict[str, torch.Tensor | None] = {}
|
| 69 |
+
for name in ("recurrent", "conv", "latent"):
|
| 70 |
+
value = field_map.get(name)
|
| 71 |
+
if value is not None and not torch.is_tensor(value):
|
| 72 |
+
raise TypeError(f"layer cache field {name!r} must be a tensor or None")
|
| 73 |
+
values[name] = _clone(value)
|
| 74 |
+
snapshot = cls(**values)
|
| 75 |
+
return None if snapshot.is_empty() else snapshot
|
| 76 |
+
|
| 77 |
+
|
| 78 |
+
@dataclass(frozen=True)
|
| 79 |
+
class RuntimeCacheSnapshot:
|
| 80 |
+
layers: tuple[LayerCacheSnapshot | None, ...] = ()
|
| 81 |
+
|
| 82 |
+
def clone(self) -> RuntimeCacheSnapshot:
|
| 83 |
+
return RuntimeCacheSnapshot(
|
| 84 |
+
layers=tuple(None if item is None else item.clone() for item in self.layers)
|
| 85 |
+
)
|
| 86 |
+
|
| 87 |
+
def is_empty(self) -> bool:
|
| 88 |
+
return all(item is None or item.is_empty() for item in self.layers)
|
| 89 |
+
|
| 90 |
+
def batch_size(self) -> int:
|
| 91 |
+
return max((item.batch_size() for item in self.layers if item is not None), default=0)
|
| 92 |
+
|
| 93 |
+
def named_snapshots(self) -> tuple[tuple[str, LayerCacheSnapshot], ...]:
|
| 94 |
+
return tuple(
|
| 95 |
+
(f"layer_{index}", item)
|
| 96 |
+
for index, item in enumerate(self.layers)
|
| 97 |
+
if item is not None and not item.is_empty()
|
| 98 |
+
)
|
| 99 |
+
|
| 100 |
+
def to_payload(self) -> dict[str, torch.Tensor]:
|
| 101 |
+
payload: dict[str, torch.Tensor] = {}
|
| 102 |
+
for prefix, snapshot in self.named_snapshots():
|
| 103 |
+
payload.update(snapshot.to_payload(prefix=prefix))
|
| 104 |
+
return payload
|
| 105 |
+
|
| 106 |
+
@classmethod
|
| 107 |
+
def from_payload(cls, payload: Mapping[str, object]) -> RuntimeCacheSnapshot:
|
| 108 |
+
fields_by_layer: dict[int, dict[str, object]] = {}
|
| 109 |
+
for raw_key, value in payload.items():
|
| 110 |
+
key = str(raw_key)
|
| 111 |
+
for suffix, field_name in _CACHE_SUFFIXES:
|
| 112 |
+
if key.endswith(suffix):
|
| 113 |
+
prefix = key[: -len(suffix)]
|
| 114 |
+
if not prefix.startswith("layer_"):
|
| 115 |
+
break
|
| 116 |
+
index = int(prefix.removeprefix("layer_"))
|
| 117 |
+
fields_by_layer.setdefault(index, {})[field_name] = value
|
| 118 |
+
break
|
| 119 |
+
else:
|
| 120 |
+
raise ValueError(f"unrecognized runtime cache payload key: {key}")
|
| 121 |
+
if not fields_by_layer:
|
| 122 |
+
return cls()
|
| 123 |
+
layers: list[LayerCacheSnapshot | None] = [None] * (max(fields_by_layer) + 1)
|
| 124 |
+
for index, field_map in fields_by_layer.items():
|
| 125 |
+
layers[index] = LayerCacheSnapshot.from_field_map(field_map)
|
| 126 |
+
return cls(layers=tuple(layers))
|
| 127 |
+
|
| 128 |
+
|
| 129 |
+
@dataclass
|
| 130 |
+
class TransformerRuntimeState:
|
| 131 |
+
lock: RLock = field(default_factory=RLock)
|
| 132 |
+
loss_chunk_size: int = 0
|
| 133 |
+
|
| 134 |
+
|
| 135 |
+
__all__ = ["LayerCacheSnapshot", "RuntimeCacheSnapshot", "TransformerRuntimeState"]
|
model_transformer_setup.py
ADDED
|
@@ -0,0 +1,92 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
import math
|
| 8 |
+
from typing import Protocol
|
| 9 |
+
|
| 10 |
+
from torch import nn
|
| 11 |
+
|
| 12 |
+
from .model_attention import CausalDepthwiseConv1d
|
| 13 |
+
from .model_blocks import AttentionResidualMixer, SophiaBlock
|
| 14 |
+
from .model_config import ModelArgs
|
| 15 |
+
from .model_ops import RMSNorm, StandardLogitMixer
|
| 16 |
+
from .model_runtime import TransformerRuntime
|
| 17 |
+
from .runtime_linear import RuntimeLinear
|
| 18 |
+
from .model_state import TransformerRuntimeState
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
class _TransformerSetup(Protocol):
|
| 22 |
+
args: ModelArgs
|
| 23 |
+
runtime: TransformerRuntime
|
| 24 |
+
runtime_state: TransformerRuntimeState
|
| 25 |
+
tok_embeddings: nn.Embedding
|
| 26 |
+
layers: nn.ModuleList
|
| 27 |
+
norm: RMSNorm
|
| 28 |
+
output: RuntimeLinear
|
| 29 |
+
head_mixer: nn.Module
|
| 30 |
+
output_attn_residual: AttentionResidualMixer
|
| 31 |
+
dropout: nn.Dropout
|
| 32 |
+
gradient_checkpointing: bool
|
| 33 |
+
gradient_checkpointing_exclude_first: int
|
| 34 |
+
gradient_checkpointing_exclude_last: int
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
def tie_word_embeddings(transformer: _TransformerSetup) -> None:
|
| 38 |
+
if tuple(transformer.tok_embeddings.weight.shape) != tuple(transformer.output.weight.shape):
|
| 39 |
+
raise ValueError("tied embedding and output projection shapes must match")
|
| 40 |
+
transformer.output.weight = transformer.tok_embeddings.weight
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
def init_weights(model: nn.Module) -> None:
|
| 44 |
+
base_std = float(getattr(getattr(model, "args", None), "initializer_range", 0.02))
|
| 45 |
+
if not math.isfinite(base_std) or base_std <= 0.0:
|
| 46 |
+
raise ValueError(f"initializer_range must be finite and > 0, got {base_std}")
|
| 47 |
+
layer_count = max(int(getattr(getattr(model, "args", None), "n_layers", 1)), 1)
|
| 48 |
+
residual_std = base_std / math.sqrt(2.0 * layer_count)
|
| 49 |
+
for module in model.modules():
|
| 50 |
+
if isinstance(module, (nn.Linear, RuntimeLinear)):
|
| 51 |
+
nn.init.normal_(module.weight, mean=0.0, std=base_std)
|
| 52 |
+
if module.bias is not None:
|
| 53 |
+
nn.init.zeros_(module.bias)
|
| 54 |
+
elif isinstance(module, nn.Embedding):
|
| 55 |
+
nn.init.normal_(module.weight, mean=0.0, std=base_std)
|
| 56 |
+
elif isinstance(module, CausalDepthwiseConv1d):
|
| 57 |
+
nn.init.normal_(module.weight, mean=0.0, std=base_std)
|
| 58 |
+
|
| 59 |
+
for layer in getattr(model, "layers", ()):
|
| 60 |
+
nn.init.normal_(layer.attn.o_proj.weight, mean=0.0, std=residual_std)
|
| 61 |
+
nn.init.normal_(layer.ffn.down_proj.weight, mean=0.0, std=residual_std)
|
| 62 |
+
nn.init.zeros_(layer.attn_residual.query.weight)
|
| 63 |
+
nn.init.zeros_(layer.ffn_residual.query.weight)
|
| 64 |
+
output_attn_residual = getattr(model, "output_attn_residual", None)
|
| 65 |
+
if output_attn_residual is not None:
|
| 66 |
+
nn.init.zeros_(output_attn_residual.query.weight)
|
| 67 |
+
|
| 68 |
+
|
| 69 |
+
def initialize_transformer(transformer: _TransformerSetup, args: ModelArgs) -> None:
|
| 70 |
+
transformer.args = args
|
| 71 |
+
transformer.runtime = TransformerRuntime(
|
| 72 |
+
transformer,
|
| 73 |
+
state=TransformerRuntimeState(),
|
| 74 |
+
)
|
| 75 |
+
transformer.runtime_state = transformer.runtime.state
|
| 76 |
+
transformer.tok_embeddings = nn.Embedding(args.vocab_size, args.dim)
|
| 77 |
+
transformer.layers = nn.ModuleList(
|
| 78 |
+
[SophiaBlock(args, layer_idx) for layer_idx in range(args.n_layers)]
|
| 79 |
+
)
|
| 80 |
+
transformer.norm = RMSNorm(args.dim, args.norm_eps)
|
| 81 |
+
transformer.output_attn_residual = AttentionResidualMixer(args)
|
| 82 |
+
transformer.output = RuntimeLinear(args.dim, args.vocab_size, bias=False)
|
| 83 |
+
tie_word_embeddings(transformer)
|
| 84 |
+
transformer.head_mixer = StandardLogitMixer()
|
| 85 |
+
transformer.dropout = nn.Dropout(float(args.dropout))
|
| 86 |
+
transformer.gradient_checkpointing = False
|
| 87 |
+
transformer.gradient_checkpointing_exclude_first = 0
|
| 88 |
+
transformer.gradient_checkpointing_exclude_last = 0
|
| 89 |
+
init_weights(transformer)
|
| 90 |
+
|
| 91 |
+
|
| 92 |
+
__all__ = ["init_weights", "initialize_transformer", "tie_word_embeddings"]
|
modeling_sophia.py
ADDED
|
@@ -0,0 +1,149 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
"""
|
| 6 |
+
Sophia model implementation.
|
| 7 |
+
|
| 8 |
+
This file contains:
|
| 9 |
+
- `SophiaForCausalLM`: Hugging Face causal language model adapter over the Sophia runtime.
|
| 10 |
+
"""
|
| 11 |
+
|
| 12 |
+
from __future__ import annotations
|
| 13 |
+
|
| 14 |
+
from typing import TYPE_CHECKING
|
| 15 |
+
|
| 16 |
+
from .sophia_runtime import Transformer as _SophiaRuntimeTransformer # noqa: F401
|
| 17 |
+
|
| 18 |
+
if TYPE_CHECKING:
|
| 19 |
+
from .sophia_runtime import * # noqa: F401,F403
|
| 20 |
+
from .model_config import * # noqa: F401,F403
|
| 21 |
+
from .hf_remote_code import * # noqa: F401,F403
|
| 22 |
+
from .hf_support import * # noqa: F401,F403
|
| 23 |
+
from .hf_cache import * # noqa: F401,F403
|
| 24 |
+
from .hf_config import * # noqa: F401,F403
|
| 25 |
+
from .hf_generation import * # noqa: F401,F403
|
| 26 |
+
from .hf_lifecycle import * # noqa: F401,F403
|
| 27 |
+
from .model_attention import * # noqa: F401,F403
|
| 28 |
+
from .model_blocks import * # noqa: F401,F403
|
| 29 |
+
from .model_ops import * # noqa: F401,F403
|
| 30 |
+
from .model_runtime import * # noqa: F401,F403
|
| 31 |
+
from .runtime_contracts import * # noqa: F401,F403
|
| 32 |
+
from .model_runtime_control import * # noqa: F401,F403
|
| 33 |
+
from .model_state import * # noqa: F401,F403
|
| 34 |
+
from .model_transformer_setup import * # noqa: F401,F403
|
| 35 |
+
from .runtime_linear import * # noqa: F401,F403
|
| 36 |
+
from .canonical_config import * # noqa: F401,F403
|
| 37 |
+
from .model_dir import * # noqa: F401,F403
|
| 38 |
+
from .loss_stats import * # noqa: F401,F403
|
| 39 |
+
from .input_mask import * # noqa: F401,F403
|
| 40 |
+
from .cache_decode import * # noqa: F401,F403
|
| 41 |
+
from .config_projection import * # noqa: F401,F403
|
| 42 |
+
from .hf_projection import * # noqa: F401,F403
|
| 43 |
+
from .runtime_backend import * # noqa: F401,F403
|
| 44 |
+
from .decoder_forward import * # noqa: F401,F403
|
| 45 |
+
from .decoder_types import * # noqa: F401,F403
|
| 46 |
+
from .decoder_loss import * # noqa: F401,F403
|
| 47 |
+
from .decoder_loss_forward import * # noqa: F401,F403
|
| 48 |
+
from .decoder_full import * # noqa: F401,F403
|
| 49 |
+
from .pretrained_bundle import * # noqa: F401,F403
|
| 50 |
+
from .decoder_output import * # noqa: F401,F403
|
| 51 |
+
from .decoder_runtime import * # noqa: F401,F403
|
| 52 |
+
from .decoder_host import * # noqa: F401,F403
|
| 53 |
+
from .sophia_decoder import * # noqa: F401,F403
|
| 54 |
+
from .semantics import * # noqa: F401,F403
|
| 55 |
+
|
| 56 |
+
import torch
|
| 57 |
+
from transformers import PreTrainedModel
|
| 58 |
+
from transformers.generation import GenerationMixin
|
| 59 |
+
from transformers.modeling_outputs import CausalLMOutputWithPast
|
| 60 |
+
|
| 61 |
+
from .runtime_backend import resolve_runtime_backend
|
| 62 |
+
from .hf_cache import SophiaCache
|
| 63 |
+
from .hf_generation import (
|
| 64 |
+
CausalLMForwardMixin,
|
| 65 |
+
GenerationCacheMixin,
|
| 66 |
+
)
|
| 67 |
+
from .hf_lifecycle import PreTrainedLifecycleMixin
|
| 68 |
+
from .hf_projection import SophiaConfig
|
| 69 |
+
from .decoder_output import (
|
| 70 |
+
format_hf_causal_lm_output as _format_hf_causal_lm_output,
|
| 71 |
+
resolve_return_dict as _resolve_return_dict,
|
| 72 |
+
)
|
| 73 |
+
from .decoder_host import DecoderHostMixin
|
| 74 |
+
from .decoder_runtime import (
|
| 75 |
+
DecoderModelMixin,
|
| 76 |
+
DecoderRecipeMixin,
|
| 77 |
+
DecoderRuntimeMixin,
|
| 78 |
+
)
|
| 79 |
+
|
| 80 |
+
|
| 81 |
+
class SophiaForCausalLM(
|
| 82 |
+
DecoderModelMixin,
|
| 83 |
+
DecoderRuntimeMixin,
|
| 84 |
+
DecoderRecipeMixin,
|
| 85 |
+
DecoderHostMixin,
|
| 86 |
+
GenerationCacheMixin,
|
| 87 |
+
CausalLMForwardMixin,
|
| 88 |
+
PreTrainedLifecycleMixin,
|
| 89 |
+
PreTrainedModel,
|
| 90 |
+
GenerationMixin,
|
| 91 |
+
):
|
| 92 |
+
config_class = SophiaConfig
|
| 93 |
+
base_model_prefix = "model"
|
| 94 |
+
_tied_weights_keys = {"model.output.weight": "model.tok_embeddings.weight"}
|
| 95 |
+
supports_gradient_checkpointing = True
|
| 96 |
+
_is_stateful = True
|
| 97 |
+
|
| 98 |
+
@classmethod
|
| 99 |
+
def _supports_default_dynamic_cache(cls) -> bool:
|
| 100 |
+
return False
|
| 101 |
+
|
| 102 |
+
def __init__(
|
| 103 |
+
self,
|
| 104 |
+
config: SophiaConfig | None = None,
|
| 105 |
+
*,
|
| 106 |
+
runtime_max_seq_len: int | None = None,
|
| 107 |
+
):
|
| 108 |
+
config = config or SophiaConfig()
|
| 109 |
+
super().__init__(config)
|
| 110 |
+
self._initialize_decoder_runtime(
|
| 111 |
+
config=config,
|
| 112 |
+
runtime_max_seq_len=runtime_max_seq_len,
|
| 113 |
+
runtime_backend=resolve_runtime_backend(),
|
| 114 |
+
gradient_checkpointing_enabled=bool(
|
| 115 |
+
getattr(self, "gradient_checkpointing", False)
|
| 116 |
+
),
|
| 117 |
+
register_tied_weights=True,
|
| 118 |
+
)
|
| 119 |
+
|
| 120 |
+
def _set_gradient_checkpointing(
|
| 121 |
+
self,
|
| 122 |
+
enable: bool = True,
|
| 123 |
+
gradient_checkpointing_func: object = None,
|
| 124 |
+
) -> None:
|
| 125 |
+
del gradient_checkpointing_func
|
| 126 |
+
self.gradient_checkpointing = bool(enable)
|
| 127 |
+
self._sync_runtime_gradient_checkpointing(enable=bool(enable))
|
| 128 |
+
|
| 129 |
+
def _resolve_runtime_return_dict(self, *, return_dict: bool | None) -> bool:
|
| 130 |
+
return _resolve_return_dict(
|
| 131 |
+
config=self.config,
|
| 132 |
+
return_dict=return_dict,
|
| 133 |
+
)
|
| 134 |
+
|
| 135 |
+
def _format_decoder_output(
|
| 136 |
+
self,
|
| 137 |
+
*,
|
| 138 |
+
loss: torch.Tensor | None,
|
| 139 |
+
logits: torch.Tensor | None,
|
| 140 |
+
cache: SophiaCache | None,
|
| 141 |
+
return_dict: bool,
|
| 142 |
+
) -> CausalLMOutputWithPast | tuple[torch.Tensor, ...]:
|
| 143 |
+
return _format_hf_causal_lm_output(
|
| 144 |
+
output_cls=CausalLMOutputWithPast,
|
| 145 |
+
loss=loss,
|
| 146 |
+
logits=logits,
|
| 147 |
+
past_key_values=cache,
|
| 148 |
+
return_dict=return_dict,
|
| 149 |
+
)
|
pretrained_bundle.py
ADDED
|
@@ -0,0 +1,118 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
from dataclasses import is_dataclass
|
| 8 |
+
from dataclasses import asdict
|
| 9 |
+
import json
|
| 10 |
+
import os
|
| 11 |
+
|
| 12 |
+
import torch
|
| 13 |
+
from safetensors.torch import load_file as load_safetensors_file
|
| 14 |
+
from safetensors.torch import save_file as save_safetensors_file
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
def save_canonical_config_bundle(
|
| 18 |
+
config: object,
|
| 19 |
+
*,
|
| 20 |
+
save_directory: str | os.PathLike,
|
| 21 |
+
) -> None:
|
| 22 |
+
save_dir = os.path.abspath(str(save_directory))
|
| 23 |
+
os.makedirs(save_dir, exist_ok=True)
|
| 24 |
+
payload = asdict(config) if is_dataclass(config) else dict(vars(config))
|
| 25 |
+
config_path = os.path.join(save_dir, "config.json")
|
| 26 |
+
with open(config_path, "w", encoding="utf-8", newline="\n") as handle:
|
| 27 |
+
json.dump(payload, handle, indent=2, sort_keys=True)
|
| 28 |
+
handle.write("\n")
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
def save_pretrained_state_dict(
|
| 32 |
+
state_dict: dict[str, torch.Tensor],
|
| 33 |
+
*,
|
| 34 |
+
save_directory: str | os.PathLike,
|
| 35 |
+
safe_serialization: bool,
|
| 36 |
+
) -> None:
|
| 37 |
+
save_dir = os.path.abspath(str(save_directory))
|
| 38 |
+
os.makedirs(save_dir, exist_ok=True)
|
| 39 |
+
export_state: dict[str, torch.Tensor] = {}
|
| 40 |
+
seen_cpu_storages: set[tuple[int, int]] = set()
|
| 41 |
+
for key, value in dict(state_dict).items():
|
| 42 |
+
tensor = value.detach().cpu()
|
| 43 |
+
storage = tensor.untyped_storage()
|
| 44 |
+
storage_key = (int(storage.data_ptr()), int(storage.nbytes()))
|
| 45 |
+
if storage_key in seen_cpu_storages:
|
| 46 |
+
tensor = tensor.clone(memory_format=torch.preserve_format)
|
| 47 |
+
storage = tensor.untyped_storage()
|
| 48 |
+
storage_key = (int(storage.data_ptr()), int(storage.nbytes()))
|
| 49 |
+
seen_cpu_storages.add(storage_key)
|
| 50 |
+
export_state[str(key)] = tensor
|
| 51 |
+
if bool(safe_serialization):
|
| 52 |
+
save_safetensors_file(
|
| 53 |
+
export_state,
|
| 54 |
+
os.path.join(save_dir, "model.safetensors"),
|
| 55 |
+
)
|
| 56 |
+
bin_path = os.path.join(save_dir, "pytorch_model.bin")
|
| 57 |
+
if os.path.exists(bin_path):
|
| 58 |
+
os.remove(bin_path)
|
| 59 |
+
return
|
| 60 |
+
torch.save(export_state, os.path.join(save_dir, "pytorch_model.bin"))
|
| 61 |
+
|
| 62 |
+
|
| 63 |
+
def load_pretrained_state_dict(
|
| 64 |
+
export_dir: str | os.PathLike,
|
| 65 |
+
) -> dict[str, torch.Tensor]:
|
| 66 |
+
resolved_export_dir = os.path.abspath(str(export_dir))
|
| 67 |
+
safetensors_path = os.path.join(resolved_export_dir, "model.safetensors")
|
| 68 |
+
pytorch_path = os.path.join(resolved_export_dir, "pytorch_model.bin")
|
| 69 |
+
if os.path.isfile(safetensors_path):
|
| 70 |
+
return load_safetensors_file(safetensors_path)
|
| 71 |
+
if os.path.isfile(pytorch_path):
|
| 72 |
+
raw_state = torch.load(pytorch_path, map_location="cpu", weights_only=True)
|
| 73 |
+
if not isinstance(raw_state, dict):
|
| 74 |
+
raise RuntimeError(
|
| 75 |
+
f"unexpected state_dict payload type: {type(raw_state)!r}"
|
| 76 |
+
)
|
| 77 |
+
return raw_state
|
| 78 |
+
raise RuntimeError(
|
| 79 |
+
"exported model directory has no supported model weights: "
|
| 80 |
+
f"{resolved_export_dir}"
|
| 81 |
+
)
|
| 82 |
+
|
| 83 |
+
|
| 84 |
+
def load_canonical_config_from_pretrained[ConfigT](
|
| 85 |
+
pretrained_model_name_or_path: str | os.PathLike,
|
| 86 |
+
*,
|
| 87 |
+
config_cls: type[ConfigT],
|
| 88 |
+
) -> ConfigT:
|
| 89 |
+
config_path = os.path.join(
|
| 90 |
+
os.path.abspath(str(pretrained_model_name_or_path)),
|
| 91 |
+
"config.json",
|
| 92 |
+
)
|
| 93 |
+
try:
|
| 94 |
+
with open(config_path, encoding="utf-8") as handle:
|
| 95 |
+
raw = json.load(handle)
|
| 96 |
+
except FileNotFoundError as exc:
|
| 97 |
+
raise RuntimeError(f"missing exported config.json: {config_path}") from exc
|
| 98 |
+
except json.JSONDecodeError as exc:
|
| 99 |
+
raise RuntimeError(f"invalid exported config.json: {config_path}") from exc
|
| 100 |
+
if not isinstance(raw, dict):
|
| 101 |
+
raise RuntimeError(
|
| 102 |
+
f"exported config.json must contain a JSON object: {config_path}"
|
| 103 |
+
)
|
| 104 |
+
from_pretrained_mapping = getattr(config_cls, "from_pretrained_mapping", None)
|
| 105 |
+
if callable(from_pretrained_mapping):
|
| 106 |
+
return from_pretrained_mapping(dict(raw))
|
| 107 |
+
from_mapping = getattr(config_cls, "from_mapping", None)
|
| 108 |
+
if callable(from_mapping):
|
| 109 |
+
return from_mapping(dict(raw))
|
| 110 |
+
return config_cls(**dict(raw))
|
| 111 |
+
|
| 112 |
+
|
| 113 |
+
__all__ = [
|
| 114 |
+
"load_canonical_config_from_pretrained",
|
| 115 |
+
"load_pretrained_state_dict",
|
| 116 |
+
"save_canonical_config_bundle",
|
| 117 |
+
"save_pretrained_state_dict",
|
| 118 |
+
]
|
runtime_backend.py
ADDED
|
@@ -0,0 +1,53 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
from dataclasses import dataclass
|
| 8 |
+
from importlib import import_module
|
| 9 |
+
from typing import TYPE_CHECKING, TypeVar
|
| 10 |
+
|
| 11 |
+
from .runtime_contracts import RuntimeBackedModel
|
| 12 |
+
|
| 13 |
+
if TYPE_CHECKING:
|
| 14 |
+
from types import ModuleType
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
RUNTIME_MODULE_NAME = ".sophia_runtime"
|
| 18 |
+
RuntimeModelArgsT = TypeVar("RuntimeModelArgsT")
|
| 19 |
+
RuntimeTransformerT = TypeVar("RuntimeTransformerT", bound=RuntimeBackedModel)
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
@dataclass(frozen=True)
|
| 23 |
+
class RuntimeBackend:
|
| 24 |
+
module: ModuleType | None
|
| 25 |
+
model_args_cls: type[RuntimeModelArgsT]
|
| 26 |
+
transformer_cls: type[RuntimeTransformerT]
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
def _import_runtime_module(module_name: str) -> ModuleType:
|
| 30 |
+
if module_name.startswith("."):
|
| 31 |
+
package = __package__
|
| 32 |
+
if not package:
|
| 33 |
+
raise ImportError(
|
| 34 |
+
"relative runtime module import requires a package context"
|
| 35 |
+
)
|
| 36 |
+
return import_module(module_name, package=package)
|
| 37 |
+
return import_module(module_name)
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
def resolve_runtime_backend() -> RuntimeBackend:
|
| 41 |
+
runtime_module = _import_runtime_module(RUNTIME_MODULE_NAME)
|
| 42 |
+
return RuntimeBackend(
|
| 43 |
+
module=runtime_module,
|
| 44 |
+
model_args_cls=runtime_module.ModelArgs,
|
| 45 |
+
transformer_cls=runtime_module.Transformer,
|
| 46 |
+
)
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
__all__ = [
|
| 50 |
+
"RUNTIME_MODULE_NAME",
|
| 51 |
+
"RuntimeBackend",
|
| 52 |
+
"resolve_runtime_backend",
|
| 53 |
+
]
|
runtime_contracts.py
ADDED
|
@@ -0,0 +1,107 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
from threading import RLock
|
| 8 |
+
from typing import Protocol, runtime_checkable
|
| 9 |
+
|
| 10 |
+
import torch
|
| 11 |
+
|
| 12 |
+
from .model_state import RuntimeCacheSnapshot
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
@runtime_checkable
|
| 16 |
+
class SupportsRuntimeStateRefreshControl(Protocol):
|
| 17 |
+
def refresh_state_buffers(self) -> None: ...
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
@runtime_checkable
|
| 21 |
+
class SupportsRuntimeStateResetControl(Protocol):
|
| 22 |
+
def reset_runtime_cache(self) -> None: ...
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
@runtime_checkable
|
| 26 |
+
class SupportsRuntimeRecipeControl(Protocol):
|
| 27 |
+
def supports_loss_chunk_size(self) -> bool: ...
|
| 28 |
+
|
| 29 |
+
def supports_checkpoint_excludes(self) -> bool: ...
|
| 30 |
+
|
| 31 |
+
def runtime_recipe_knobs(self) -> tuple[int, int, int]: ...
|
| 32 |
+
|
| 33 |
+
def apply_runtime_recipe_knobs(
|
| 34 |
+
self,
|
| 35 |
+
*,
|
| 36 |
+
loss_chunk_size: int,
|
| 37 |
+
gradient_checkpointing_exclude_first: int,
|
| 38 |
+
gradient_checkpointing_exclude_last: int,
|
| 39 |
+
) -> tuple[int, int, int]: ...
|
| 40 |
+
|
| 41 |
+
def ensure_runtime_max_seq_len(self, max_seq_len: int) -> None: ...
|
| 42 |
+
|
| 43 |
+
def runtime_max_seq_len(self) -> int: ...
|
| 44 |
+
|
| 45 |
+
def sync_runtime_batch_capacity(self, max_batch_size: int) -> int: ...
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
@runtime_checkable
|
| 49 |
+
class RuntimeHost(
|
| 50 |
+
SupportsRuntimeStateRefreshControl,
|
| 51 |
+
SupportsRuntimeStateResetControl,
|
| 52 |
+
SupportsRuntimeRecipeControl,
|
| 53 |
+
Protocol,
|
| 54 |
+
):
|
| 55 |
+
def rebuild_runtime_buffers(self) -> None: ...
|
| 56 |
+
|
| 57 |
+
def replay_with_cache(
|
| 58 |
+
self,
|
| 59 |
+
input_ids: torch.Tensor,
|
| 60 |
+
start_pos: int = 0,
|
| 61 |
+
return_all_logits: bool = True,
|
| 62 |
+
) -> tuple[torch.Tensor, torch.Tensor | None]: ...
|
| 63 |
+
|
| 64 |
+
def forward_with_last_hidden(
|
| 65 |
+
self,
|
| 66 |
+
input_ids: torch.Tensor,
|
| 67 |
+
start_pos: int = 0,
|
| 68 |
+
return_all_logits: bool = True,
|
| 69 |
+
) -> tuple[torch.Tensor, torch.Tensor | None]: ...
|
| 70 |
+
|
| 71 |
+
def cache_dump(
|
| 72 |
+
self,
|
| 73 |
+
device: str = "cpu",
|
| 74 |
+
*,
|
| 75 |
+
cache_pos: int | None = None,
|
| 76 |
+
batch_size: int | None = None,
|
| 77 |
+
) -> RuntimeCacheSnapshot: ...
|
| 78 |
+
|
| 79 |
+
def cache_load(self, cache_snapshot: RuntimeCacheSnapshot) -> None: ...
|
| 80 |
+
|
| 81 |
+
|
| 82 |
+
@runtime_checkable
|
| 83 |
+
class RuntimeBackedModel(Protocol):
|
| 84 |
+
runtime: RuntimeHost
|
| 85 |
+
runtime_lock: RLock
|
| 86 |
+
|
| 87 |
+
|
| 88 |
+
@runtime_checkable
|
| 89 |
+
class Runtime(Protocol):
|
| 90 |
+
@property
|
| 91 |
+
def runtime_model(self) -> RuntimeBackedModel: ...
|
| 92 |
+
|
| 93 |
+
@property
|
| 94 |
+
def runtime_host(self) -> RuntimeHost: ...
|
| 95 |
+
|
| 96 |
+
@property
|
| 97 |
+
def runtime_lock(self) -> RLock: ...
|
| 98 |
+
|
| 99 |
+
|
| 100 |
+
__all__ = [
|
| 101 |
+
"RuntimeBackedModel",
|
| 102 |
+
"RuntimeHost",
|
| 103 |
+
"Runtime",
|
| 104 |
+
"SupportsRuntimeRecipeControl",
|
| 105 |
+
"SupportsRuntimeStateRefreshControl",
|
| 106 |
+
"SupportsRuntimeStateResetControl",
|
| 107 |
+
]
|
runtime_linear.py
ADDED
|
@@ -0,0 +1,46 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
"""Runtime linear layers used by the model."""
|
| 6 |
+
|
| 7 |
+
from __future__ import annotations
|
| 8 |
+
|
| 9 |
+
import torch
|
| 10 |
+
import torch.nn.functional as functional
|
| 11 |
+
from torch import nn
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
class RuntimeLinear(nn.Module):
|
| 15 |
+
"""Plain linear layer used across the runtime model."""
|
| 16 |
+
|
| 17 |
+
def __init__(
|
| 18 |
+
self,
|
| 19 |
+
in_features: int,
|
| 20 |
+
out_features: int,
|
| 21 |
+
*,
|
| 22 |
+
bias: bool = False,
|
| 23 |
+
) -> None:
|
| 24 |
+
super().__init__()
|
| 25 |
+
self.in_features = int(in_features)
|
| 26 |
+
self.out_features = int(out_features)
|
| 27 |
+
self.weight = nn.Parameter(torch.empty(self.out_features, self.in_features))
|
| 28 |
+
self.bias = nn.Parameter(torch.empty(self.out_features)) if bool(bias) else None
|
| 29 |
+
self.reset_parameters()
|
| 30 |
+
|
| 31 |
+
def reset_parameters(self) -> None:
|
| 32 |
+
nn.init.normal_(self.weight, mean=0.0, std=0.02)
|
| 33 |
+
if self.bias is not None:
|
| 34 |
+
nn.init.zeros_(self.bias)
|
| 35 |
+
|
| 36 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 37 |
+
return functional.linear(x, self.weight, self.bias)
|
| 38 |
+
|
| 39 |
+
def extra_repr(self) -> str:
|
| 40 |
+
return (
|
| 41 |
+
f"in_features={self.in_features}, out_features={self.out_features}, "
|
| 42 |
+
f"bias={self.bias is not None}"
|
| 43 |
+
)
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
__all__ = ["RuntimeLinear"]
|
semantics.py
ADDED
|
@@ -0,0 +1,135 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
|
| 2 |
+
# Exported for HuggingFace trust_remote_code loading.
|
| 3 |
+
# This file is intentionally self-contained.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
from collections.abc import Mapping
|
| 8 |
+
import math
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
PINNED_SEMANTIC_ADAM_EPS = 1e-8
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
def canonicalize_model_values(values: Mapping[str, object]) -> dict[str, object]:
|
| 15 |
+
normalized: dict[str, object] = dict(values)
|
| 16 |
+
integer_fields = (
|
| 17 |
+
"vocab_size",
|
| 18 |
+
"dim",
|
| 19 |
+
"n_layers",
|
| 20 |
+
"num_heads",
|
| 21 |
+
"head_dim",
|
| 22 |
+
"ffn_hidden",
|
| 23 |
+
"kda_decay_rank",
|
| 24 |
+
"kda_output_gate_rank",
|
| 25 |
+
"mla_q_rank",
|
| 26 |
+
"mla_kv_rank",
|
| 27 |
+
"short_conv_kernel",
|
| 28 |
+
"attn_res_block_size",
|
| 29 |
+
"max_seq_len",
|
| 30 |
+
"max_batch_size",
|
| 31 |
+
)
|
| 32 |
+
float_fields = (
|
| 33 |
+
"kda_decay_lower_bound",
|
| 34 |
+
"kda_dt_min",
|
| 35 |
+
"kda_dt_max",
|
| 36 |
+
"kda_dt_floor",
|
| 37 |
+
"kda_a_log_init",
|
| 38 |
+
"norm_eps",
|
| 39 |
+
"dropout",
|
| 40 |
+
"initializer_range",
|
| 41 |
+
"situ_gate_softcap",
|
| 42 |
+
"situ_up_softcap",
|
| 43 |
+
)
|
| 44 |
+
for name in integer_fields:
|
| 45 |
+
if name in normalized:
|
| 46 |
+
normalized[name] = int(normalized[name])
|
| 47 |
+
for name in float_fields:
|
| 48 |
+
if name in normalized:
|
| 49 |
+
normalized[name] = float(normalized[name])
|
| 50 |
+
if "kda_backend" in normalized:
|
| 51 |
+
normalized["kda_backend"] = str(normalized["kda_backend"])
|
| 52 |
+
if "kda_output_gate_full_rank" in normalized:
|
| 53 |
+
normalized["kda_output_gate_full_rank"] = bool(
|
| 54 |
+
normalized["kda_output_gate_full_rank"]
|
| 55 |
+
)
|
| 56 |
+
return normalized
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
def validate_model_values(values: Mapping[str, object]) -> None:
|
| 60 |
+
positive_ints = (
|
| 61 |
+
"vocab_size",
|
| 62 |
+
"dim",
|
| 63 |
+
"n_layers",
|
| 64 |
+
"num_heads",
|
| 65 |
+
"head_dim",
|
| 66 |
+
"ffn_hidden",
|
| 67 |
+
"kda_decay_rank",
|
| 68 |
+
"kda_output_gate_rank",
|
| 69 |
+
"mla_q_rank",
|
| 70 |
+
"mla_kv_rank",
|
| 71 |
+
"short_conv_kernel",
|
| 72 |
+
"attn_res_block_size",
|
| 73 |
+
"max_seq_len",
|
| 74 |
+
"max_batch_size",
|
| 75 |
+
)
|
| 76 |
+
for name in positive_ints:
|
| 77 |
+
value = int(values[name])
|
| 78 |
+
if value <= 0:
|
| 79 |
+
raise ValueError(f"{name} must be > 0, got {value}")
|
| 80 |
+
|
| 81 |
+
lower_bound = float(values["kda_decay_lower_bound"])
|
| 82 |
+
if not -5.0 <= lower_bound < 0.0:
|
| 83 |
+
raise ValueError(
|
| 84 |
+
"kda_decay_lower_bound must be in [-5, 0), "
|
| 85 |
+
f"got {lower_bound}"
|
| 86 |
+
)
|
| 87 |
+
dt_min = float(values["kda_dt_min"])
|
| 88 |
+
dt_max = float(values["kda_dt_max"])
|
| 89 |
+
dt_floor = float(values["kda_dt_floor"])
|
| 90 |
+
a_log_init = float(values["kda_a_log_init"])
|
| 91 |
+
for name, value in (
|
| 92 |
+
("kda_dt_min", dt_min),
|
| 93 |
+
("kda_dt_max", dt_max),
|
| 94 |
+
("kda_dt_floor", dt_floor),
|
| 95 |
+
):
|
| 96 |
+
if not math.isfinite(value) or value <= 0.0:
|
| 97 |
+
raise ValueError(f"{name} must be finite and > 0, got {value}")
|
| 98 |
+
if dt_min > dt_max:
|
| 99 |
+
raise ValueError(
|
| 100 |
+
f"kda_dt_min must be <= kda_dt_max, got {dt_min} > {dt_max}"
|
| 101 |
+
)
|
| 102 |
+
if dt_floor > dt_min:
|
| 103 |
+
raise ValueError(
|
| 104 |
+
f"kda_dt_floor must be <= kda_dt_min, got {dt_floor} > {dt_min}"
|
| 105 |
+
)
|
| 106 |
+
if not math.isfinite(a_log_init):
|
| 107 |
+
raise ValueError(f"kda_a_log_init must be finite, got {a_log_init}")
|
| 108 |
+
norm_eps = float(values["norm_eps"])
|
| 109 |
+
if not math.isfinite(norm_eps) or norm_eps <= 0.0:
|
| 110 |
+
raise ValueError(f"norm_eps must be finite and > 0, got {norm_eps}")
|
| 111 |
+
dropout = float(values["dropout"])
|
| 112 |
+
if not 0.0 <= dropout < 1.0:
|
| 113 |
+
raise ValueError(f"dropout must be in [0, 1), got {dropout}")
|
| 114 |
+
initializer_range = float(values["initializer_range"])
|
| 115 |
+
if not math.isfinite(initializer_range) or initializer_range <= 0.0:
|
| 116 |
+
raise ValueError(
|
| 117 |
+
f"initializer_range must be finite and > 0, got {initializer_range}"
|
| 118 |
+
)
|
| 119 |
+
for name in ("situ_gate_softcap", "situ_up_softcap"):
|
| 120 |
+
value = float(values[name])
|
| 121 |
+
if not math.isfinite(value) or value <= 0.0:
|
| 122 |
+
raise ValueError(f"{name} must be finite and > 0, got {value}")
|
| 123 |
+
backend = str(values.get("kda_backend", "auto"))
|
| 124 |
+
if backend not in {"auto", "reference", "fla"}:
|
| 125 |
+
raise ValueError(
|
| 126 |
+
"kda_backend must be one of auto/reference/fla, "
|
| 127 |
+
f"got {backend!r}"
|
| 128 |
+
)
|
| 129 |
+
|
| 130 |
+
|
| 131 |
+
__all__ = [
|
| 132 |
+
"PINNED_SEMANTIC_ADAM_EPS",
|
| 133 |
+
"canonicalize_model_values",
|
| 134 |
+
"validate_model_values",
|
| 135 |
+
]
|