Lobakkang commited on
Commit
59412c2
·
verified ·
1 Parent(s): d092869

Upload folder using huggingface_hub

Browse files
__pycache__/export_to_hf.cpython-312.pyc ADDED
Binary file (8.69 kB). View file
 
__pycache__/modeling_taonet.cpython-312.pyc CHANGED
Binary files a/__pycache__/modeling_taonet.cpython-312.pyc and b/__pycache__/modeling_taonet.cpython-312.pyc differ
 
__pycache__/taonet_model.cpython-312.pyc CHANGED
Binary files a/__pycache__/taonet_model.cpython-312.pyc and b/__pycache__/taonet_model.cpython-312.pyc differ
 
__pycache__/tokenization_taonet.cpython-312.pyc CHANGED
Binary files a/__pycache__/tokenization_taonet.cpython-312.pyc and b/__pycache__/tokenization_taonet.cpython-312.pyc differ
 
export_to_hf.py CHANGED
@@ -4,6 +4,13 @@ import json
4
  import shutil
5
  from pathlib import Path
6
 
 
 
 
 
 
 
 
7
 
8
  def normalize_checkpoint(checkpoint):
9
  if isinstance(checkpoint, dict):
@@ -26,6 +33,18 @@ def infer_special_token_paths(repo_dir):
26
  return repo_dir / "tokenizer" / "tokenizer.model", subdir_metadata, repo_dir / "tokenizer" / "tokenizer.vocab"
27
 
28
 
 
 
 
 
 
 
 
 
 
 
 
 
29
  def write_clean_tokenizer_metadata(repo_dir, special_tokens):
30
  tokenizer_config = {
31
  "backend": "custom",
@@ -67,8 +86,6 @@ def write_clean_tokenizer_metadata(repo_dir, special_tokens):
67
 
68
 
69
  def main():
70
- import torch
71
-
72
  from configuration_taonet import TaoNetConfig
73
  from modeling_taonet import TaoNetForCausalLM
74
  from tokenization_taonet import TaoNetTokenizer
@@ -78,6 +95,7 @@ def main():
78
 
79
  checkpoint = torch.load(checkpoint_path, map_location="cpu")
80
  model_state, train_config = normalize_checkpoint(checkpoint)
 
81
  model_config = dict(train_config.get("model", {}))
82
 
83
  metadata_model_path, metadata_path, vocab_path = infer_special_token_paths(repo_dir)
@@ -100,8 +118,43 @@ def main():
100
  }
101
 
102
  model = TaoNetForCausalLM(hf_config)
103
- missing, unexpected = model.model.load_state_dict(model_state, strict=False)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
104
  model.tie_weights()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
105
  model.save_pretrained(repo_dir, safe_serialization=False)
106
 
