LakshyAAAgrawal commited on
Commit
87cf8ce
·
verified ·
1 Parent(s): ba6a297

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,196 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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_rw_perstep_g1_ans256
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: 82.7
29
+ name: Accuracy
30
+ ---
31
+
32
+ # r11_rw_perstep_g1_ans256
33
+
34
+ **Best reward-weighted — correct-only per-step distillation (82.7%)**
35
+
36
+ - **Best reward-weighted result**: 82.7% on GSM8k
37
+ - Uses reward-weighted teacher (average over CORRECT rollouts only)
38
+ - Per-step distillation at every latent step with γ=1.0
39
+ - Trained with max_answer_len=256
40
+
41
+ ## Overview
42
+
43
+ This model implements **Continuous Thought Distillation (CODI)** — an autoregressive latent
44
+ reasoning loop that processes K=6 continuous thought steps before generating a text answer.
45
+ Teacher hidden states are extracted from multiple chain-of-thought rollouts generated by the
46
+ base model and distilled into the latent representations.
47
+
48
+ **Key idea**: Instead of distilling from a single reasoning trace, we aggregate hidden states
49
+ from multiple rollouts (16 per problem). The teacher signal is the reward-weighted average of CORRECT rollouts only.
50
+
51
+ ## Architecture
52
+
53
+ - **Base model**: Qwen3-1.7B (1.7B parameters)
54
+ - **Fine-tuning**: LoRA (rank=32, alpha=16) on q/k/v/o/gate/up/down_proj
55
+ - **Projection head**: Linear(2048, 2048) → GELU → Linear(2048, 2048) → LayerNorm(2048)
56
+ - **Latent steps**: K=6 autoregressive continuous thought steps
57
+ - **Inference**: Process prompt → K latent steps via ProjectionHead + KV cache → greedy text generation
58
+
59
+ ## Training Details
60
+
61
+ | Parameter | Value |
62
+ |---|---|
63
+ | Mode | `codi_reward_weighted` |
64
+ | Per-step distillation | `True` |
65
+ | Distillation γ | 1.0 |
66
+ | Learning rate | 0.0002 |
67
+ | Epochs | 3 |
68
+ | Batch size (per GPU) | 1 |
69
+ | Gradient accumulation | 16 |
70
+ | Effective batch size | 128 (across 8 GPUs) |
71
+ | Max answer length | 256 |
72
+ | Latent steps (K) | 6 |
73
+ | Task | GSM8k (7,473 training problems) |
74
+ | Rollouts per problem | 16 |
75
+ | **GSM8k test accuracy** | **82.7%** |
76
+
77
+ ## How to Use
78
+
79
+ ### Requirements
80
+
81
+ ```bash
82
+ pip install torch transformers
83
+ ```
84
+
85
+ ### Inference Code
86
+
87
+ ```python
88
+ import torch
89
+ import torch.nn as nn
90
+ from transformers import AutoModelForCausalLM, AutoTokenizer
91
+
92
+ # Load model
93
+ model_name = "LakshyAAAgrawal/continuous-thought-r11_rw_perstep_g1_ans256"
94
+ tokenizer = AutoTokenizer.from_pretrained(model_name)
95
+ model = AutoModelForCausalLM.from_pretrained(
96
+ model_name, torch_dtype=torch.bfloat16, device_map="auto"
97
+ )
98
+ model.eval()
99
+
100
+ # Load projection head
101
+ class ProjectionHead(nn.Module):
102
+ def __init__(self, hidden_size):
103
+ super().__init__()
104
+ self.mlp = nn.Sequential(
105
+ nn.Linear(hidden_size, hidden_size),
106
+ nn.GELU(),
107
+ nn.Linear(hidden_size, hidden_size),
108
+ nn.LayerNorm(hidden_size),
109
+ )
110
+ def forward(self, x):
111
+ return self.mlp(x)
112
+
113
+ proj = ProjectionHead(model.config.hidden_size)
114
+ proj.load_state_dict(torch.load(
115
+ hf_hub_download(model_name, "projection_head.pt"), map_location="cpu"
116
+ ))
117
+ proj = proj.to(model.dtype).to(model.device).eval()
118
+
119
+ # Generate with latent reasoning
120
+ 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?"
121
+ 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}"}]
122
+ prompt = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
123
+ inputs = tokenizer(prompt, return_tensors="pt").to(model.device)
124
+
125
+ # Step 1: Process prompt
126
+ with torch.no_grad():
127
+ out = model(**inputs, output_hidden_states=True, use_cache=True)
128
+ past_kv = out.past_key_values
129
+ latent = out.hidden_states[-1][:, -1, :] # last token hidden state
130
+
131
+ # Step 2: K latent reasoning steps
132
+ mask = inputs["attention_mask"].clone()
133
+ for k in range(6):
134
+ latent = proj(latent)
135
+ mask = torch.cat([mask, torch.ones(1, 1, device=mask.device, dtype=mask.dtype)], dim=1)
136
+ out = model(inputs_embeds=latent.unsqueeze(1), attention_mask=mask,
137
+ past_key_values=past_kv, output_hidden_states=True, use_cache=True)
138
+ past_kv = out.past_key_values
139
+ latent = out.hidden_states[-1][:, -1, :]
140
+
141
+ # Step 3: Greedy text generation
142
+ next_token = out.logits[:, -1, :].argmax(dim=-1)
143
+ generated = [next_token]
144
+ for _ in range(2047):
145
+ if next_token.item() == tokenizer.eos_token_id:
146
+ break
147
+ mask = torch.cat([mask, torch.ones(1, 1, device=mask.device, dtype=mask.dtype)], dim=1)
148
+ out = model(input_ids=next_token.unsqueeze(0), attention_mask=mask,
149
+ past_key_values=past_kv, use_cache=True)
150
+ past_kv = out.past_key_values
151
+ next_token = out.logits[:, -1, :].argmax(dim=-1)
152
+ generated.append(next_token)
153
+
154
+ response = tokenizer.decode(torch.cat(generated), skip_special_tokens=True)
155
+ print(response)
156
+ ```
157
+
158
+ ### Evaluation
159
+
160
+ To evaluate on the full GSM8k test set, use the evaluation script from our repository:
161
+
162
+ ```bash
163
+ python evaluate.py \
164
+ --model_dir LakshyAAAgrawal/continuous-thought-r11_rw_perstep_g1_ans256 \
165
+ --mode codi_reward_weighted \
166
+ --output results/r11_rw_perstep_g1_ans256.json \
167
+ --max_new_tokens 2048 \
168
+ --num_latent 6
169
+ ```
170
+
171
+ **Important**: Use `max_new_tokens=2048` for evaluation. The model generates verbose
172
+ chain-of-thought text after the latent steps, requiring more tokens than standard models.
173
+
174
+ ## Results Comparison
175
+
176
+ | Model | Mode | Per-step | γ | ans_len | GSM8k Accuracy |
177
+ |---|---|---|---|---|---|
178
+ | **This model** | **codi_reward_weighted** | **True** | **1.0** | **256** | **82.7%** |
179
+ | Qwen3-1.7B (base) | — | — | — | — | 77.3% |
180
+ | Discrete SFT | sft | — | — | — | 80.7% |
181
+ | CODI RW final-step | rw | no | 1.0 | 128 | 81.0% |
182
+ | CODI Uniform per-step ans256 | uniform | yes | 2.0 | 256 | 83.2% |
183
+ | CODI RW per-step ans256 | rw | yes | 1.0 | 256 | 82.7% |
184
+
185
+ ## Citation
186
+
187
+ If you use this model, please cite:
188
+
189
+ ```bibtex
190
+ @misc{continuous-thought-2025,
191
+ title={Continuous Thought Distillation: Reward-Weighted Multi-Trace Reasoning},
192
+ author={Lakshya Agrawal},
193
+ year={2025},
194
+ url={https://huggingface.co/LakshyAAAgrawal/continuous-thought-r11_rw_perstep_g1_ans256}
195
+ }
196
+ ```
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:b59c1ed931b3aef61396f2b9f7176ae5891cb98695607983f1902b6ed32a4098
3
+ size 4063515640
projection_head.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:88ecc9fa08f454c7c8909d81fdafd299b85441540c5e2d6f576b2e2d75c06aee
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_reward_weighted",
3
+ "model": "Qwen/Qwen3-1.7B",
4
+ "rollouts": "data/rollouts.jsonl",
5
+ "teacher_states": "data/teacher_states_v3.pt",
6
+ "output_dir": "checkpoints/r11_rw_perstep_g1_ans256",
7
+ "epochs": 3,
8
+ "batch_size": 1,
9
+ "grad_accum": 16,
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": 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": 0.5783941572043985,
24
+ "total_time_s": 2959.5815694332123
25
+ }