rjz123 commited on
Commit
5b97bd2
·
verified ·
1 Parent(s): dca9747

Upload folder using huggingface_hub

Browse files
Files changed (3) hide show
  1. README.md +22 -0
  2. colar_coding.ckpt +3 -0
  3. hparams.yaml +167 -0
README.md ADDED
@@ -0,0 +1,22 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ base_model: unsloth/Llama-3.2-1B-Instruct
3
+ library_name: peft
4
+ pipeline_tag: text-generation
5
+ tags:
6
+ - latent-reasoning
7
+ - colar
8
+ - research
9
+ ---
10
+ # colar-coding-l1b
11
+
12
+ single-track CoLaR, coding 域(符号执行/coding_mix), warm-start 自 colar-gsm, compress=5
13
+
14
+ - **Base model:** `unsloth/Llama-3.2-1B-Instruct`
15
+ - **Files:** `colar_coding.ckpt`, `hparams.yaml`
16
+
17
+ ## Loading (PyTorch-Lightning checkpoint — NOT AutoModel-loadable)
18
+ Weights live under the top-level key `['state_dict']` and only fit the custom CoLaR scaffold (base LLM + `[PAD]` resize + r128 q/v LoRA + a `LatentPolicy` MLP), loaded `strict=False`. Load the base separately and splice this `state_dict` in. Runtime env:
19
+ ```
20
+ COLAR_BASE=<base> COLAR_CKPT=colar-gsm/colar_best.ckpt COLAR_EMB_STD=0.018 COLAR_COMPRESS=5 COLAR_MAXLAT=64 TORCH_FORCE_NO_WEIGHTS_ONLY_LOAD=1
21
+ ```
22
+ `TORCH_FORCE_NO_WEIGHTS_ONLY_LOAD=1` is required for these older Lightning ckpts.
colar_coding.ckpt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:e03131e8382eba1670857eec1979c89225c439b62160e6fdba2b39e7aedb0af2
3
+ size 121741061
hparams.yaml ADDED
@@ -0,0 +1,167 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ model_kwargs:
2
+ model_id: Llama-3.2-1B-Instruct
3
+ sft_method: colar
4
+ chat_template: false
5
+ do_lora: true
6
+ lora_config:
7
+ r: 128
8
+ lora_alpha: 32
9
+ latent_cot_config:
10
+ ce_weight: 1
11
+ embed_modeling_weight: 1
12
+ embed_modeling_loss: mse
13
+ entropy_weight: 0
14
+ pred_embed_forward_weight: 0
15
+ max_compression_factor: 5
16
+ pred_compressed_cot: true
17
+ sqrt_mean: true
18
+ latent_policy_config:
19
+ lp_determinisitc: false
20
+ lp_intermediate_size: 2048
21
+ latent_generation_config:
22
+ max_n_latent_forward: 64
23
+ latent_temperature: 1.0
24
+ compression_factor: 5
25
+ answer_generation_config:
26
+ max_new_tokens: 16
27
+ do_sample: true
28
+ top_p: 0.9
29
+ temperature: 1.0
30
+ do_rl: false
31
+ rl_config:
32
+ average_per_token_loss: false
33
+ random_speed_in_group: false
34
+ filter_dataset: false
35
+ exp_batch_size: 8
36
+ group_size: 8
37
+ punish_latent_length: false
38
+ clip_grad_norm: 1.0
39
+ clip_eps: 0.2
40
+ use_latent_loss: true
41
+ use_answer_loss: true
42
+ n_train_samples_per_epoch: 512
43
+ training_kwargs:
44
+ optimizer:
45
+ target: torch.optim.AdamW
46
+ lr: 0.0001
47
+ weight_decay: 0.01
48
+ use_scheduler: false
49
+ scheduler:
50
+ target: constant_schedule_with_warmup
51
+ warmup_steps: 1000
52
+ all_config:
53
+ trainer:
54
+ target: lightning.pytorch.trainer.Trainer
55
+ devices:
56
+ - 0
57
+ max_steps: -1
58
+ check_val_every_n_epoch: 5
59
+ log_every_n_steps: 10
60
+ num_sanity_val_steps: 2
61
+ gradient_clip_val: 1.0
62
+ reload_dataloaders_every_n_epochs: 0
63
+ accumulate_grad_batches: 4
64
+ precision: bf16-mixed
65
+ use_distributed_sampler: true
66
+ strategy: auto
67
+ logger:
68
+ target: lightning.pytorch.loggers.TensorBoardLogger
69
+ save_dir: logs/colar
70
+ name: qsa-coding_mix
71
+ version: 20260802-200100_983467_coding_gsmwarm
72
+ max_epochs: 25
73
+ callbacks:
74
+ - target: lightning.pytorch.callbacks.ModelCheckpoint
75
+ save_last: true
76
+ save_top_k: 3
77
+ mode: max
78
+ monitor: monitor
79
+ auto_insert_metric_name: false
80
+ filename: epoch{epoch}__step{step}__monitor{monitor:.3f}
81
+ save_weights_only: true
82
+ seed: null
83
+ model:
84
+ target: src.models.colar.LitCoLaR
85
+ model_kwargs:
86
+ model_id: Llama-3.2-1B-Instruct
87
+ sft_method: colar
88
+ chat_template: false
89
+ do_lora: true
90
+ lora_config:
91
+ r: 128
92
+ lora_alpha: 32
93
+ latent_cot_config:
94
+ ce_weight: 1
95
+ embed_modeling_weight: 1
96
+ embed_modeling_loss: mse
97
+ entropy_weight: 0
98
+ pred_embed_forward_weight: 0
99
+ max_compression_factor: 5
100
+ pred_compressed_cot: true
101
+ sqrt_mean: true
102
+ latent_policy_config:
103
+ lp_determinisitc: false
104
+ lp_intermediate_size: 2048
105
+ latent_generation_config:
106
+ max_n_latent_forward: 64
107
+ latent_temperature: 1.0
108
+ compression_factor: 5
109
+ answer_generation_config:
110
+ max_new_tokens: 16
111
+ do_sample: true
112
+ top_p: 0.9
113
+ temperature: 1.0
114
+ do_rl: false
115
+ rl_config:
116
+ average_per_token_loss: false
117
+ random_speed_in_group: false
118
+ filter_dataset: false
119
+ exp_batch_size: 8
120
+ group_size: 8
121
+ punish_latent_length: false
122
+ clip_grad_norm: 1.0
123
+ clip_eps: 0.2
124
+ use_latent_loss: true
125
+ use_answer_loss: true
126
+ n_train_samples_per_epoch: 512
127
+ training_kwargs:
128
+ optimizer:
129
+ target: torch.optim.AdamW
130
+ lr: 0.0001
131
+ weight_decay: 0.01
132
+ use_scheduler: false
133
+ scheduler:
134
+ target: constant_schedule_with_warmup
135
+ warmup_steps: 1000
136
+ dataloader:
137
+ batch_size: 4
138
+ val_batch_size: 32
139
+ num_workers: 8
140
+ pin_memory: true
141
+ persistent_workers: true
142
+ data_module:
143
+ target: src.datasets.qsa.QSADataModule
144
+ dataset_name: coding_mix
145
+ tiny_dataset: false
146
+ epoch_scaling: 1
147
+ args:
148
+ model: colar
149
+ dataset: qsa
150
+ trainer: default
151
+ devices: '0'
152
+ no_log: false
153
+ log_suffix: coding_gsmwarm
154
+ resume_ckpt_path: null
155
+ load_ckpt_path: /content/colar_hf/logs/colar/qsa-gsm/colar-final/checkpoints/colar_best.ckpt
156
+ workspace_path: /content/ws
157
+ do_test: false
158
+ test_ckpt_path: ''
159
+ test_times: 5
160
+ seed: 0
161
+ unkown_args:
162
+ dataset_name: coding_mix
163
+ model_id: Llama-3.2-1B-Instruct
164
+ batch_size: '4'
165
+ accumulate_grad_batches: '4'
166
+ max_epochs: '25'
167
+ check_val_every_n_epoch: '5'