Upload folder using huggingface_hub
Browse files- .gitattributes +1 -0
- README.md +195 -0
- chat_template.jinja +89 -0
- config.json +63 -0
- generation_config.json +13 -0
- model.safetensors +3 -0
- projection_head.pt +3 -0
- tokenizer.json +3 -0
- tokenizer_config.json +29 -0
- training_config.json +25 -0
.gitattributes
CHANGED
|
@@ -33,3 +33,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
|
| 33 |
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
|
|
|
|
|
| 33 |
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
| 36 |
+
tokenizer.json filter=lfs diff=lfs merge=lfs -text
|
README.md
ADDED
|
@@ -0,0 +1,195 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
language:
|
| 3 |
+
- en
|
| 4 |
+
license: apache-2.0
|
| 5 |
+
base_model: Qwen/Qwen3-1.7B
|
| 6 |
+
tags:
|
| 7 |
+
- continuous-thought
|
| 8 |
+
- latent-reasoning
|
| 9 |
+
- distillation
|
| 10 |
+
- gsm8k
|
| 11 |
+
- codi
|
| 12 |
+
datasets:
|
| 13 |
+
- openai/gsm8k
|
| 14 |
+
metrics:
|
| 15 |
+
- accuracy
|
| 16 |
+
model-index:
|
| 17 |
+
- name: r11_single_perstep_g1
|
| 18 |
+
results:
|
| 19 |
+
- task:
|
| 20 |
+
type: math-reasoning
|
| 21 |
+
name: Grade School Math
|
| 22 |
+
dataset:
|
| 23 |
+
name: GSM8k
|
| 24 |
+
type: openai/gsm8k
|
| 25 |
+
split: test
|
| 26 |
+
metrics:
|
| 27 |
+
- type: accuracy
|
| 28 |
+
value: 80.4
|
| 29 |
+
name: Accuracy
|
| 30 |
+
---
|
| 31 |
+
|
| 32 |
+
# r11_single_perstep_g1
|
| 33 |
+
|
| 34 |
+
**Single-trace per-step baseline (80.4%)**
|
| 35 |
+
|
| 36 |
+
- Single best rollout teacher (no multi-rollout aggregation)
|
| 37 |
+
- Per-step distillation at every latent step with γ=1.0
|
| 38 |
+
- Baseline for measuring multi-rollout benefit
|
| 39 |
+
|
| 40 |
+
## Overview
|
| 41 |
+
|
| 42 |
+
This model implements **Continuous Thought Distillation (CODI)** — an autoregressive latent
|
| 43 |
+
reasoning loop that processes K=6 continuous thought steps before generating a text answer.
|
| 44 |
+
Teacher hidden states are extracted from multiple chain-of-thought rollouts generated by the
|
| 45 |
+
base model and distilled into the latent representations.
|
| 46 |
+
|
| 47 |
+
**Key idea**: Instead of distilling from a single reasoning trace, we aggregate hidden states
|
| 48 |
+
from multiple rollouts (16 per problem). The teacher signal is the hidden state from the single best rollout.
|
| 49 |
+
|
| 50 |
+
## Architecture
|
| 51 |
+
|
| 52 |
+
- **Base model**: Qwen3-1.7B (1.7B parameters)
|
| 53 |
+
- **Fine-tuning**: LoRA (rank=32, alpha=16) on q/k/v/o/gate/up/down_proj
|
| 54 |
+
- **Projection head**: Linear(2048, 2048) → GELU → Linear(2048, 2048) → LayerNorm(2048)
|
| 55 |
+
- **Latent steps**: K=6 autoregressive continuous thought steps
|
| 56 |
+
- **Inference**: Process prompt → K latent steps via ProjectionHead + KV cache → greedy text generation
|
| 57 |
+
|
| 58 |
+
## Training Details
|
| 59 |
+
|
| 60 |
+
| Parameter | Value |
|
| 61 |
+
|---|---|
|
| 62 |
+
| Mode | `codi_single` |
|
| 63 |
+
| Per-step distillation | `True` |
|
| 64 |
+
| Distillation γ | 1.0 |
|
| 65 |
+
| Learning rate | 0.0002 |
|
| 66 |
+
| Epochs | 3 |
|
| 67 |
+
| Batch size (per GPU) | 2 |
|
| 68 |
+
| Gradient accumulation | 8 |
|
| 69 |
+
| Effective batch size | 128 (across 8 GPUs) |
|
| 70 |
+
| Max answer length | 128 |
|
| 71 |
+
| Latent steps (K) | 6 |
|
| 72 |
+
| Task | GSM8k (7,473 training problems) |
|
| 73 |
+
| Rollouts per problem | 16 |
|
| 74 |
+
| **GSM8k test accuracy** | **80.4%** |
|
| 75 |
+
|
| 76 |
+
## How to Use
|
| 77 |
+
|
| 78 |
+
### Requirements
|
| 79 |
+
|
| 80 |
+
```bash
|
| 81 |
+
pip install torch transformers
|
| 82 |
+
```
|
| 83 |
+
|
| 84 |
+
### Inference Code
|
| 85 |
+
|
| 86 |
+
```python
|
| 87 |
+
import torch
|
| 88 |
+
import torch.nn as nn
|
| 89 |
+
from transformers import AutoModelForCausalLM, AutoTokenizer
|
| 90 |
+
|
| 91 |
+
# Load model
|
| 92 |
+
model_name = "LakshyAAAgrawal/continuous-thought-r11_single_perstep_g1"
|
| 93 |
+
tokenizer = AutoTokenizer.from_pretrained(model_name)
|
| 94 |
+
model = AutoModelForCausalLM.from_pretrained(
|
| 95 |
+
model_name, torch_dtype=torch.bfloat16, device_map="auto"
|
| 96 |
+
)
|
| 97 |
+
model.eval()
|
| 98 |
+
|
| 99 |
+
# Load projection head
|
| 100 |
+
class ProjectionHead(nn.Module):
|
| 101 |
+
def __init__(self, hidden_size):
|
| 102 |
+
super().__init__()
|
| 103 |
+
self.mlp = nn.Sequential(
|
| 104 |
+
nn.Linear(hidden_size, hidden_size),
|
| 105 |
+
nn.GELU(),
|
| 106 |
+
nn.Linear(hidden_size, hidden_size),
|
| 107 |
+
nn.LayerNorm(hidden_size),
|
| 108 |
+
)
|
| 109 |
+
def forward(self, x):
|
| 110 |
+
return self.mlp(x)
|
| 111 |
+
|
| 112 |
+
proj = ProjectionHead(model.config.hidden_size)
|
| 113 |
+
proj.load_state_dict(torch.load(
|
| 114 |
+
hf_hub_download(model_name, "projection_head.pt"), map_location="cpu"
|
| 115 |
+
))
|
| 116 |
+
proj = proj.to(model.dtype).to(model.device).eval()
|
| 117 |
+
|
| 118 |
+
# Generate with latent reasoning
|
| 119 |
+
question = "Natalia sold clips to 48 of her friends in April, and then she sold half as many clips in May. How many clips did Natalia sell altogether in April and May?"
|
| 120 |
+
messages = [{"role": "user", "content": f"Solve the following math problem step by step. Show your work and put your final numerical answer after ####.\n\n{question}"}]
|
| 121 |
+
prompt = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
|
| 122 |
+
inputs = tokenizer(prompt, return_tensors="pt").to(model.device)
|
| 123 |
+
|
| 124 |
+
# Step 1: Process prompt
|
| 125 |
+
with torch.no_grad():
|
| 126 |
+
out = model(**inputs, output_hidden_states=True, use_cache=True)
|
| 127 |
+
past_kv = out.past_key_values
|
| 128 |
+
latent = out.hidden_states[-1][:, -1, :] # last token hidden state
|
| 129 |
+
|
| 130 |
+
# Step 2: K latent reasoning steps
|
| 131 |
+
mask = inputs["attention_mask"].clone()
|
| 132 |
+
for k in range(6):
|
| 133 |
+
latent = proj(latent)
|
| 134 |
+
mask = torch.cat([mask, torch.ones(1, 1, device=mask.device, dtype=mask.dtype)], dim=1)
|
| 135 |
+
out = model(inputs_embeds=latent.unsqueeze(1), attention_mask=mask,
|
| 136 |
+
past_key_values=past_kv, output_hidden_states=True, use_cache=True)
|
| 137 |
+
past_kv = out.past_key_values
|
| 138 |
+
latent = out.hidden_states[-1][:, -1, :]
|
| 139 |
+
|
| 140 |
+
# Step 3: Greedy text generation
|
| 141 |
+
next_token = out.logits[:, -1, :].argmax(dim=-1)
|
| 142 |
+
generated = [next_token]
|
| 143 |
+
for _ in range(2047):
|
| 144 |
+
if next_token.item() == tokenizer.eos_token_id:
|
| 145 |
+
break
|
| 146 |
+
mask = torch.cat([mask, torch.ones(1, 1, device=mask.device, dtype=mask.dtype)], dim=1)
|
| 147 |
+
out = model(input_ids=next_token.unsqueeze(0), attention_mask=mask,
|
| 148 |
+
past_key_values=past_kv, use_cache=True)
|
| 149 |
+
past_kv = out.past_key_values
|
| 150 |
+
next_token = out.logits[:, -1, :].argmax(dim=-1)
|
| 151 |
+
generated.append(next_token)
|
| 152 |
+
|
| 153 |
+
response = tokenizer.decode(torch.cat(generated), skip_special_tokens=True)
|
| 154 |
+
print(response)
|
| 155 |
+
```
|
| 156 |
+
|
| 157 |
+
### Evaluation
|
| 158 |
+
|
| 159 |
+
To evaluate on the full GSM8k test set, use the evaluation script from our repository:
|
| 160 |
+
|
| 161 |
+
```bash
|
| 162 |
+
python evaluate.py \
|
| 163 |
+
--model_dir LakshyAAAgrawal/continuous-thought-r11_single_perstep_g1 \
|
| 164 |
+
--mode codi_single \
|
| 165 |
+
--output results/r11_single_perstep_g1.json \
|
| 166 |
+
--max_new_tokens 2048 \
|
| 167 |
+
--num_latent 6
|
| 168 |
+
```
|
| 169 |
+
|
| 170 |
+
**Important**: Use `max_new_tokens=2048` for evaluation. The model generates verbose
|
| 171 |
+
chain-of-thought text after the latent steps, requiring more tokens than standard models.
|
| 172 |
+
|
| 173 |
+
## Results Comparison
|
| 174 |
+
|
| 175 |
+
| Model | Mode | Per-step | γ | ans_len | GSM8k Accuracy |
|
| 176 |
+
|---|---|---|---|---|---|
|
| 177 |
+
| **This model** | **codi_single** | **True** | **1.0** | **128** | **80.4%** |
|
| 178 |
+
| Qwen3-1.7B (base) | — | — | — | — | 77.3% |
|
| 179 |
+
| Discrete SFT | sft | — | — | — | 80.7% |
|
| 180 |
+
| CODI RW final-step | rw | no | 1.0 | 128 | 81.0% |
|
| 181 |
+
| CODI Uniform per-step ans256 | uniform | yes | 2.0 | 256 | 83.2% |
|
| 182 |
+
| CODI RW per-step ans256 | rw | yes | 1.0 | 256 | 82.7% |
|
| 183 |
+
|
| 184 |
+
## Citation
|
| 185 |
+
|
| 186 |
+
If you use this model, please cite:
|
| 187 |
+
|
| 188 |
+
```bibtex
|
| 189 |
+
@misc{continuous-thought-2025,
|
| 190 |
+
title={Continuous Thought Distillation: Reward-Weighted Multi-Trace Reasoning},
|
| 191 |
+
author={Lakshya Agrawal},
|
| 192 |
+
year={2025},
|
| 193 |
+
url={https://huggingface.co/LakshyAAAgrawal/continuous-thought-r11_single_perstep_g1}
|
| 194 |
+
}
|
| 195 |
+
```
|
chat_template.jinja
ADDED
|
@@ -0,0 +1,89 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{%- if tools %}
|
| 2 |
+
{{- '<|im_start|>system\n' }}
|
| 3 |
+
{%- if messages[0].role == 'system' %}
|
| 4 |
+
{{- messages[0].content + '\n\n' }}
|
| 5 |
+
{%- endif %}
|
| 6 |
+
{{- "# Tools\n\nYou may call one or more functions to assist with the user query.\n\nYou are provided with function signatures within <tools></tools> XML tags:\n<tools>" }}
|
| 7 |
+
{%- for tool in tools %}
|
| 8 |
+
{{- "\n" }}
|
| 9 |
+
{{- tool | tojson }}
|
| 10 |
+
{%- endfor %}
|
| 11 |
+
{{- "\n</tools>\n\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\n<tool_call>\n{\"name\": <function-name>, \"arguments\": <args-json-object>}\n</tool_call><|im_end|>\n" }}
|
| 12 |
+
{%- else %}
|
| 13 |
+
{%- if messages[0].role == 'system' %}
|
| 14 |
+
{{- '<|im_start|>system\n' + messages[0].content + '<|im_end|>\n' }}
|
| 15 |
+
{%- endif %}
|
| 16 |
+
{%- endif %}
|
| 17 |
+
{%- set ns = namespace(multi_step_tool=true, last_query_index=messages|length - 1) %}
|
| 18 |
+
{%- for message in messages[::-1] %}
|
| 19 |
+
{%- set index = (messages|length - 1) - loop.index0 %}
|
| 20 |
+
{%- if ns.multi_step_tool and message.role == "user" and message.content is string and not(message.content.startswith('<tool_response>') and message.content.endswith('</tool_response>')) %}
|
| 21 |
+
{%- set ns.multi_step_tool = false %}
|
| 22 |
+
{%- set ns.last_query_index = index %}
|
| 23 |
+
{%- endif %}
|
| 24 |
+
{%- endfor %}
|
| 25 |
+
{%- for message in messages %}
|
| 26 |
+
{%- if message.content is string %}
|
| 27 |
+
{%- set content = message.content %}
|
| 28 |
+
{%- else %}
|
| 29 |
+
{%- set content = '' %}
|
| 30 |
+
{%- endif %}
|
| 31 |
+
{%- if (message.role == "user") or (message.role == "system" and not loop.first) %}
|
| 32 |
+
{{- '<|im_start|>' + message.role + '\n' + content + '<|im_end|>' + '\n' }}
|
| 33 |
+
{%- elif message.role == "assistant" %}
|
| 34 |
+
{%- set reasoning_content = '' %}
|
| 35 |
+
{%- if message.reasoning_content is string %}
|
| 36 |
+
{%- set reasoning_content = message.reasoning_content %}
|
| 37 |
+
{%- else %}
|
| 38 |
+
{%- if '</think>' in content %}
|
| 39 |
+
{%- set reasoning_content = content.split('</think>')[0].rstrip('\n').split('<think>')[-1].lstrip('\n') %}
|
| 40 |
+
{%- set content = content.split('</think>')[-1].lstrip('\n') %}
|
| 41 |
+
{%- endif %}
|
| 42 |
+
{%- endif %}
|
| 43 |
+
{%- if loop.index0 > ns.last_query_index %}
|
| 44 |
+
{%- if loop.last or (not loop.last and reasoning_content) %}
|
| 45 |
+
{{- '<|im_start|>' + message.role + '\n<think>\n' + reasoning_content.strip('\n') + '\n</think>\n\n' + content.lstrip('\n') }}
|
| 46 |
+
{%- else %}
|
| 47 |
+
{{- '<|im_start|>' + message.role + '\n' + content }}
|
| 48 |
+
{%- endif %}
|
| 49 |
+
{%- else %}
|
| 50 |
+
{{- '<|im_start|>' + message.role + '\n' + content }}
|
| 51 |
+
{%- endif %}
|
| 52 |
+
{%- if message.tool_calls %}
|
| 53 |
+
{%- for tool_call in message.tool_calls %}
|
| 54 |
+
{%- if (loop.first and content) or (not loop.first) %}
|
| 55 |
+
{{- '\n' }}
|
| 56 |
+
{%- endif %}
|
| 57 |
+
{%- if tool_call.function %}
|
| 58 |
+
{%- set tool_call = tool_call.function %}
|
| 59 |
+
{%- endif %}
|
| 60 |
+
{{- '<tool_call>\n{"name": "' }}
|
| 61 |
+
{{- tool_call.name }}
|
| 62 |
+
{{- '", "arguments": ' }}
|
| 63 |
+
{%- if tool_call.arguments is string %}
|
| 64 |
+
{{- tool_call.arguments }}
|
| 65 |
+
{%- else %}
|
| 66 |
+
{{- tool_call.arguments | tojson }}
|
| 67 |
+
{%- endif %}
|
| 68 |
+
{{- '}\n</tool_call>' }}
|
| 69 |
+
{%- endfor %}
|
| 70 |
+
{%- endif %}
|
| 71 |
+
{{- '<|im_end|>\n' }}
|
| 72 |
+
{%- elif message.role == "tool" %}
|
| 73 |
+
{%- if loop.first or (messages[loop.index0 - 1].role != "tool") %}
|
| 74 |
+
{{- '<|im_start|>user' }}
|
| 75 |
+
{%- endif %}
|
| 76 |
+
{{- '\n<tool_response>\n' }}
|
| 77 |
+
{{- content }}
|
| 78 |
+
{{- '\n</tool_response>' }}
|
| 79 |
+
{%- if loop.last or (messages[loop.index0 + 1].role != "tool") %}
|
| 80 |
+
{{- '<|im_end|>\n' }}
|
| 81 |
+
{%- endif %}
|
| 82 |
+
{%- endif %}
|
| 83 |
+
{%- endfor %}
|
| 84 |
+
{%- if add_generation_prompt %}
|
| 85 |
+
{{- '<|im_start|>assistant\n' }}
|
| 86 |
+
{%- if enable_thinking is defined and enable_thinking is false %}
|
| 87 |
+
{{- '<think>\n\n</think>\n\n' }}
|
| 88 |
+
{%- endif %}
|
| 89 |
+
{%- endif %}
|
config.json
ADDED
|
@@ -0,0 +1,63 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"architectures": [
|
| 3 |
+
"Qwen3ForCausalLM"
|
| 4 |
+
],
|
| 5 |
+
"attention_bias": false,
|
| 6 |
+
"attention_dropout": 0.0,
|
| 7 |
+
"bos_token_id": 151643,
|
| 8 |
+
"dtype": "bfloat16",
|
| 9 |
+
"eos_token_id": 151645,
|
| 10 |
+
"head_dim": 128,
|
| 11 |
+
"hidden_act": "silu",
|
| 12 |
+
"hidden_size": 2048,
|
| 13 |
+
"initializer_range": 0.02,
|
| 14 |
+
"intermediate_size": 6144,
|
| 15 |
+
"layer_types": [
|
| 16 |
+
"full_attention",
|
| 17 |
+
"full_attention",
|
| 18 |
+
"full_attention",
|
| 19 |
+
"full_attention",
|
| 20 |
+
"full_attention",
|
| 21 |
+
"full_attention",
|
| 22 |
+
"full_attention",
|
| 23 |
+
"full_attention",
|
| 24 |
+
"full_attention",
|
| 25 |
+
"full_attention",
|
| 26 |
+
"full_attention",
|
| 27 |
+
"full_attention",
|
| 28 |
+
"full_attention",
|
| 29 |
+
"full_attention",
|
| 30 |
+
"full_attention",
|
| 31 |
+
"full_attention",
|
| 32 |
+
"full_attention",
|
| 33 |
+
"full_attention",
|
| 34 |
+
"full_attention",
|
| 35 |
+
"full_attention",
|
| 36 |
+
"full_attention",
|
| 37 |
+
"full_attention",
|
| 38 |
+
"full_attention",
|
| 39 |
+
"full_attention",
|
| 40 |
+
"full_attention",
|
| 41 |
+
"full_attention",
|
| 42 |
+
"full_attention",
|
| 43 |
+
"full_attention"
|
| 44 |
+
],
|
| 45 |
+
"max_position_embeddings": 40960,
|
| 46 |
+
"max_window_layers": 28,
|
| 47 |
+
"model_type": "qwen3",
|
| 48 |
+
"num_attention_heads": 16,
|
| 49 |
+
"num_hidden_layers": 28,
|
| 50 |
+
"num_key_value_heads": 8,
|
| 51 |
+
"pad_token_id": null,
|
| 52 |
+
"rms_norm_eps": 1e-06,
|
| 53 |
+
"rope_parameters": {
|
| 54 |
+
"rope_theta": 1000000,
|
| 55 |
+
"rope_type": "default"
|
| 56 |
+
},
|
| 57 |
+
"sliding_window": null,
|
| 58 |
+
"tie_word_embeddings": true,
|
| 59 |
+
"transformers_version": "5.3.0",
|
| 60 |
+
"use_cache": true,
|
| 61 |
+
"use_sliding_window": false,
|
| 62 |
+
"vocab_size": 151936
|
| 63 |
+
}
|
generation_config.json
ADDED
|
@@ -0,0 +1,13 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"bos_token_id": 151643,
|
| 3 |
+
"do_sample": true,
|
| 4 |
+
"eos_token_id": [
|
| 5 |
+
151645,
|
| 6 |
+
151643
|
| 7 |
+
],
|
| 8 |
+
"pad_token_id": 151643,
|
| 9 |
+
"temperature": 0.6,
|
| 10 |
+
"top_k": 20,
|
| 11 |
+
"top_p": 0.95,
|
| 12 |
+
"transformers_version": "5.3.0"
|
| 13 |
+
}
|
model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:b59f22809db59a437b741422a3d2b349f44b2c6d8b25383d34f767016bf033f0
|
| 3 |
+
size 4063515640
|
projection_head.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:60e3973048cd1c8f14ec58c8d59f3b60dd11f2b84ae21f361b3820cfdca073a2
|
| 3 |
+
size 16796853
|
tokenizer.json
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:be75606093db2094d7cd20f3c2f385c212750648bd6ea4fb2bf507a6a4c55506
|
| 3 |
+
size 11422650
|
tokenizer_config.json
ADDED
|
@@ -0,0 +1,29 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"add_prefix_space": false,
|
| 3 |
+
"backend": "tokenizers",
|
| 4 |
+
"bos_token": null,
|
| 5 |
+
"clean_up_tokenization_spaces": false,
|
| 6 |
+
"eos_token": "<|im_end|>",
|
| 7 |
+
"errors": "replace",
|
| 8 |
+
"extra_special_tokens": [
|
| 9 |
+
"<|im_start|>",
|
| 10 |
+
"<|im_end|>",
|
| 11 |
+
"<|object_ref_start|>",
|
| 12 |
+
"<|object_ref_end|>",
|
| 13 |
+
"<|box_start|>",
|
| 14 |
+
"<|box_end|>",
|
| 15 |
+
"<|quad_start|>",
|
| 16 |
+
"<|quad_end|>",
|
| 17 |
+
"<|vision_start|>",
|
| 18 |
+
"<|vision_end|>",
|
| 19 |
+
"<|vision_pad|>",
|
| 20 |
+
"<|image_pad|>",
|
| 21 |
+
"<|video_pad|>"
|
| 22 |
+
],
|
| 23 |
+
"is_local": false,
|
| 24 |
+
"model_max_length": 131072,
|
| 25 |
+
"pad_token": "<|endoftext|>",
|
| 26 |
+
"split_special_tokens": false,
|
| 27 |
+
"tokenizer_class": "Qwen2Tokenizer",
|
| 28 |
+
"unk_token": null
|
| 29 |
+
}
|
training_config.json
ADDED
|
@@ -0,0 +1,25 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"mode": "codi_single",
|
| 3 |
+
"model": "Qwen/Qwen3-1.7B",
|
| 4 |
+
"rollouts": "data/rollouts.jsonl",
|
| 5 |
+
"teacher_states": "data/teacher_states_v3.pt",
|
| 6 |
+
"output_dir": "checkpoints/r11_single_perstep_g1",
|
| 7 |
+
"epochs": 3,
|
| 8 |
+
"batch_size": 2,
|
| 9 |
+
"grad_accum": 8,
|
| 10 |
+
"lr": 0.0002,
|
| 11 |
+
"gamma": 1.0,
|
| 12 |
+
"warmup_ratio": 0.03,
|
| 13 |
+
"num_latent": 6,
|
| 14 |
+
"max_prompt_len": 512,
|
| 15 |
+
"max_answer_len": 128,
|
| 16 |
+
"max_len": 1024,
|
| 17 |
+
"answer_format": "full",
|
| 18 |
+
"per_step_distill": true,
|
| 19 |
+
"log_every": 10,
|
| 20 |
+
"seed": 42,
|
| 21 |
+
"lora_rank": 32,
|
| 22 |
+
"lora_alpha": 16,
|
| 23 |
+
"best_loss": 0.6426934776270491,
|
| 24 |
+
"total_time_s": 1477.4382388591766
|
| 25 |
+
}
|