107
  tokenizer = TaoNetTokenizer(
@@ -115,13 +168,9 @@ def main():
115
  shutil.copyfile(metadata_path, repo_dir / "tokenizer.special_tokens.json")
116
  shutil.copyfile(vocab_path, repo_dir / "tokenizer.vocab")
117
 
118
- if missing:
119
- print("Missing keys while loading model state:")
120
- for key in missing:
121
- print(f" - {key}")
122
- if unexpected:
123
- print("Unexpected keys while loading model state:")
124
- for key in unexpected:
125
  print(f" - {key}")
126
  print(f"Saved Hugging Face package to: {repo_dir}")
127
 
 
4
  import shutil
5
  from pathlib import Path
6
 
7
+ import torch
8
+
9
+
10
+ IGNORED_CHECKPOINT_SUFFIXES = (
11
+ ".rotary.inv_freq",
12
+ )
13
+
14
 
15
  def normalize_checkpoint(checkpoint):
16
  if isinstance(checkpoint, dict):
 
33
  return repo_dir / "tokenizer" / "tokenizer.model", subdir_metadata, repo_dir / "tokenizer" / "tokenizer.vocab"
34
 
35
 
36
+ def sanitize_model_state(model_state):
37
+ """Drop deterministic non-persistent buffers that should not participate in HF weight export."""
38
+ sanitized = {}
39
+ ignored = []
40
+ for key, value in model_state.items():
41
+ if key.endswith(IGNORED_CHECKPOINT_SUFFIXES):
42
+ ignored.append(key)
43
+ continue
44
+ sanitized[key] = value
45
+ return sanitized, ignored
46
+
47
+
48
  def write_clean_tokenizer_metadata(repo_dir, special_tokens):
49
  tokenizer_config = {
50
  "backend": "custom",
 
86
 
87
 
88
  def main():
 
 
89
  from configuration_taonet import TaoNetConfig
90
  from modeling_taonet import TaoNetForCausalLM
91
  from tokenization_taonet import TaoNetTokenizer
 
95
 
96
  checkpoint = torch.load(checkpoint_path, map_location="cpu")
97
  model_state, train_config = normalize_checkpoint(checkpoint)
98
+ model_state, ignored_keys = sanitize_model_state(model_state)
99
  model_config = dict(train_config.get("model", {}))
100
 
101
  metadata_model_path, metadata_path, vocab_path = infer_special_token_paths(repo_dir)
 
118
  }
119
 
120
  model = TaoNetForCausalLM(hf_config)
121
+ current_state = model.model.state_dict()
122
+ missing = sorted(set(current_state) - set(model_state))
123
+ unexpected = sorted(set(model_state) - set(current_state))
124
+ if missing or unexpected:
125
+ if missing:
126
+ print("Missing keys while loading model state:")
127
+ for key in missing:
128
+ print(f" - {key}")
129
+ if unexpected:
130
+ print("Unexpected keys while loading model state:")
131
+ for key in unexpected:
132
+ print(f" - {key}")
133
+ raise ValueError("Checkpoint/model key mismatch detected. Refusing to export a partial model.")
134
+
135
+ model.model.load_state_dict(model_state, strict=True)
136
  model.tie_weights()
137
+
138
+ exported_state = model.model.state_dict()
139
+ mismatched_tensors = []
140
+ for key, value in model_state.items():
141
+ exported_value = exported_state[key]
142
+ if value.shape != exported_value.shape:
143
+ mismatched_tensors.append((key, "shape"))
144
+ continue
145
+ if value.dtype.is_floating_point:
146
+ if not torch.equal(value, exported_value.to(dtype=value.dtype)):
147
+ mismatched_tensors.append((key, "value"))
148
+ else:
149
+ if not torch.equal(value, exported_value):
150
+ mismatched_tensors.append((key, "value"))
151
+
152
+ if mismatched_tensors:
153
+ print("Tensor mismatches detected after strict load:")
154
+ for key, mismatch_type in mismatched_tensors[:20]:
155
+ print(f" - {key} ({mismatch_type})")
156
+ raise ValueError("Checkpoint tensors do not match the HF wrapper after load.")
157
+
158
  model.save_pretrained(repo_dir, safe_serialization=False)
159
 
160
  tokenizer = TaoNetTokenizer(
 
168
  shutil.copyfile(metadata_path, repo_dir / "tokenizer.special_tokens.json")
169
  shutil.copyfile(vocab_path, repo_dir / "tokenizer.vocab")
170
 
171
+ if ignored_keys:
172
+ print("Ignored non-persistent checkpoint buffers:")
173
+ for key in ignored_keys:
 
 
 
 
174
  print(f" - {key}")
175
  print(f"Saved Hugging Face package to: {repo_dir}")
176
 
verify_export_weights.py ADDED
@@ -0,0 +1,82 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Verify that TaoTrain checkpoint weights match the exported HF package exactly."""
2
+
3
+ from pathlib import Path
4
+
5
+ import torch
6
+
7
+ from export_to_hf import normalize_checkpoint, sanitize_model_state
8
+ from modeling_taonet import TaoNetForCausalLM
9
+
10
+
11
+ def summarize_tensor_diff(left: torch.Tensor, right: torch.Tensor) -> str:
12
+ if left.shape != right.shape:
13
+ return f"shape mismatch: {tuple(left.shape)} != {tuple(right.shape)}"
14
+ if left.dtype != right.dtype:
15
+ right = right.to(dtype=left.dtype)
16
+ if left.dtype.is_floating_point:
17
+ diff = (left - right).abs()
18
+ return f"max_abs_diff={diff.max().item():.8g}, mean_abs_diff={diff.mean().item():.8g}"
19
+ unequal = (left != right).sum().item()
20
+ return f"unequal_values={unequal}"
21
+
22
+
23
+ def main():
24
+ repo_dir = Path(__file__).resolve().parent
25
+ checkpoint_path = repo_dir / "checkpoints" / "sft" / "final_model.pt"
26
+
27
+ checkpoint = torch.load(checkpoint_path, map_location="cpu")
28
+ checkpoint_state, _ = normalize_checkpoint(checkpoint)
29
+ checkpoint_state, ignored_keys = sanitize_model_state(checkpoint_state)
30
+
31
+ model = TaoNetForCausalLM.from_pretrained(str(repo_dir))
32
+ exported_state = model.model.state_dict()
33
+
34
+ checkpoint_keys = set(checkpoint_state)
35
+ exported_keys = set(exported_state)
36
+
37
+ missing = sorted(exported_keys - checkpoint_keys)
38
+ unexpected = sorted(checkpoint_keys - exported_keys)
39
+
40
+ print(f"checkpoint tensors: {len(checkpoint_keys)}")
41
+ print(f"exported tensors: {len(exported_keys)}")
42
+ if ignored_keys:
43
+ print(f"ignored checkpoint buffers: {len(ignored_keys)}")
44
+
45
+ if missing:
46
+ print("\nKeys present in exported model but missing from checkpoint:")
47
+ for key in missing[:50]:
48
+ print(f" - {key}")
49
+
50
+ if unexpected:
51
+ print("\nKeys present in checkpoint but missing from exported model:")
52
+ for key in unexpected[:50]:
53
+ print(f" - {key}")
54
+
55
+ common_keys = sorted(checkpoint_keys & exported_keys)
56
+ mismatches = []
57
+ exact_matches = 0
58
+ for key in common_keys:
59
+ left = checkpoint_state[key].cpu()
60
+ right = exported_state[key].cpu()
61
+ if left.shape == right.shape and torch.equal(left, right.to(dtype=left.dtype) if right.dtype != left.dtype else right):
62
+ exact_matches += 1
63
+ continue
64
+ mismatches.append((key, summarize_tensor_diff(left, right)))
65
+
66
+ print(f"\nexact tensor matches: {exact_matches}/{len(common_keys)}")
67
+ print(f"tensor mismatches: {len(mismatches)}")
68
+
69
+ if mismatches:
70
+ print("\nFirst mismatches:")
71
+ for key, summary in mismatches[:50]:
72
+ print(f" - {key}: {summary}")
73
+ raise SystemExit(1)
74
+
75
+ if missing or unexpected:
76
+ raise SystemExit(1)
77
+
78
+ print("\nWeight verification passed.")
79
+
80
+
81
+ if __name__ == "__main__":
82
+ main()