RonanMcGovern commited on
Commit
7366f82
·
verified ·
1 Parent(s): 66bb47c

Upload via push_to_hf.py

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +37 -0
  2. all_config.yaml +69 -0
  3. evaluator_ARC_step_110598/submission.json +0 -0
  4. evaluator_ARC_step_129031/submission.json +0 -0
  5. evaluator_ARC_step_147464/submission.json +0 -0
  6. evaluator_ARC_step_165897/submission.json +0 -0
  7. evaluator_ARC_step_18433/submission.json +0 -0
  8. evaluator_ARC_step_184330/submission.json +0 -0
  9. evaluator_ARC_step_202763/submission.json +0 -0
  10. evaluator_ARC_step_221196/submission.json +0 -0
  11. evaluator_ARC_step_239629/submission.json +0 -0
  12. evaluator_ARC_step_258062/submission.json +0 -0
  13. evaluator_ARC_step_276495/submission.json +0 -0
  14. evaluator_ARC_step_294928/submission.json +0 -0
  15. evaluator_ARC_step_313361/submission.json +0 -0
  16. evaluator_ARC_step_331794/submission.json +0 -0
  17. evaluator_ARC_step_350227/submission.json +0 -0
  18. evaluator_ARC_step_36866/submission.json +0 -0
  19. evaluator_ARC_step_368660/submission.json +0 -0
  20. evaluator_ARC_step_387093/submission.json +0 -0
  21. evaluator_ARC_step_405526/submission.json +0 -0
  22. evaluator_ARC_step_423959/submission.json +0 -0
  23. evaluator_ARC_step_442392/submission.json +0 -0
  24. evaluator_ARC_step_460825/submission.json +0 -0
  25. evaluator_ARC_step_479258/submission.json +0 -0
  26. evaluator_ARC_step_497691/submission.json +0 -0
  27. evaluator_ARC_step_516124/submission.json +0 -0
  28. evaluator_ARC_step_534557/submission.json +0 -0
  29. evaluator_ARC_step_55299/submission.json +0 -0
  30. evaluator_ARC_step_552990/submission.json +0 -0
  31. evaluator_ARC_step_571423/submission.json +0 -0
  32. evaluator_ARC_step_589856/submission.json +0 -0
  33. evaluator_ARC_step_608289/submission.json +0 -0
  34. evaluator_ARC_step_626722/submission.json +0 -0
  35. evaluator_ARC_step_645155/submission.json +0 -0
  36. evaluator_ARC_step_663588/submission.json +0 -0
  37. evaluator_ARC_step_682021/submission.json +0 -0
  38. evaluator_ARC_step_73732/submission.json +0 -0
  39. evaluator_ARC_step_92165/submission.json +0 -0
  40. losses.py +103 -0
  41. step_110598 +3 -0
  42. step_129031 +3 -0
  43. step_147464 +3 -0
  44. step_165897 +3 -0
  45. step_18433 +3 -0
  46. step_184330 +3 -0
  47. step_202763 +3 -0
  48. step_221196 +3 -0
  49. step_239629 +3 -0
  50. step_258062 +3 -0
.gitattributes CHANGED
@@ -34,3 +34,40 @@ saved_model/**/* 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
  arc2u-pretrain/identifiers.json 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
  arc2u-pretrain/identifiers.json filter=lfs diff=lfs merge=lfs -text
