LakshyAAAgrawal commited on
Commit
0135ce6
·
verified ·
1 Parent(s): 745d676

Upload folder using huggingface_hub

Browse files
.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,184 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ language:
3
+ - en
4
+ license: apache-2.0
5
+ library_name: transformers
6
+ base_model: Qwen/Qwen3-1.7B
7
+ tags:
8
+ - qthink
9
+ - latent-reasoning
10
+ - distillation
11
+ - qwen3
12
+ - lora
13
+ - tooluse
14
+ datasets:
15
+ - tooluse
16
+ pipeline_tag: text-generation
17
+ model-index:
18
+ - name: QThink-Qwen3-1.7B-Tooluse
19
+ results:
20
+ - task:
21
+ type: text-generation
22
+ name: Tooluse
23
+ dataset:
24
+ name: Tooluse
25
+ type: tooluse
26
+ split: test
27
+ metrics:
28
+ - type: accuracy
29
+ value: 48.5% EM
30
+ name: Accuracy
31
+ ---
32
+
33
+ # QThink-Qwen3-1.7B-Tooluse
34
+
35
+ **QThink: Parallel Latent Reasoning via Per-Step Distillation of Multiple Rollouts**
36
+
37
+ This model replaces explicit chain-of-thought (`<think>...</think>`) with **6 latent forward passes** through a learned projection head, achieving **48.5% EM** on Tooluse.
38
+
39
+ ## How QThink Works
40
+
41
+ Instead of generating thousands of reasoning tokens, QThink:
42
+
43
+ 1. **Processes the prompt** through the base model
44
+ 2. **Runs K=6 latent steps**: each step applies a ProjectionHead (`Linear(2048,2048) → GELU → Linear(2048,2048) → LayerNorm(2048)`) to the hidden state, then feeds it back through the model via `inputs_embeds` + `past_key_values`
45
+ 3. **Generates the answer** directly from the last latent step's hidden state — no `<think>` block needed
46
+
47
+ ### Training: Per-Step Distillation from Multiple Rollouts
48
+
49
+ 1. Generate G=16 chain-of-thought rollouts per training problem using the base model
50
+ 2. For each rollout, extract hidden states at K=6 evenly-spaced positions within the `<think>` block
51
+ 3. Average these hidden states across **all rollouts** (uniform — including incorrect ones) at each step
52
+ 4. Train each latent step to match the corresponding teacher state via L1 loss, jointly with cross-entropy on the answer
53
+
54
+ The key innovations:
55
+ - **Per-step distillation**: Every latent step gets direct supervision, not just the final one
56
+ - **Uniform multi-rollout teachers**: Averaging over ALL rollouts (correct + incorrect) outperforms using only correct rollouts
57
+ - **ans256 training**: Training with longer answer targets (+3pp improvement)
58
+
59
+ ## Results
60
+
61
+ | Model | Exact Match | Action Accuracy |
62
+ |-------|------------|----------------|
63
+ | **QThink uniform per-step (ours)** | **48.5%** | 67.6% |
64
+ | Base Qwen3-1.7B | 47.1% | 75.0% |
65
+ | SFT | 45.6% | 73.5% |
66
+ | QThink RW final-step (CODI) | 42.6% | 67.6% |
67
+
68
+
69
+ ### Cross-Benchmark Results
70
+
71
+ QThink uniform per-step is the **best model on all 3 benchmarks**:
72
+
73
+ | Benchmark | QThink (ours) | SFT | Base | CODI |
74
+ |-----------|--------------|-----|------|------|
75
+ | GSM8k | **83.2%** | 80.7% | 77.3% | 80.4% |
76
+ | MATH-500 | **43.6%** | 38.2% | 33.6% | 28.0% |
77
+ | Tooluse | **48.5%** | 45.6% | 47.1% | 42.6% |
78
+
79
+ ## Model Details
80
+
81
+ | Parameter | Value |
82
+ |-----------|-------|
83
+ | Base model | [Qwen/Qwen3-1.7B](https://huggingface.co/Qwen/Qwen3-1.7B) |
84
+ | Fine-tuning | LoRA (rank=32, alpha=16) |
85
+ | Distillation mode | Uniform (all rollouts) |
86
+ | Per-step distillation | Yes (K=6 steps) |
87
+ | Distillation weight (γ) | 2.0 |
88
+ | Learning rate | 0.0002 |
89
+ | Epochs | 3 |
90
+ | Batch size × grad accum × GPUs | 1 × 16 × 8 = 128 effective |
91
+ | Max answer length | 256 tokens |
92
+ | Max prompt length | 1024 tokens |
93
+ | Rollouts per problem | 16 |
94
+ | Dataset | Tooluse (from [SDPO](https://github.com/lasgroup/SDPO)) — 4,046 train problems, 68 test problems |
95
+
96
+ ## Architecture
97
+
98
+ The checkpoint contains:
99
+ - **`model.safetensors`**: Full Qwen3-1.7B weights with merged LoRA adapters
100
+ - **`projection_head.pt`**: ProjectionHead weights (PyTorch state dict)
101
+ - `mlp.0`: Linear(2048 → 2048) + bias
102
+ - `mlp.2`: Linear(2048 → 2048) + bias
103
+ - `mlp.3`: LayerNorm(2048)
104
+
105
+ ## Usage
106
+
107
+ ```python
108
+ import torch
109
+ import torch.nn as nn
110
+ from transformers import AutoModelForCausalLM, AutoTokenizer
111
+ from huggingface_hub import hf_hub_download
112
+
113
+ class ProjectionHead(nn.Module):
114
+ def __init__(self, hidden_size=2048):
115
+ super().__init__()
116
+ self.mlp = nn.Sequential(
117
+ nn.Linear(hidden_size, hidden_size),
118
+ nn.GELU(),
119
+ nn.Linear(hidden_size, hidden_size),
120
+ nn.LayerNorm(hidden_size),
121
+ )
122
+ def forward(self, x):
123
+ return self.mlp(x)
124
+
125
+ # Load model and projection head
126
+ repo_id = "LakshyAAAgrawal/QThink-Qwen3-1.7B-Tooluse"
127
+ model = AutoModelForCausalLM.from_pretrained(repo_id, torch_dtype=torch.bfloat16, device_map="auto")
128
+ tokenizer = AutoTokenizer.from_pretrained(repo_id)
129
+ proj = ProjectionHead(2048).to(model.device).to(torch.bfloat16)
130
+ proj.load_state_dict(torch.load(
131
+ hf_hub_download(repo_id, "projection_head.pt"), map_location=model.device
132
+ ))
133
+ proj.eval()
134
+ model.eval()
135
+
136
+ # Prepare input
137
+ question = "What's the weather like in San Francisco?"
138
+ messages = [{"role": "user", "content": question}]
139
+ text = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True, enable_thinking=False)
140
+ inputs = tokenizer(text, return_tensors="pt").to(model.device)
141
+
142
+ with torch.no_grad():
143
+ # Step 1: Process prompt
144
+ out = model(**inputs, output_hidden_states=True, use_cache=True)
145
+ past_kv = out.past_key_values
146
+ latent = out.hidden_states[-1][:, -1, :]
147
+
148
+ # Step 2: K=6 latent reasoning steps
149
+ mask = inputs["attention_mask"].clone()
150
+ for k in range(6):
151
+ latent = proj(latent)
152
+ mask = torch.cat([mask, mask.new_ones(1, 1)], dim=1)
153
+ out = model(inputs_embeds=latent.unsqueeze(1), attention_mask=mask,
154
+ past_key_values=past_kv, output_hidden_states=True, use_cache=True)
155
+ past_kv = out.past_key_values
156
+ latent = out.hidden_states[-1][:, -1, :]
157
+
158
+ # Step 3: Greedy decode answer
159
+ next_token = out.logits[:, -1, :].argmax(dim=-1)
160
+ tokens = [next_token]
161
+ eos_id = tokenizer.eos_token_id
162
+ for _ in range(2047):
163
+ if next_token.item() == eos_id:
164
+ break
165
+ mask = torch.cat([mask, mask.new_ones(1, 1)], dim=1)
166
+ out = model(input_ids=next_token.unsqueeze(0), attention_mask=mask,
167
+ past_key_values=past_kv, use_cache=True)
168
+ past_kv = out.past_key_values
169
+ next_token = out.logits[:, -1, :].argmax(dim=-1)
170
+ tokens.append(next_token)
171
+
172
+ print(tokenizer.decode(torch.cat(tokens), skip_special_tokens=True))
173
+ ```
174
+
175
+ ## Citation
176
+
177
+ ```bibtex
178
+ @misc{qthink2025,
179
+ title={QThink: Parallel Latent Reasoning via Per-Step Distillation of Multiple Rollouts},
180
+ author={Lakshya Agrawal},
181
+ year={2025},
182
+ url={https://huggingface.co/LakshyAAAgrawal/QThink-Qwen3-1.7B-Tooluse}
183
+ }
184
+ ```
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:861c41dfdd21044334d14b35d7cce7484a53aff89e21a44188819afe6595ab4c
3
+ size 4063515640
projection_head.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:fc867470fe41a65ec91a64a2f87b2087a79a6125eb5edaddc77d8e27a44ea0a4
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_uniform",
3
+ "model": "Qwen/Qwen3-1.7B",
4
+ "rollouts": "data/tooluse_rollouts.jsonl",
5
+ "teacher_states": "data/tooluse_teacher_states.pt",
6
+ "output_dir": "checkpoints/r14_tooluse_qthink_uniform_perstep_g2_ans256",
7
+ "epochs": 3,
8
+ "batch_size": 1,
9
+ "grad_accum": 16,
10
+ "lr": 0.0002,
11
+ "gamma": 2.0,
12
+ "warmup_ratio": 0.03,
13
+ "num_latent": 6,
14
+ "max_prompt_len": 1024,
15
+ "max_answer_len": 256,
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": 1.1106567853995464,
24
+ "total_time_s": 1602.5439190864563
25
+ }