37
+ step_110598 filter=lfs diff=lfs merge=lfs -text
38
+ step_129031 filter=lfs diff=lfs merge=lfs -text
39
+ step_147464 filter=lfs diff=lfs merge=lfs -text
40
+ step_165897 filter=lfs diff=lfs merge=lfs -text
41
+ step_18433 filter=lfs diff=lfs merge=lfs -text
42
+ step_184330 filter=lfs diff=lfs merge=lfs -text
43
+ step_202763 filter=lfs diff=lfs merge=lfs -text
44
+ step_221196 filter=lfs diff=lfs merge=lfs -text
45
+ step_239629 filter=lfs diff=lfs merge=lfs -text
46
+ step_258062 filter=lfs diff=lfs merge=lfs -text
47
+ step_276495 filter=lfs diff=lfs merge=lfs -text
48
+ step_294928 filter=lfs diff=lfs merge=lfs -text
49
+ step_313361 filter=lfs diff=lfs merge=lfs -text
50
+ step_331794 filter=lfs diff=lfs merge=lfs -text
51
+ step_350227 filter=lfs diff=lfs merge=lfs -text
52
+ step_36866 filter=lfs diff=lfs merge=lfs -text
53
+ step_368660 filter=lfs diff=lfs merge=lfs -text
54
+ step_387093 filter=lfs diff=lfs merge=lfs -text
55
+ step_405526 filter=lfs diff=lfs merge=lfs -text
56
+ step_423959 filter=lfs diff=lfs merge=lfs -text
57
+ step_442392 filter=lfs diff=lfs merge=lfs -text
58
+ step_460825 filter=lfs diff=lfs merge=lfs -text
59
+ step_479258 filter=lfs diff=lfs merge=lfs -text
60
+ step_497691 filter=lfs diff=lfs merge=lfs -text
61
+ step_516124 filter=lfs diff=lfs merge=lfs -text
62
+ step_534557 filter=lfs diff=lfs merge=lfs -text
63
+ step_55299 filter=lfs diff=lfs merge=lfs -text
64
+ step_552990 filter=lfs diff=lfs merge=lfs -text
65
+ step_571423 filter=lfs diff=lfs merge=lfs -text
66
+ step_589856 filter=lfs diff=lfs merge=lfs -text
67
+ step_608289 filter=lfs diff=lfs merge=lfs -text
68
+ step_626722 filter=lfs diff=lfs merge=lfs -text
69
+ step_645155 filter=lfs diff=lfs merge=lfs -text
70
+ step_663588 filter=lfs diff=lfs merge=lfs -text
71
+ step_682021 filter=lfs diff=lfs merge=lfs -text
72
+ step_73732 filter=lfs diff=lfs merge=lfs -text
73
+ step_92165 filter=lfs diff=lfs merge=lfs -text
all_config.yaml ADDED
@@ -0,0 +1,69 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ arch:
2
+ H_cycles: 3
3
+ H_layers: 0
4
+ L_cycles: 4
5
+ L_layers: 2
6
+ expansion: 4
7
+ forward_dtype: bfloat16
8
+ grid_noise_fraction: 0.0
9
+ grid_noise_prob: 0.0
10
+ grid_token_dropout: 0.0
11
+ halt_exploration_prob: 0.1
12
+ halt_max_steps: 16
13
+ halt_max_steps_eval: null
14
+ hidden_size: 512
15
+ loss:
16
+ loss_type: stablemax_cross_entropy
17
+ name: losses@ACTLossHead
18
+ mlp_t: false
19
+ name: recursive_reasoning.trm@TinyRecursiveReasoningModel_ACTV1
20
+ no_ACT_continue: true
21
+ num_heads: 8
22
+ pos_encodings: rope
23
+ puzzle_emb_dropout: 0.0
24
+ puzzle_emb_len: 1
25
+ puzzle_emb_ndim: 512
26
+ beta1: 0.9
27
+ beta2: 0.95
28
+ checkpoint_every_eval: true
29
+ checkpoint_every_n_steps: null
30
+ checkpoint_path: checkpoints/Arc2concept-aug-1000-ACT-torch/pretrain_arc2_rearc_200k
31
+ data_paths:
32
+ - data/rearc-pretrain
33
+ - data/arc2u-pretrain
34
+ data_paths_test:
35
+ - data/arc2u-pretrain
36
+ dataloader_num_workers: 1
37
+ dataloader_persistent_workers: true
38
+ dataloader_pin_memory: true
39
+ dataloader_prefetch_factor: 8
40
+ ema: true
41
+ ema_rate: 0.999
42
+ entity: trelis
43
+ epochs: 100000
44
+ eval_global_batch_size: null
45
+ eval_interval: 2500
46
+ eval_max_augmentations: null
47
+ eval_save_outputs: []
48
+ evaluators:
49
+ - name: arc@ARC
50
+ freeze_weights: false
51
+ freeze_weights_epochs: null
52
+ global_batch_size: 768
53
+ grid_noise_fraction: null
54
+ grid_noise_prob: null
55
+ halt_max_steps_eval: null
56
+ load_checkpoint: checkpoints/Arc2concept-aug-1000-ACT-torch/pretrain_arc2_rearc_100k/step_737320
57
+ lr: 0.0001
58
+ lr_min_ratio: 1.0
59
+ lr_warmup_steps: 2000
60
+ max_eval_examples_per_puzzle: null
61
+ max_examples_per_puzzle: 5
62
+ min_eval_interval: 0
63
+ project_name: Arc2concept-aug-1000-ACT-torch
64
+ puzzle_emb_lr: 0.01
65
+ puzzle_emb_reinit_strategy: mean
66
+ puzzle_emb_weight_decay: 0.1
67
+ run_name: pretrain_arc2_rearc_200k
68
+ seed: 0
69
+ weight_decay: 0.1
evaluator_ARC_step_110598/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_129031/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_147464/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_165897/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_18433/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_184330/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_202763/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_221196/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_239629/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_258062/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_276495/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_294928/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_313361/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_331794/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_350227/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_36866/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_368660/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_387093/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_405526/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_423959/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_442392/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_460825/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_479258/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_497691/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_516124/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_534557/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_55299/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_552990/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_571423/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_589856/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_608289/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_626722/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_645155/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_663588/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_682021/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_73732/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
evaluator_ARC_step_92165/submission.json ADDED
The diff for this file is too large to render. See raw diff
 
losses.py ADDED
@@ -0,0 +1,103 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Any, Tuple, Dict, Sequence, Optional
2
+
3
+ import torch
4
+ import torch.nn.functional as F
5
+ from torch import nn
6
+ import math
7
+
8
+ IGNORE_LABEL_ID = -100
9
+
10
+
11
+ def s(x, epsilon=1e-30):
12
+ return torch.where(
13
+ x<0,
14
+ 1/(1-x+ epsilon),
15
+ x + 1
16
+ )
17
+
18
+
19
+ def log_stablemax(x, dim=-1):
20
+ s_x = s(x)
21
+ return torch.log(s_x/torch.sum(s_x, dim=dim, keepdim=True))
22
+
23
+
24
+ def stablemax_cross_entropy(logits, labels, ignore_index: int = -100, valid_mask=None):
25
+ logprobs = log_stablemax(logits.to(torch.float64), dim=-1)
26
+
27
+ if valid_mask is None:
28
+ valid_mask = (labels != ignore_index)
29
+ transformed_labels = torch.where(valid_mask, labels, 0)
30
+ prediction_logprobs = torch.gather(logprobs, index=transformed_labels.to(torch.long).unsqueeze(-1), dim=-1).squeeze(-1)
31
+
32
+ return -torch.where(valid_mask, prediction_logprobs, 0)
33
+
34
+
35
+ def softmax_cross_entropy(logits, labels, ignore_index: int = -100):
36
+ # Cast logits to f32
37
+ # Flatten logits
38
+ return F.cross_entropy(logits.to(torch.float32).view(-1, logits.shape[-1]), labels.to(torch.long).view(-1), ignore_index=ignore_index, reduction="none").view(labels.shape)
39
+
40
+
41
+ class ACTLossHead(nn.Module):
42
+ def __init__(self, model: nn.Module, loss_type: str):
43
+ super().__init__()
44
+ self.model = model
45
+ self.loss_fn = globals()[loss_type]
46
+
47
+ def initial_carry(self, *args, **kwargs):
48
+ return self.model.initial_carry(*args, **kwargs) # type: ignore
49
+
50
+ def forward(
51
+ self,
52
+ return_keys: Sequence[str],
53
+ # Model args
54
+ **model_kwargs,
55
+ ) -> Tuple[Any, torch.Tensor, Dict[str, torch.Tensor], Optional[Dict[str, torch.Tensor]], torch.Tensor]:
56
+ # Model logits
57
+ # B x SeqLen x D
58
+ new_carry, outputs = self.model(**model_kwargs)
59
+ labels = new_carry.current_data["labels"]
60
+
61
+ with torch.no_grad():
62
+ # Preds
63
+ outputs["preds"] = torch.argmax(outputs["logits"], dim=-1)
64
+
65
+ # Correctness
66
+ mask = (labels != IGNORE_LABEL_ID)
67
+ loss_counts = mask.sum(-1)
68
+ loss_divisor = loss_counts.clamp_min(1).unsqueeze(-1) # Avoid NaNs in division
69
+
70
+ is_correct = mask & (torch.argmax(outputs["logits"], dim=-1) == labels)
71
+ seq_is_correct = is_correct.sum(-1) == loss_counts
72
+
73
+ # Metrics (halted)
74
+ valid_metrics = new_carry.halted & (loss_counts > 0)
75
+ metrics = {
76
+ "count": valid_metrics.sum(),
77
+
78
+ "accuracy": torch.where(valid_metrics, (is_correct.to(torch.float32) / loss_divisor).sum(-1), 0).sum(),
79
+ "exact_accuracy": (valid_metrics & seq_is_correct).sum(),
80
+
81
+ "q_halt_accuracy": (valid_metrics & ((outputs["q_halt_logits"] >= 0) == seq_is_correct)).sum(),
82
+ "steps": torch.where(valid_metrics, new_carry.steps, 0).sum(),
83
+ }
84
+
85
+ # Losses
86
+
87
+ lm_loss = (self.loss_fn(outputs["logits"], labels, ignore_index=IGNORE_LABEL_ID, valid_mask=mask) / loss_divisor).sum()
88
+ q_halt_loss = F.binary_cross_entropy_with_logits(outputs["q_halt_logits"], seq_is_correct.to(outputs["q_halt_logits"].dtype), reduction="sum")
89
+ metrics.update({
90
+ "lm_loss": lm_loss.detach(),
91
+ "q_halt_loss": q_halt_loss.detach(),
92
+ })
93
+ # Q continue (bootstrapping target loss); Alexia: This fits Q-learning, but seems totally unecessary
94
+ q_continue_loss = 0
95
+ if "target_q_continue" in outputs:
96
+ q_continue_loss = F.binary_cross_entropy_with_logits(outputs["q_continue_logits"], outputs["target_q_continue"], reduction="sum")
97
+
98
+ metrics["q_continue_loss"] = q_continue_loss.detach()
99
+ # Filter outputs for return
100
+ detached_outputs = {k: outputs[k].detach() for k in return_keys if k in outputs}
101
+
102
+ return new_carry, lm_loss + 0.5 * (q_halt_loss + q_continue_loss), metrics, detached_outputs, new_carry.halted.all()
103
+
step_110598 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:7edcd45e7ebadd1d10050381a878f04d44ae1dfae2d5ece017669a0d89ebabd2
3
+ size 1748938058
step_129031 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:36e2c64ce5ef980038e178554d3251e03524b115e504901f04cb610fdda297e7
3
+ size 1748938058
step_147464 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:7317cf63718343cf27d84cfbe7c1056966c8759c63a7201c36e4f49fcd39a71b
3
+ size 1748938058
step_165897 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:3946d8eecb83fbf3606e8b94b55647f2e14723e57864c21390b3d67375938632
3
+ size 1748938058
step_18433 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:3a3b0f8ff3f212c9ccf862f765fa35ed8e97a561186210d65dcdc0839daa2628
3
+ size 1748938037
step_184330 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:12e22621c1b20f8cb10032e5bd46f5d2ca266ba2b21bb16ca81c089ead05406d
3
+ size 1748938058
step_202763 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:65defc860cca8f25f810df742041aaa39b4f249c5c8dc756846b9aeeeccd86b0
3
+ size 1748938058
step_221196 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:a8dcc1c95875909546a331a535b58a0c7598e2d1c2923af608b99c5f992aaeff
3
+ size 1748938058
step_239629 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:efee3c38cb1d29778f552de11f1170eb6499518a8905fe118f46e6e8cfe9f31b
3
+ size 1748938058
step_258062 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:d542a33dc419deed6ca2bc6214546a71db5432f46bc24d2b758a2708383be15c
3
+ size 1748938058