E6E831728 commited on
Commit
21a2150
·
verified ·
1 Parent(s): 458999f

Delete training_original

Browse files
training_original/__pycache__/classic_model.cpython-311.pyc DELETED
Binary file (25.6 kB)
 
training_original/__pycache__/model_16_dim_bin.cpython-311.pyc DELETED
Binary file (9.1 kB)
 
training_original/binary16_train.py DELETED
@@ -1,997 +0,0 @@
1
- #!/usr/bin/env python3
2
-
3
- import argparse
4
- import contextlib
5
- import json
6
- import math
7
- import os
8
- import random
9
- import time
10
- from pathlib import Path
11
-
12
- import numpy as np
13
- import torch
14
- import torch.distributed as dist
15
- from torch.nn.parallel import DistributedDataParallel as DDP
16
- from transformers import AutoTokenizer
17
-
18
- from classic_train import (
19
- Logger,
20
- TokenShardStore,
21
- all_reduce_sum,
22
- cleanup_distributed,
23
- configure_optimizer,
24
- estimate_loss,
25
- generate_diagnostics,
26
- log_input_representation_diagnostics,
27
- get_learning_rate,
28
- get_rank,
29
- get_world_size,
30
- human_tokens,
31
- is_main_process,
32
- save_model_safetensors,
33
- load_checkpoint,
34
- save_checkpoint,
35
- setup_distributed,
36
- timestamp,
37
- )
38
-
39
- from model_16_dim_bin import (
40
- Binary16Config,
41
- Binary16ForCausalLM,
42
- )
43
-
44
-
45
- def parse_args():
46
- parser = argparse.ArgumentParser()
47
-
48
- # Paths
49
- parser.add_argument("--data_dir", required=True)
50
- parser.add_argument("--output_dir", required=True)
51
- parser.add_argument(
52
- "--tokenizer",
53
- default="HuggingFaceTB/SmolLM2-135M",
54
- )
55
- parser.add_argument("--tokenizer_revision", default=None)
56
- parser.add_argument("--resume", default=None)
57
-
58
- # Architecture: держать идентично classic run.
59
- parser.add_argument("--d_model", type=int, default=960)
60
- parser.add_argument("--n_layer", type=int, default=24)
61
- parser.add_argument("--n_head", type=int, default=15)
62
- parser.add_argument(
63
- "--ffn_multiplier",
64
- type=float,
65
- default=8.0 / 3.0,
66
- )
67
- parser.add_argument("--multiple_of", type=int, default=256)
68
- parser.add_argument("--sequence_length", type=int, default=2048)
69
- parser.add_argument("--rope_theta", type=float, default=10000.0)
70
- parser.add_argument("--dropout", type=float, default=0.0)
71
- parser.add_argument("--rms_norm_eps", type=float, default=1e-5)
72
- parser.add_argument(
73
- "--gradient_checkpointing",
74
- action="store_true",
75
- )
76
- parser.add_argument("--compile", action="store_true")
77
-
78
- # Binary input
79
- parser.add_argument(
80
- "--binary_encoding",
81
- choices=("zero_one", "bipolar"),
82
- default="zero_one",
83
- help=(
84
- "zero_one reproduces the original 0/1 experiment. "
85
- "Do not mix encodings inside the main comparison."
86
- ),
87
- )
88
- parser.add_argument(
89
- "--binary_scale",
90
- type=float,
91
- default=1.0,
92
- )
93
- parser.add_argument(
94
- "--binary_permutation_seed",
95
- type=int,
96
- default=None,
97
- help=(
98
- "None gives canonical token-ID bits. A seed creates a fixed "
99
- "injective random reassignment of token IDs to 16-bit codes."
100
- ),
101
- )
102
-
103
- # Training
104
- parser.add_argument(
105
- "--micro_batch_size",
106
- type=int,
107
- default=2,
108
- )
109
- parser.add_argument(
110
- "--gradient_accumulation_steps",
111
- type=int,
112
- default=16,
113
- )
114
- parser.add_argument("--max_steps", type=int, default=200_000)
115
- parser.add_argument(
116
- "--max_tokens",
117
- type=int,
118
- default=0,
119
- )
120
- parser.add_argument(
121
- "--learning_rate",
122
- type=float,
123
- default=3e-4,
124
- )
125
- parser.add_argument(
126
- "--min_learning_rate",
127
- type=float,
128
- default=3e-5,
129
- )
130
- parser.add_argument("--warmup_steps", type=int, default=2000)
131
- parser.add_argument("--weight_decay", type=float, default=0.1)
132
- parser.add_argument("--beta1", type=float, default=0.9)
133
- parser.add_argument("--beta2", type=float, default=0.95)
134
- parser.add_argument("--grad_clip", type=float, default=1.0)
135
- parser.add_argument("--seed", type=int, default=42)
136
-
137
- # Dataset cache
138
- parser.add_argument(
139
- "--train_cache_shards",
140
- type=int,
141
- default=2,
142
- )
143
- parser.add_argument(
144
- "--valid_cache_shards",
145
- type=int,
146
- default=2,
147
- )
148
-
149
- # Diagnostics
150
- parser.add_argument(
151
- "--diagnostic_interval_seconds",
152
- type=float,
153
- default=3600.0,
154
- )
155
- parser.add_argument(
156
- "--log_interval_steps",
157
- type=int,
158
- default=20,
159
- )
160
- parser.add_argument(
161
- "--eval_batches",
162
- type=int,
163
- default=32,
164
- )
165
- parser.add_argument(
166
- "--eval_batch_size",
167
- type=int,
168
- default=2,
169
- )
170
- parser.add_argument(
171
- "--generation_tokens",
172
- type=int,
173
- default=24,
174
- )
175
- parser.add_argument(
176
- "--save_interval_steps",
177
- type=int,
178
- default=2000,
179
- )
180
- parser.add_argument(
181
- "--batches_per_shard",
182
- type=int,
183
- default=256,
184
- help=(
185
- "Number of microbatches sampled from one resident shard "
186
- "before loading the next shard."
187
- ),
188
- )
189
- parser.add_argument(
190
- "--latest_save_interval_steps",
191
- type=int,
192
- default=2000,
193
- )
194
-
195
- parser.add_argument(
196
- "--milestone_save_interval_steps",
197
- type=int,
198
- default=20000,
199
- )
200
-
201
- return parser.parse_args()
202
-
203
-
204
- def verify_binary_interface(
205
- model: Binary16ForCausalLM,
206
- vocab_size: int,
207
- device: torch.device,
208
- ):
209
- embedding = model.token_embeddings
210
-
211
- if any(True for _ in embedding.parameters()):
212
- raise RuntimeError(
213
- "FixedBinary16Embedding unexpectedly contains parameters"
214
- )
215
-
216
- expected_shape = (vocab_size, 16)
217
-
218
- if tuple(embedding.codebook.shape) != expected_shape:
219
- raise RuntimeError(
220
- f"Invalid codebook shape: "
221
- f"{tuple(embedding.codebook.shape)} != {expected_shape}"
222
- )
223
-
224
- if embedding.codebook.requires_grad:
225
- raise RuntimeError("Binary codebook must not require gradients")
226
-
227
- # Проверяем первые canonical codes только если permutation выключена.
228
- if model.config.binary_permutation_seed is None:
229
- expected = torch.tensor(
230
- [
231
- [0, 0, 0, 0, 0, 0, 0, 0,
232
- 0, 0, 0, 0, 0, 0, 0, 0],
233
- [1, 0, 0, 0, 0, 0, 0, 0,
234
- 0, 0, 0, 0, 0, 0, 0, 0],
235
- [0, 1, 0, 0, 0, 0, 0, 0,
236
- 0, 0, 0, 0, 0, 0, 0, 0],
237
- [1, 1, 0, 0, 0, 0, 0, 0,
238
- 0, 0, 0, 0, 0, 0, 0, 0],
239
- ],
240
- dtype=torch.float32,
241
- device=embedding.codebook.device,
242
- )
243
-
244
- actual = embedding.codebook[:4]
245
-
246
- if model.config.binary_encoding == "bipolar":
247
- expected = expected.mul(2).sub(1)
248
-
249
- if not torch.equal(actual, expected):
250
- raise RuntimeError(
251
- "Canonical binary code sanity check failed"
252
- )
253
-
254
- ids = torch.tensor(
255
- [[0, min(1, vocab_size - 1), vocab_size - 1]],
256
- device=device,
257
- dtype=torch.long,
258
- )
259
-
260
- output = embedding(ids)
261
-
262
- expected_output_shape = (
263
- 1,
264
- 3,
265
- model.config.d_model,
266
- )
267
-
268
- if tuple(output.shape) != expected_output_shape:
269
- raise RuntimeError(
270
- f"Invalid lifted shape: "
271
- f"{tuple(output.shape)} != {expected_output_shape}"
272
- )
273
-
274
-
275
- def find_nonfinite_gradients(model, max_names=20):
276
- bad = []
277
-
278
- for name, parameter in model.named_parameters():
279
- gradient = parameter.grad
280
-
281
- if gradient is None:
282
- continue
283
-
284
- finite = torch.isfinite(gradient)
285
-
286
- if not bool(finite.all()):
287
- bad.append(
288
- {
289
- "name": name,
290
- "shape": tuple(gradient.shape),
291
- "dtype": str(gradient.dtype),
292
- "nan": int(torch.isnan(gradient).sum().item()),
293
- "inf": int(torch.isinf(gradient).sum().item()),
294
- }
295
- )
296
-
297
- if len(bad) >= max_names:
298
- break
299
-
300
- return bad
301
-
302
-
303
- def grad_norm_for_named_parameters(named_parameters):
304
- squares = []
305
-
306
- for _, parameter in named_parameters:
307
- if parameter.grad is None:
308
- continue
309
-
310
- gradient = parameter.grad.detach().float()
311
- squares.append(gradient.pow(2).sum())
312
-
313
- if not squares:
314
- return 0.0
315
-
316
- return torch.sqrt(torch.stack(squares).sum()).item()
317
-
318
-
319
- def main():
320
- args = parse_args()
321
-
322
- distributed, local_rank, device = setup_distributed()
323
- rank = get_rank()
324
- world_size = get_world_size()
325
-
326
- if device.type != "cuda":
327
- raise RuntimeError("CUDA is required")
328
-
329
- torch.backends.cuda.matmul.allow_tf32 = True
330
- torch.backends.cudnn.allow_tf32 = True
331
- torch.backends.cuda.enable_flash_sdp(True)
332
- torch.backends.cuda.enable_mem_efficient_sdp(True)
333
- torch.backends.cuda.enable_math_sdp(True)
334
-
335
- seed = args.seed + rank
336
- random.seed(seed)
337
- np.random.seed(seed)
338
- torch.manual_seed(seed)
339
- torch.cuda.manual_seed_all(seed)
340
-
341
- Path(args.output_dir).mkdir(
342
- parents=True,
343
- exist_ok=True,
344
- )
345
-
346
- logger = Logger(
347
- os.path.join(args.output_dir, "train.log")
348
- )
349
-
350
- tokenizer = AutoTokenizer.from_pretrained(
351
- args.tokenizer,
352
- revision=args.tokenizer_revision,
353
- use_fast=True,
354
- )
355
-
356
- vocab_size = len(tokenizer)
357
-
358
- if vocab_size > 65536:
359
- raise ValueError(
360
- f"Tokenizer length {vocab_size} exceeds 16-bit capacity"
361
- )
362
-
363
- if args.d_model % 16 != 0:
364
- raise ValueError(
365
- "d_model must be divisible by 16 for parameter-free tiling"
366
- )
367
-
368
- config = Binary16Config(
369
- vocab_size=vocab_size,
370
- d_model=args.d_model,
371
- n_layer=args.n_layer,
372
- n_head=args.n_head,
373
- ffn_multiplier=args.ffn_multiplier,
374
- multiple_of=args.multiple_of,
375
- block_size=args.sequence_length,
376
- rope_theta=args.rope_theta,
377
- dropout=args.dropout,
378
- rms_norm_eps=args.rms_norm_eps,
379
- pad_token_id=tokenizer.pad_token_id,
380
- bos_token_id=tokenizer.bos_token_id,
381
- eos_token_id=tokenizer.eos_token_id,
382
- tie_word_embeddings=False,
383
- binary_dim=16,
384
- binary_encoding=args.binary_encoding,
385
- binary_scale=args.binary_scale,
386
- binary_permutation_seed=args.binary_permutation_seed,
387
- )
388
-
389
- raw_model = Binary16ForCausalLM(config)
390
-
391
- if args.gradient_checkpointing:
392
- raw_model.gradient_checkpointing = True
393
-
394
- raw_model.to(device)
395
-
396
- if is_main_process():
397
- first_parameter = next(raw_model.parameters())
398
-
399
- logger.log(
400
- "[dtype] "
401
- f"parameter_dtype={first_parameter.dtype}, "
402
- f"optimizer_master_expected=float32"
403
- )
404
-
405
- if next(raw_model.parameters()).dtype != torch.float32:
406
- raise RuntimeError(
407
- "Model parameters must remain FP32; "
408
- "BF16 should be enabled only through autocast"
409
- )
410
-
411
- verify_binary_interface(
412
- model=raw_model,
413
- vocab_size=vocab_size,
414
- device=device,
415
- )
416
-
417
- parameter_counts = raw_model.count_parameters()
418
-
419
- if is_main_process():
420
- logger.log(
421
- "[model] "
422
- + ", ".join(
423
- f"{key}={value:,}"
424
- for key, value in parameter_counts.items()
425
- )
426
- )
427
- logger.log(
428
- "[binary_input] "
429
- f"encoding={args.binary_encoding}, "
430
- f"scale={args.binary_scale}, "
431
- f"permutation_seed={args.binary_permutation_seed}, "
432
- f"codebook_shape={tuple(raw_model.token_embeddings.codebook.shape)}, "
433
- f"codebook_is_parameter=False"
434
- )
435
- logger.log(
436
- "[config] "
437
- + json.dumps(
438
- config.to_dict(),
439
- ensure_ascii=False,
440
- sort_keys=True,
441
- )
442
- )
443
- logger.log(
444
- f"[run] world_size={world_size}, "
445
- f"micro_batch={args.micro_batch_size}, "
446
- f"grad_accum={args.gradient_accumulation_steps}, "
447
- f"seq={args.sequence_length}, "
448
- f"global_tokens_per_step="
449
- #f"{world_size * args.micro_batch_size * args.gradient_accumulation_steps * (args.sequence_length - 1):,}"
450
- f"{world_size * args.micro_batch_size * args.gradient_accumulation_steps * args.sequence_length:,}"
451
- )
452
-
453
- optimizer = configure_optimizer(
454
- raw_model,
455
- learning_rate=args.learning_rate,
456
- weight_decay=args.weight_decay,
457
- betas=(args.beta1, args.beta2),
458
- fused=True,
459
- )
460
-
461
- start_step = 0
462
- tokens_seen = 0
463
-
464
- if args.resume is not None:
465
- start_step, tokens_seen = load_checkpoint(
466
- args.resume,
467
- raw_model,
468
- optimizer,
469
- device,
470
- )
471
-
472
- logger.log(
473
- f"[resume] path={args.resume}, "
474
- f"step={start_step}, "
475
- f"tokens_seen={tokens_seen:,}"
476
- )
477
-
478
- model = raw_model
479
-
480
- if args.compile:
481
- model = torch.compile(
482
- model,
483
- mode="max-autotune",
484
- dynamic=False,
485
- )
486
-
487
- if distributed:
488
- '''model = DDP(
489
- model,
490
- device_ids=[local_rank],
491
- output_device=local_rank,
492
- broadcast_buffers=False,
493
- gradient_as_bucket_view=True,
494
- static_graph=not args.gradient_checkpointing,
495
- )'''
496
- model = DDP(
497
- model,
498
- device_ids=[local_rank],
499
- output_device=local_rank,
500
- broadcast_buffers=False,
501
- gradient_as_bucket_view=True,
502
- static_graph=False,
503
- find_unused_parameters=False,
504
- )
505
-
506
- train_data = TokenShardStore(
507
- data_dir=args.data_dir,
508
- split="train",
509
- sequence_length=args.sequence_length,
510
- cache_shards=args.train_cache_shards,
511
- seed=args.seed,
512
- batches_per_shard=args.batches_per_shard,
513
- )
514
-
515
- valid_data = TokenShardStore(
516
- data_dir=args.data_dir,
517
- split="valid",
518
- sequence_length=args.sequence_length,
519
- cache_shards=args.valid_cache_shards,
520
- seed=args.seed + 10_000,
521
- batches_per_shard=max(
522
- args.batches_per_shard,
523
- args.eval_batches,
524
- ),
525
- )
526
-
527
- train_generator = torch.Generator(device="cpu")
528
- train_generator.manual_seed(
529
- args.seed + rank * 100_003
530
- )
531
-
532
- prompts = [
533
- # English factual completion
534
- "London is the capital of",
535
- "The capital of France is",
536
- "The largest planet in the Solar System is",
537
- "Water freezes at",
538
- "The chemical symbol for gold is",
539
- "The Pacific Ocean is",
540
- "The human heart pumps",
541
- "The Second World War ended in",
542
- "The author of Romeo and Juliet was",
543
- "A triangle has",
544
-
545
- # English continuation and grammar
546
- "Once upon a time, there was",
547
- "The scientist opened the laboratory door and",
548
- "When the rain finally stopped,",
549
- "She went to the store because",
550
- "If I had known about the problem,",
551
- "The old house on the hill",
552
- "Although the experiment failed,",
553
- "In order to solve this problem, we need to",
554
- "The main difference between cats and dogs is",
555
- "This article explains how to",
556
-
557
- # Definitions and explanations
558
- "Photosynthesis is the process by which",
559
- "Gravity is a force that",
560
- "A computer program is",
561
- "Democracy can be defined as",
562
- "Machine learning is used to",
563
- "The purpose of a database is to",
564
- "An ecosystem consists of",
565
- "Inflation occurs when",
566
- "The Internet allows people to",
567
- "Energy cannot be created or destroyed, but",
568
-
569
- # Arithmetic and symbolic patterns
570
- "2 + 2 =",
571
- "10 - 3 =",
572
- "6 * 7 =",
573
- "12 / 4 =",
574
- "1, 2, 3, 4,",
575
- "2, 4, 6, 8,",
576
- "The square root of 9 is",
577
- "If x = 5, then x + 2 =",
578
- "One hundred divided by ten equals",
579
- "The next number after 99 is",
580
- ]
581
-
582
- interval_loss_sum = 0.0
583
- interval_loss_tokens = 0
584
- interval_start_tokens = tokens_seen
585
- interval_start_time = time.monotonic()
586
- last_diagnostic_time = time.monotonic()
587
-
588
- if is_main_process():
589
- logger.log(f"[start] {timestamp()}")
590
-
591
- log_input_representation_diagnostics(
592
- raw_model=raw_model,
593
- tokenizer=tokenizer,
594
- logger=logger,
595
- device=device,
596
- step=start_step,
597
- )
598
-
599
- model.train()
600
- optimizer.zero_grad(set_to_none=True)
601
-
602
- for step in range(start_step, args.max_steps):
603
- lr = get_learning_rate(
604
- step=step,
605
- max_steps=args.max_steps,
606
- warmup_steps=args.warmup_steps,
607
- learning_rate=args.learning_rate,
608
- min_learning_rate=args.min_learning_rate,
609
- )
610
-
611
- for group in optimizer.param_groups:
612
- group["lr"] = lr
613
-
614
- step_loss_sum = 0.0
615
- step_loss_tokens = 0
616
-
617
- for micro_step in range(
618
- args.gradient_accumulation_steps
619
- ):
620
- batch = train_data.sample_batch(
621
- batch_size=args.micro_batch_size,
622
- generator=train_generator,
623
- ).to(device, non_blocking=True)
624
-
625
- # forward() сам сдвигает logits и labels на один токен.
626
- #input_ids = batch[:, :-1]
627
- #labels = input_ids
628
-
629
- input_ids = batch[:, :-1]
630
- labels = batch[:, 1:]
631
-
632
- should_sync = (
633
- micro_step
634
- == args.gradient_accumulation_steps - 1
635
- )
636
-
637
- sync_context = contextlib.nullcontext()
638
-
639
- if distributed and not should_sync:
640
- sync_context = model.no_sync()
641
-
642
- with sync_context:
643
- with torch.autocast(
644
- device_type="cuda",
645
- dtype=torch.bfloat16,
646
- ):
647
- outputs = model(
648
- input_ids=input_ids,
649
- labels=labels,
650
- return_dict=True,
651
- )
652
-
653
- loss = (
654
- outputs.loss
655
- / args.gradient_accumulation_steps
656
- )
657
-
658
- loss.backward()
659
-
660
- # Первый label не имеет предшествующего logit после внутреннего shift.
661
-
662
- #local_tokens = labels[:, 1:].numel()
663
-
664
- local_tokens = labels.numel()
665
-
666
- step_loss_sum += (
667
- float(outputs.loss.detach())
668
- * local_tokens
669
- )
670
- step_loss_tokens += local_tokens
671
-
672
- bad_gradients = find_nonfinite_gradients(raw_model)
673
-
674
- local_bad = torch.tensor(
675
- [1 if bad_gradients else 0],
676
- device=device,
677
- dtype=torch.int32,
678
- )
679
-
680
- if distributed:
681
- dist.all_reduce(
682
- local_bad,
683
- op=dist.ReduceOp.MAX,
684
- )
685
-
686
- if int(local_bad.item()) != 0:
687
- if bad_gradients:
688
- logger.log(
689
- "[nonfinite_gradients] "
690
- + json.dumps(
691
- bad_gradients,
692
- ensure_ascii=False,
693
- )
694
- )
695
-
696
- emergency_path = os.path.join(
697
- args.output_dir,
698
- f"checkpoint_nonfinite_step_{step + 1:07d}.pt",
699
- )
700
-
701
- save_checkpoint(
702
- path=emergency_path,
703
- raw_model=raw_model,
704
- optimizer=optimizer,
705
- step=step,
706
- tokens_seen=tokens_seen,
707
- args=args,
708
- )
709
-
710
- raise RuntimeError(
711
- f"Non-finite gradients at step {step + 1}"
712
- )
713
-
714
- input_norm = grad_norm_for_named_parameters(
715
- raw_model.token_embeddings.named_parameters()
716
- )
717
-
718
- body_norm = grad_norm_for_named_parameters(
719
- (
720
- (name, parameter)
721
- for name, parameter in raw_model.named_parameters()
722
- if not name.startswith("token_embeddings.")
723
- and not name.startswith("lm_head.")
724
- )
725
- )
726
-
727
- output_norm = grad_norm_for_named_parameters(
728
- raw_model.lm_head.named_parameters()
729
- )
730
-
731
- #logger.log(
732
- # f"[grad_groups] input={input_norm:.3f}, "
733
- # f"body={body_norm:.3f}, output={output_norm:.3f}"
734
- #)
735
-
736
- if args.grad_clip > 0:
737
- grad_norm = torch.nn.utils.clip_grad_norm_(
738
- raw_model.parameters(),
739
- args.grad_clip,
740
- error_if_nonfinite=True,
741
- )
742
- else:
743
- grad_norm = torch.tensor(
744
- float("nan"),
745
- device=device,
746
- )
747
-
748
- optimizer.step()
749
- optimizer.zero_grad(set_to_none=True)
750
-
751
- global_step_tokens = step_loss_tokens * world_size
752
- tokens_seen += global_step_tokens
753
-
754
- loss_stats = torch.tensor(
755
- [step_loss_sum, step_loss_tokens],
756
- device=device,
757
- dtype=torch.float64,
758
- )
759
- all_reduce_sum(loss_stats)
760
-
761
- global_loss_sum = float(loss_stats[0].item())
762
- global_loss_tokens = int(loss_stats[1].item())
763
-
764
- interval_loss_sum += global_loss_sum
765
- interval_loss_tokens += global_loss_tokens
766
-
767
- completed_step = step + 1
768
-
769
- if (
770
- completed_step % args.log_interval_steps == 0
771
- and is_main_process()
772
- ):
773
- mean_step_loss = (
774
- global_loss_sum / global_loss_tokens
775
- )
776
-
777
- logger.log(
778
- f"step {completed_step:07d}: "
779
- f"loss {mean_step_loss:.4f}, "
780
- f"lr {lr:.8f}, "
781
- f"grad_norm {float(grad_norm):.4f}, "
782
- f"tokens_seen={human_tokens(tokens_seen)}, "
783
- f"{timestamp()}"
784
- )
785
-
786
- now = time.monotonic()
787
-
788
- '''diagnostic_due = (
789
- now - last_diagnostic_time
790
- >= args.diagnostic_interval_seconds
791
- )'''
792
-
793
- diagnostic_due = (
794
- completed_step % 700 == 0
795
- )
796
-
797
- final_step = completed_step == args.max_steps
798
-
799
- token_budget_reached = (
800
- args.max_tokens > 0
801
- and tokens_seen >= args.max_tokens
802
- )
803
-
804
- if diagnostic_due or final_step or token_budget_reached:
805
- if distributed:
806
- dist.barrier()
807
-
808
- elapsed = max(
809
- now - interval_start_time,
810
- 1e-9,
811
- )
812
-
813
- interval_tokens = (
814
- tokens_seen - interval_start_tokens
815
- )
816
-
817
- tokens_per_second = (
818
- interval_tokens / elapsed
819
- )
820
-
821
- train_loss = (
822
- interval_loss_sum
823
- / max(interval_loss_tokens, 1)
824
- )
825
-
826
- val_loss, val_tokens = estimate_loss(
827
- model=model,
828
- dataset=valid_data,
829
- device=device,
830
- batch_size=args.eval_batch_size,
831
- eval_batches=args.eval_batches,
832
- seed=args.seed + 123_456,
833
- )
834
-
835
- if is_main_process():
836
- train_ppl = math.exp(
837
- min(train_loss, 20.0)
838
- )
839
-
840
- val_ppl = math.exp(
841
- min(val_loss, 20.0)
842
- )
843
-
844
- logger.log(
845
- f"step {completed_step:07d}: "
846
- f"train loss {train_loss:.4f}, "
847
- f"val loss {val_loss:.4f}, "
848
- f"Train PPL {train_ppl:.3f}, "
849
- f"Val PPL {val_ppl:.3f}, "
850
- f"tokens_seen={human_tokens(tokens_seen)}, "
851
- f"val_tokens={human_tokens(val_tokens)}, "
852
- f"toks/s={tokens_per_second:,.0f}, "
853
- f"{timestamp()}"
854
- )
855
-
856
- should_log_input_repr = (
857
- completed_step == 1
858
- or completed_step % args.milestone_save_interval_steps == 0
859
- or final_step
860
- or token_budget_reached
861
- )
862
-
863
- if should_log_input_repr:
864
- log_input_representation_diagnostics(
865
- raw_model=raw_model,
866
- tokenizer=tokenizer,
867
- logger=logger,
868
- device=device,
869
- step=completed_step,
870
- )
871
-
872
- greedy_prompts = prompts
873
- sample_prompts = prompts
874
-
875
- greedy_generations = generate_diagnostics(
876
- raw_model=raw_model,
877
- tokenizer=tokenizer,
878
- prompts=greedy_prompts,
879
- device=device,
880
- max_new_tokens=args.generation_tokens,
881
- temperature=0.0,
882
- top_k=None,
883
- )
884
-
885
- sample_generations = generate_diagnostics(
886
- raw_model=raw_model,
887
- tokenizer=tokenizer,
888
- prompts=sample_prompts,
889
- device=device,
890
- max_new_tokens=args.generation_tokens,
891
- temperature=0.8,
892
- top_k=50,
893
- )
894
-
895
- for prompt, text in zip(greedy_prompts, greedy_generations):
896
- logger.log(
897
- f"step {completed_step:07d}: "
898
- f"[greedy] prompt={prompt!r} => {text}"
899
- )
900
-
901
- for prompt, text in zip(sample_prompts, sample_generations):
902
- logger.log(
903
- f"step {completed_step:07d}: "
904
- f"[sample t=0.8 k=50] prompt={prompt!r} => {text}"
905
- )
906
-
907
- logger.log(
908
- f"step {completed_step:07d}: "
909
- f"LR: {lr:.8f}, "
910
- f"opt_step: {completed_step}, "
911
- f"{timestamp()}"
912
- )
913
-
914
- if distributed:
915
- dist.barrier()
916
-
917
- last_diagnostic_time = now
918
- interval_start_time = now
919
- interval_start_tokens = tokens_seen
920
- interval_loss_sum = 0.0
921
- interval_loss_tokens = 0
922
-
923
- save_latest = (
924
- completed_step % args.latest_save_interval_steps == 0
925
- or final_step
926
- or token_budget_reached
927
- )
928
-
929
- save_milestone = (
930
- completed_step % args.milestone_save_interval_steps == 0
931
- or final_step
932
- or token_budget_reached
933
- )
934
-
935
- if save_latest or save_milestone:
936
- if distributed:
937
- dist.barrier()
938
-
939
- if save_latest:
940
- save_checkpoint(
941
- path=os.path.join(
942
- args.output_dir,
943
- "checkpoint_latest.pt",
944
- ),
945
- raw_model=raw_model,
946
- optimizer=optimizer,
947
- step=completed_step,
948
- tokens_seen=tokens_seen,
949
- args=args,
950
- )
951
-
952
- save_model_safetensors(
953
- path=os.path.join(
954
- args.output_dir,
955
- "model_latest.safetensors",
956
- ),
957
- raw_model=raw_model,
958
- )
959
-
960
- if save_milestone:
961
- save_checkpoint(
962
- path=os.path.join(
963
- args.output_dir,
964
- f"checkpoint_{completed_step:07d}.pt",
965
- ),
966
- raw_model=raw_model,
967
- optimizer=optimizer,
968
- step=completed_step,
969
- tokens_seen=tokens_seen,
970
- args=args,
971
- )
972
-
973
- save_model_safetensors(
974
- path=os.path.join(
975
- args.output_dir,
976
- f"model_{completed_step:07d}.safetensors",
977
- ),
978
- raw_model=raw_model,
979
- )
980
-
981
- if save_latest or save_milestone:
982
- if distributed:
983
- dist.barrier()
984
-
985
- if token_budget_reached:
986
- if is_main_process():
987
- logger.log(
988
- "[stop] max_tokens reached: "
989
- f"{tokens_seen:,}"
990
- )
991
- break
992
-
993
- cleanup_distributed()
994
-
995
-
996
- if __name__ == "__main__":
997
- main()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
training_original/classic_model.py DELETED
@@ -1,594 +0,0 @@
1
- #!/usr/bin/env python3
2
-
3
- import math
4
- from typing import Optional
5
-
6
- import torch
7
- import torch.nn as nn
8
- import torch.nn.functional as F
9
-
10
- from transformers import PretrainedConfig, PreTrainedModel
11
- from transformers.modeling_outputs import CausalLMOutput
12
-
13
-
14
- class ClassicConfig(PretrainedConfig):
15
- model_type = "classic_causal_lm"
16
-
17
- def __init__(
18
- self,
19
- vocab_size: int = 49152,
20
- d_model: int = 960,
21
- n_layer: int = 24,
22
- n_head: int = 15,
23
- ffn_multiplier: float = 8.0 / 3.0,
24
- multiple_of: int = 256,
25
- block_size: int = 2048,
26
- rope_theta: float = 10000.0,
27
- dropout: float = 0.0,
28
- rms_norm_eps: float = 1e-5,
29
- initializer_range: float = 0.02,
30
- pad_token_id: Optional[int] = None,
31
- bos_token_id: Optional[int] = None,
32
- eos_token_id: Optional[int] = None,
33
- tie_word_embeddings: bool = False,
34
- attention_bias: bool = False,
35
- mlp_bias: bool = False,
36
- **kwargs,
37
- ):
38
- super().__init__(
39
- pad_token_id=pad_token_id,
40
- bos_token_id=bos_token_id,
41
- eos_token_id=eos_token_id,
42
- tie_word_embeddings=tie_word_embeddings,
43
- **kwargs,
44
- )
45
-
46
- if d_model % n_head != 0:
47
- raise ValueError("d_model must be divisible by n_head")
48
-
49
- head_dim = d_model // n_head
50
- if head_dim % 2 != 0:
51
- raise ValueError("RoPE requires an even head_dim")
52
-
53
- self.vocab_size = vocab_size
54
- self.d_model = d_model
55
- self.hidden_size = d_model
56
- self.n_layer = n_layer
57
- self.num_hidden_layers = n_layer
58
- self.n_head = n_head
59
- self.num_attention_heads = n_head
60
- self.head_dim = head_dim
61
-
62
- self.ffn_multiplier = ffn_multiplier
63
- self.multiple_of = multiple_of
64
- self.block_size = block_size
65
- self.max_position_embeddings = block_size
66
- self.rope_theta = rope_theta
67
-
68
- self.dropout = dropout
69
- self.rms_norm_eps = rms_norm_eps
70
- self.initializer_range = initializer_range
71
- self.attention_bias = attention_bias
72
- self.mlp_bias = mlp_bias
73
- self.use_cache = False
74
-
75
-
76
- class RMSNorm(nn.Module):
77
- def __init__(self, dim: int, eps: float):
78
- super().__init__()
79
- self.weight = nn.Parameter(torch.ones(dim))
80
- self.eps = eps
81
-
82
- def forward(self, x: torch.Tensor) -> torch.Tensor:
83
- dtype = x.dtype
84
- x_float = x.float()
85
- x_norm = x_float * torch.rsqrt(
86
- x_float.pow(2).mean(dim=-1, keepdim=True) + self.eps
87
- )
88
- return (x_norm * self.weight.float()).to(dtype)
89
-
90
-
91
- def rotate_half(x: torch.Tensor) -> torch.Tensor:
92
- x1 = x[..., ::2]
93
- x2 = x[..., 1::2]
94
- return torch.stack((-x2, x1), dim=-1).flatten(-2)
95
-
96
-
97
- class RotaryEmbedding(nn.Module):
98
- def __init__(self, dim: int, max_position: int, theta: float):
99
- super().__init__()
100
-
101
- inv_freq = 1.0 / (
102
- theta ** (
103
- torch.arange(0, dim, 2, dtype=torch.float32) / dim
104
- )
105
- )
106
- positions = torch.arange(max_position, dtype=torch.float32)
107
- freqs = torch.outer(positions, inv_freq)
108
-
109
- # Повторяем каждую частоту для пары even/odd.
110
- emb = torch.repeat_interleave(freqs, repeats=2, dim=-1)
111
-
112
- self.register_buffer(
113
- "cos_cached",
114
- emb.cos(),
115
- persistent=False,
116
- )
117
- self.register_buffer(
118
- "sin_cached",
119
- emb.sin(),
120
- persistent=False,
121
- )
122
-
123
- def forward(
124
- self,
125
- q: torch.Tensor,
126
- k: torch.Tensor,
127
- position_ids: Optional[torch.Tensor] = None,
128
- ):
129
- # q, k: [B, H, T, Dh]
130
- T = q.shape[-2]
131
-
132
- if position_ids is None:
133
- cos = self.cos_cached[:T][None, None, :, :]
134
- sin = self.sin_cached[:T][None, None, :, :]
135
- else:
136
- # position_ids: [B, T]
137
- cos = self.cos_cached[position_ids][:, None, :, :]
138
- sin = self.sin_cached[position_ids][:, None, :, :]
139
-
140
- cos = cos.to(device=q.device, dtype=q.dtype)
141
- sin = sin.to(device=q.device, dtype=q.dtype)
142
-
143
- q = q * cos + rotate_half(q) * sin
144
- k = k * cos + rotate_half(k) * sin
145
- return q, k
146
-
147
-
148
- class CausalSelfAttention(nn.Module):
149
- def __init__(self, config: ClassicConfig):
150
- super().__init__()
151
-
152
- self.d_model = config.d_model
153
- self.n_head = config.n_head
154
- self.head_dim = config.d_model // config.n_head
155
- self.dropout_p = config.dropout
156
-
157
- self.q_proj = nn.Linear(
158
- config.d_model,
159
- config.d_model,
160
- bias=config.attention_bias,
161
- )
162
- self.k_proj = nn.Linear(
163
- config.d_model,
164
- config.d_model,
165
- bias=config.attention_bias,
166
- )
167
- self.v_proj = nn.Linear(
168
- config.d_model,
169
- config.d_model,
170
- bias=config.attention_bias,
171
- )
172
- self.o_proj = nn.Linear(
173
- config.d_model,
174
- config.d_model,
175
- bias=config.attention_bias,
176
- )
177
-
178
- self.rope = RotaryEmbedding(
179
- dim=self.head_dim,
180
- max_position=config.block_size,
181
- theta=config.rope_theta,
182
- )
183
-
184
- def forward(
185
- self,
186
- x: torch.Tensor,
187
- attention_mask: Optional[torch.Tensor] = None,
188
- position_ids: Optional[torch.Tensor] = None,
189
- ) -> torch.Tensor:
190
- B, T, C = x.shape
191
-
192
- q = self.q_proj(x).view(
193
- B, T, self.n_head, self.head_dim
194
- ).transpose(1, 2)
195
-
196
- k = self.k_proj(x).view(
197
- B, T, self.n_head, self.head_dim
198
- ).transpose(1, 2)
199
-
200
- v = self.v_proj(x).view(
201
- B, T, self.n_head, self.head_dim
202
- ).transpose(1, 2)
203
-
204
- q, k = self.rope(q, k, position_ids=position_ids)
205
-
206
- dropout_p = self.dropout_p if self.training else 0.0
207
-
208
- if attention_mask is None or bool(attention_mask.all()):
209
- # is_causal=True позволяет PyTorch выбрать Flash SDP kernel.
210
- out = F.scaled_dot_product_attention(
211
- q,
212
- k,
213
- v,
214
- attn_mask=None,
215
- dropout_p=dropout_p,
216
- is_causal=True,
217
- )
218
- else:
219
- if attention_mask.shape != (B, T):
220
- raise ValueError(
221
- f"attention_mask must be {(B, T)}, "
222
- f"got {tuple(attention_mask.shape)}"
223
- )
224
-
225
- causal = torch.ones(
226
- (T, T),
227
- device=x.device,
228
- dtype=torch.bool,
229
- ).tril()
230
-
231
- # В SDPA bool=True означает, что элемент разрешён.
232
- allowed = (
233
- causal[None, None, :, :]
234
- & attention_mask[:, None, None, :].bool()
235
- )
236
-
237
- out = F.scaled_dot_product_attention(
238
- q,
239
- k,
240
- v,
241
- attn_mask=allowed,
242
- dropout_p=dropout_p,
243
- is_causal=False,
244
- )
245
-
246
- out = out.transpose(1, 2).contiguous().view(B, T, C)
247
- return self.o_proj(out)
248
-
249
-
250
- def round_up(value: int, multiple: int) -> int:
251
- return multiple * math.ceil(value / multiple)
252
-
253
-
254
- class SwiGLU(nn.Module):
255
- def __init__(self, config: ClassicConfig):
256
- super().__init__()
257
-
258
- hidden_dim = round_up(
259
- int(config.ffn_multiplier * config.d_model),
260
- config.multiple_of,
261
- )
262
-
263
- self.hidden_dim = hidden_dim
264
-
265
- self.gate_proj = nn.Linear(
266
- config.d_model,
267
- hidden_dim,
268
- bias=config.mlp_bias,
269
- )
270
- self.up_proj = nn.Linear(
271
- config.d_model,
272
- hidden_dim,
273
- bias=config.mlp_bias,
274
- )
275
- self.down_proj = nn.Linear(
276
- hidden_dim,
277
- config.d_model,
278
- bias=config.mlp_bias,
279
- )
280
- self.dropout = nn.Dropout(config.dropout)
281
-
282
- def forward(self, x: torch.Tensor) -> torch.Tensor:
283
- x = F.silu(self.gate_proj(x)) * self.up_proj(x)
284
- return self.dropout(self.down_proj(x))
285
-
286
-
287
- class TransformerBlock(nn.Module):
288
- def __init__(self, config: ClassicConfig):
289
- super().__init__()
290
-
291
- self.input_norm = RMSNorm(
292
- config.d_model,
293
- eps=config.rms_norm_eps,
294
- )
295
- self.post_attention_norm = RMSNorm(
296
- config.d_model,
297
- eps=config.rms_norm_eps,
298
- )
299
-
300
- self.attention = CausalSelfAttention(config)
301
- self.mlp = SwiGLU(config)
302
-
303
- def forward(
304
- self,
305
- x: torch.Tensor,
306
- attention_mask: Optional[torch.Tensor] = None,
307
- position_ids: Optional[torch.Tensor] = None,
308
- ) -> torch.Tensor:
309
- x = x + self.attention(
310
- self.input_norm(x),
311
- attention_mask=attention_mask,
312
- position_ids=position_ids,
313
- )
314
- x = x + self.mlp(self.post_attention_norm(x))
315
- return x
316
-
317
-
318
- class ClassicForCausalLM(PreTrainedModel):
319
- config_class = ClassicConfig
320
- main_input_name = "input_ids"
321
- supports_gradient_checkpointing = True
322
-
323
- def __init__(self, config: ClassicConfig):
324
- super().__init__(config)
325
-
326
- self.token_embeddings = nn.Embedding(
327
- config.vocab_size,
328
- config.d_model,
329
- )
330
-
331
- self.layers = nn.ModuleList(
332
- [TransformerBlock(config) for _ in range(config.n_layer)]
333
- )
334
-
335
- self.final_norm = RMSNorm(
336
- config.d_model,
337
- eps=config.rms_norm_eps,
338
- )
339
-
340
- self.lm_head = nn.Linear(
341
- config.d_model,
342
- config.vocab_size,
343
- bias=False,
344
- )
345
-
346
- self.gradient_checkpointing = False
347
- self.post_init()
348
-
349
- residual_std = (
350
- config.initializer_range
351
- / math.sqrt(2 * config.n_layer)
352
- )
353
-
354
- for layer in self.layers:
355
- nn.init.normal_(
356
- layer.attention.o_proj.weight,
357
- mean=0.0,
358
- std=residual_std,
359
- )
360
- nn.init.normal_(
361
- layer.mlp.down_proj.weight,
362
- mean=0.0,
363
- std=residual_std,
364
- )
365
-
366
- if config.tie_word_embeddings:
367
- self.tie_weights()
368
-
369
- def _init_weights(self, module):
370
- if isinstance(module, nn.Linear):
371
- nn.init.normal_(
372
- module.weight,
373
- mean=0.0,
374
- std=self.config.initializer_range,
375
- )
376
- if module.bias is not None:
377
- nn.init.zeros_(module.bias)
378
-
379
- elif isinstance(module, nn.Embedding):
380
- nn.init.normal_(
381
- module.weight,
382
- mean=0.0,
383
- std=self.config.initializer_range,
384
- )
385
-
386
- def _set_gradient_checkpointing(
387
- self,
388
- module,
389
- value=False,
390
- enable=None,
391
- gradient_checkpointing_func=None,
392
- ):
393
- if enable is not None:
394
- value = enable
395
- self.gradient_checkpointing = value
396
-
397
- def get_input_embeddings(self):
398
- return self.token_embeddings
399
-
400
- def set_input_embeddings(self, value):
401
- self.token_embeddings = value
402
-
403
- def get_output_embeddings(self):
404
- return self.lm_head
405
-
406
- def set_output_embeddings(self, value):
407
- self.lm_head = value
408
-
409
- def count_parameters(self):
410
- total = sum(p.numel() for p in self.parameters())
411
- trainable = sum(
412
- p.numel() for p in self.parameters() if p.requires_grad
413
- )
414
- input_params = self.token_embeddings.weight.numel()
415
- output_params = self.lm_head.weight.numel()
416
-
417
- return {
418
- "total": total,
419
- "trainable": trainable,
420
- "input": input_params,
421
- "output": output_params,
422
- "body": total - input_params - output_params,
423
- }
424
-
425
- def forward(
426
- self,
427
- input_ids: torch.Tensor,
428
- attention_mask: Optional[torch.Tensor] = None,
429
- labels: Optional[torch.Tensor] = None,
430
- position_ids: Optional[torch.Tensor] = None,
431
- return_dict: Optional[bool] = None,
432
- output_logits: bool = True,
433
- **kwargs,
434
- ):
435
- if input_ids is None:
436
- raise ValueError("input_ids must be provided")
437
-
438
- B, T = input_ids.shape
439
-
440
- if T > self.config.block_size:
441
- raise ValueError(
442
- f"Sequence length {T} exceeds "
443
- f"block_size={self.config.block_size}"
444
- )
445
-
446
- if attention_mask is not None:
447
- if attention_mask.shape != (B, T):
448
- raise ValueError(
449
- f"attention_mask must be {(B, T)}, "
450
- f"got {tuple(attention_mask.shape)}"
451
- )
452
-
453
- x = self.token_embeddings(input_ids)
454
-
455
- for layer in self.layers:
456
- if self.gradient_checkpointing and self.training:
457
- def custom_forward(
458
- hidden_states,
459
- current_layer=layer,
460
- ):
461
- return current_layer(
462
- hidden_states,
463
- attention_mask=attention_mask,
464
- position_ids=position_ids,
465
- )
466
-
467
- x = torch.utils.checkpoint.checkpoint(
468
- custom_forward,
469
- x,
470
- use_reentrant=False,
471
- )
472
- else:
473
- x = layer(
474
- x,
475
- attention_mask=attention_mask,
476
- position_ids=position_ids,
477
- )
478
-
479
- x = self.final_norm(x)
480
- logits = self.lm_head(x) if output_logits or labels is not None else None
481
-
482
- loss = None
483
- if labels is not None:
484
- loss_labels = labels.contiguous().clone()
485
-
486
- if attention_mask is not None:
487
- loss_labels.masked_fill_(
488
- attention_mask.eq(0),
489
- -100,
490
- )
491
-
492
- loss = F.cross_entropy(
493
- logits.float().reshape(-1, self.config.vocab_size),
494
- loss_labels.reshape(-1),
495
- ignore_index=-100,
496
- )
497
-
498
- '''shift_logits = logits[:, :-1, :].contiguous()
499
- shift_labels = labels[:, 1:].contiguous().clone()
500
-
501
- if attention_mask is not None:
502
- shift_labels.masked_fill_(
503
- attention_mask[:, 1:].eq(0),
504
- -100,
505
- )
506
-
507
- loss = F.cross_entropy(
508
- shift_logits.float().view(-1, self.config.vocab_size),
509
- shift_labels.view(-1),
510
- ignore_index=-100,
511
- )'''
512
-
513
- if labels is not None and labels.shape != input_ids.shape:
514
- raise ValueError(
515
- f"labels shape {tuple(labels.shape)} must equal "
516
- f"input_ids shape {tuple(input_ids.shape)}"
517
- )
518
-
519
- return_dict = (
520
- self.config.use_return_dict
521
- if return_dict is None
522
- else return_dict
523
- )
524
-
525
- if not return_dict:
526
- output = (logits,) if output_logits else tuple()
527
- return ((loss,) + output) if loss is not None else output
528
-
529
- return CausalLMOutput(
530
- loss=loss,
531
- logits=logits if output_logits else None,
532
- )
533
-
534
- @torch.no_grad()
535
- def generate_simple(
536
- self,
537
- input_ids: torch.Tensor,
538
- max_new_tokens: int,
539
- temperature: float = 0.0,
540
- top_k: Optional[int] = None,
541
- eos_token_id: Optional[int] = None,
542
- ):
543
- was_training = self.training
544
- self.eval()
545
-
546
- for _ in range(max_new_tokens):
547
- x = input_ids[:, -self.config.block_size:]
548
-
549
- logits = self(
550
- input_ids=x,
551
- return_dict=True,
552
- ).logits[:, -1, :]
553
-
554
- if temperature <= 0:
555
- next_token = torch.argmax(
556
- logits,
557
- dim=-1,
558
- keepdim=True,
559
- )
560
- else:
561
- logits = logits / temperature
562
-
563
- if top_k is not None:
564
- values, _ = torch.topk(
565
- logits,
566
- min(top_k, logits.size(-1)),
567
- )
568
- cutoff = values[:, [-1]]
569
- logits = logits.masked_fill(
570
- logits < cutoff,
571
- float("-inf"),
572
- )
573
-
574
- probs = F.softmax(logits.float(), dim=-1)
575
- next_token = torch.multinomial(
576
- probs,
577
- num_samples=1,
578
- )
579
-
580
- input_ids = torch.cat(
581
- [input_ids, next_token],
582
- dim=1,
583
- )
584
-
585
- if (
586
- eos_token_id is not None
587
- and bool((next_token == eos_token_id).all())
588
- ):
589
- break
590
-
591
- if was_training:
592
- self.train()
593
-
594
- return input_ids
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
training_original/classic_train.py DELETED
@@ -1,1779 +0,0 @@
1
- #!/usr/bin/env python3
2
-
3
- import argparse
4
- import contextlib
5
- import glob
6
- import json
7
- import math
8
- import os
9
- import random
10
- import time
11
- import glob
12
- from dataclasses import asdict, dataclass
13
- from datetime import datetime
14
- from pathlib import Path
15
- from typing import Dict, List, Optional, Tuple
16
-
17
- import numpy as np
18
- import torch
19
- import torch.distributed as dist
20
- from safetensors.torch import save_file
21
- from torch.nn.parallel import DistributedDataParallel as DDP
22
- from transformers import AutoTokenizer
23
-
24
- from classic_model import ClassicConfig, ClassicForCausalLM
25
-
26
-
27
- # ------------------------------------------------------------
28
- # Distributed utilities
29
- # ------------------------------------------------------------
30
-
31
- def distributed_is_initialized() -> bool:
32
- return dist.is_available() and dist.is_initialized()
33
-
34
-
35
- def get_rank() -> int:
36
- return dist.get_rank() if distributed_is_initialized() else 0
37
-
38
-
39
- def get_world_size() -> int:
40
- return dist.get_world_size() if distributed_is_initialized() else 1
41
-
42
-
43
- def is_main_process() -> bool:
44
- return get_rank() == 0
45
-
46
-
47
- def setup_distributed():
48
- world_size = int(os.environ.get("WORLD_SIZE", "1"))
49
- distributed = world_size > 1
50
-
51
- if distributed:
52
- local_rank = int(os.environ["LOCAL_RANK"])
53
- torch.cuda.set_device(local_rank)
54
-
55
- dist.init_process_group(
56
- backend="nccl",
57
- init_method="env://",
58
- )
59
-
60
- device = torch.device("cuda", local_rank)
61
- else:
62
- local_rank = 0
63
- device = torch.device(
64
- "cuda" if torch.cuda.is_available() else "cpu"
65
- )
66
-
67
- return distributed, local_rank, device
68
-
69
-
70
- def cleanup_distributed():
71
- if distributed_is_initialized():
72
- dist.barrier()
73
- dist.destroy_process_group()
74
-
75
-
76
- def all_reduce_sum(value: torch.Tensor) -> torch.Tensor:
77
- if distributed_is_initialized():
78
- dist.all_reduce(value, op=dist.ReduceOp.SUM)
79
- return value
80
-
81
-
82
- # ------------------------------------------------------------
83
- # Logging
84
- # ------------------------------------------------------------
85
-
86
- class Logger:
87
- def __init__(self, path: str):
88
- self.path = path
89
-
90
- if is_main_process():
91
- Path(path).parent.mkdir(
92
- parents=True,
93
- exist_ok=True,
94
- )
95
-
96
- def log(self, message: str):
97
- if not is_main_process():
98
- return
99
-
100
- print(message, flush=True)
101
-
102
- with open(self.path, "a", encoding="utf-8") as f:
103
- f.write(message.rstrip() + "\n")
104
-
105
-
106
- def timestamp() -> str:
107
- return datetime.now().strftime("%Y-%m-%d %H:%M")
108
-
109
-
110
- def human_tokens(value: int) -> str:
111
- if value >= 1_000_000_000:
112
- return f"{value / 1_000_000_000:.3f}B"
113
- if value >= 1_000_000:
114
- return f"{value / 1_000_000:.3f}M"
115
- if value >= 1_000:
116
- return f"{value / 1_000:.3f}K"
117
- return str(value)
118
-
119
-
120
- # ------------------------------------------------------------
121
- # Dataset
122
- # ------------------------------------------------------------
123
-
124
- class TokenShardStore:
125
- """
126
- Stateful sequential shard sampler для файлов:
127
-
128
- {split}_u16_*.pt
129
- {split}_doc_offsets_u64_*.pt
130
-
131
- Каждый DDP rank:
132
- 1. получает собственный порядок shards;
133
- 2. загружает один shard;
134
- 3. использует его для batches_per_shard microbatches;
135
- 4. только затем переходит к следующему shard.
136
-
137
- Это предотвращает постоянный torch.load() с NFS и резкий
138
- дисбаланс между DDP ranks.
139
-
140
- Samples не пересекают границы документов.
141
- """
142
-
143
- def __init__(
144
- self,
145
- data_dir: str,
146
- split: str,
147
- sequence_length: int,
148
- cache_shards: int = 2,
149
- seed: int = 42,
150
- batches_per_shard: int = 256,
151
- min_long_documents: int = 1,
152
- ):
153
- if split not in ("train", "valid"):
154
- raise ValueError("split must be 'train' or 'valid'")
155
-
156
- if sequence_length < 1:
157
- raise ValueError("sequence_length must be positive")
158
-
159
- if batches_per_shard < 1:
160
- raise ValueError("batches_per_shard must be positive")
161
-
162
- self.data_dir = data_dir
163
- self.split = split
164
- self.sequence_length = int(sequence_length)
165
- self.required_tokens = self.sequence_length + 1
166
- self.cache_shards = max(1, int(cache_shards))
167
- self.seed = int(seed)
168
- self.batches_per_shard = int(batches_per_shard)
169
- self.min_long_documents = int(min_long_documents)
170
-
171
- # Значения берутся после dist.init_process_group().
172
- self.rank = get_rank()
173
- self.world_size = get_world_size()
174
-
175
- token_pattern = os.path.join(
176
- data_dir,
177
- f"{split}_u16_*.pt",
178
- )
179
-
180
- self.token_paths = sorted(glob.glob(token_pattern))
181
-
182
- if not self.token_paths:
183
- raise FileNotFoundError(
184
- f"No token shards found: {token_pattern}"
185
- )
186
-
187
- self.offset_paths = []
188
-
189
- for token_path in self.token_paths:
190
- token_name = os.path.basename(token_path)
191
- suffix = token_name[len(f"{split}_u16_"):]
192
-
193
- offset_name = (
194
- f"{split}_doc_offsets_u64_{suffix}"
195
- )
196
- offset_path = os.path.join(
197
- data_dir,
198
- offset_name,
199
- )
200
-
201
- if not os.path.exists(offset_path):
202
- raise FileNotFoundError(
203
- f"Missing offsets file: {offset_path}"
204
- )
205
-
206
- self.offset_paths.append(offset_path)
207
-
208
- self.num_shards = len(self.token_paths)
209
-
210
- # Небольшой LRU cache. Обычно активен один shard, предыдущий
211
- # может остаться для дешёвого возврата после reshuffle.
212
- self._cache = {}
213
- self._cache_order = []
214
-
215
- # Stateful cursor.
216
- self._epoch = 0
217
- self._shard_order = []
218
- self._shard_cursor = 0
219
-
220
- self._current_shard_index = None
221
- self._current_tokens = None
222
- self._current_offsets = None
223
- self._current_valid_docs = None
224
- self._batches_on_current_shard = 0
225
-
226
- self._build_shard_order()
227
- self._advance_to_next_usable_shard()
228
-
229
- def __len__(self):
230
- return self.num_shards
231
-
232
- def _build_shard_order(self):
233
- """
234
- Создаёт общий deterministic permutation, затем каждый rank
235
- начинает с собственного смещения.
236
-
237
- При world_size=2 rank 0 и rank 1 обычно читают разные shards.
238
- """
239
- generator = torch.Generator(device="cpu")
240
- generator.manual_seed(
241
- self.seed + self._epoch * 1_000_003
242
- )
243
-
244
- permutation = torch.randperm(
245
- self.num_shards,
246
- generator=generator,
247
- ).tolist()
248
-
249
- if self.num_shards > 1:
250
- shift = self.rank % self.num_shards
251
- permutation = (
252
- permutation[shift:]
253
- + permutation[:shift]
254
- )
255
-
256
- self._shard_order = permutation
257
- self._shard_cursor = 0
258
- self._epoch += 1
259
-
260
- def _touch_cache(self, shard_index: int):
261
- if shard_index in self._cache_order:
262
- self._cache_order.remove(shard_index)
263
-
264
- self._cache_order.append(shard_index)
265
-
266
- while len(self._cache_order) > self.cache_shards:
267
- old_index = self._cache_order.pop(0)
268
-
269
- if old_index == self._current_shard_index:
270
- # Активный shard не удаляем.
271
- self._cache_order.append(old_index)
272
- break
273
-
274
- self._cache.pop(old_index, None)
275
-
276
- def _load_shard(
277
- self,
278
- shard_index: int,
279
- ):
280
- if shard_index in self._cache:
281
- self._touch_cache(shard_index)
282
- return self._cache[shard_index]
283
-
284
- tokens = torch.load(
285
- self.token_paths[shard_index],
286
- map_location="cpu",
287
- weights_only=True,
288
- )
289
-
290
- offsets = torch.load(
291
- self.offset_paths[shard_index],
292
- map_location="cpu",
293
- weights_only=True,
294
- )
295
-
296
- if tokens.ndim != 1:
297
- raise RuntimeError(
298
- f"Token tensor must be 1D: "
299
- f"{self.token_paths[shard_index]}"
300
- )
301
-
302
- if offsets.ndim != 1:
303
- raise RuntimeError(
304
- f"Offsets tensor must be 1D: "
305
- f"{self.offset_paths[shard_index]}"
306
- )
307
-
308
- if tokens.dtype != torch.uint16:
309
- raise TypeError(
310
- f"Expected uint16 tokens, got {tokens.dtype}: "
311
- f"{self.token_paths[shard_index]}"
312
- )
313
-
314
- if offsets.dtype not in (
315
- torch.uint64,
316
- torch.int64,
317
- ):
318
- raise TypeError(
319
- f"Expected uint64/int64 offsets, "
320
- f"got {offsets.dtype}: "
321
- f"{self.offset_paths[shard_index]}"
322
- )
323
-
324
- if offsets.numel() < 2:
325
- raise RuntimeError(
326
- f"Offsets contain no documents: "
327
- f"{self.offset_paths[shard_index]}"
328
- )
329
-
330
- if int(offsets[0]) != 0:
331
- raise RuntimeError(
332
- f"First offset must be zero: "
333
- f"{self.offset_paths[shard_index]}"
334
- )
335
-
336
- if int(offsets[-1]) != tokens.numel():
337
- raise RuntimeError(
338
- f"Final offset/token mismatch in shard "
339
- f"{shard_index}: "
340
- f"{int(offsets[-1])} != {tokens.numel()}"
341
- )
342
-
343
- # Рассчитывается один раз при загрузке shard.
344
- # valid_docs содержит индексы документов, которые достаточно
345
- # длинные для sequence_length + target token.
346
- lengths = (
347
- offsets[1:].to(torch.int64)
348
- - offsets[:-1].to(torch.int64)
349
- )
350
-
351
- valid_docs = torch.nonzero(
352
- lengths >= self.required_tokens,
353
- as_tuple=False,
354
- ).flatten()
355
-
356
- result = (
357
- tokens,
358
- offsets,
359
- valid_docs,
360
- )
361
-
362
- self._cache[shard_index] = result
363
- self._touch_cache(shard_index)
364
-
365
- return result
366
-
367
- def _next_shard_index(self) -> int:
368
- if self._shard_cursor >= len(self._shard_order):
369
- self._build_shard_order()
370
-
371
- shard_index = self._shard_order[
372
- self._shard_cursor
373
- ]
374
- self._shard_cursor += 1
375
-
376
- return int(shard_index)
377
-
378
- def _advance_to_next_usable_shard(self):
379
- """
380
- Переключается на следующий shard с достаточным количеством
381
- длинных документов.
382
-
383
- Ограничение attempts предотвращает бесконечный цикл, если
384
- sequence_length слишком велик для всех документов.
385
- """
386
- attempts = 0
387
- max_attempts = max(
388
- self.num_shards * 2,
389
- 1,
390
- )
391
-
392
- while attempts < max_attempts:
393
- attempts += 1
394
- shard_index = self._next_shard_index()
395
-
396
- tokens, offsets, valid_docs = self._load_shard(
397
- shard_index
398
- )
399
-
400
- if valid_docs.numel() < self.min_long_documents:
401
- continue
402
-
403
- self._current_shard_index = shard_index
404
- self._current_tokens = tokens
405
- self._current_offsets = offsets
406
- self._current_valid_docs = valid_docs
407
- self._batches_on_current_shard = 0
408
- return
409
-
410
- raise RuntimeError(
411
- "Could not find a usable shard containing documents "
412
- f"with at least {self.required_tokens} tokens. "
413
- "Reduce sequence_length or inspect document offsets."
414
- )
415
-
416
- def _ensure_current_shard(self):
417
- if self._current_tokens is None:
418
- self._advance_to_next_usable_shard()
419
- return
420
-
421
- if (
422
- self._batches_on_current_shard
423
- >= self.batches_per_shard
424
- ):
425
- self._advance_to_next_usable_shard()
426
-
427
- def sample_batch(
428
- self,
429
- batch_size: int,
430
- generator: torch.Generator,
431
- ) -> torch.Tensor:
432
- """
433
- Возвращает LongTensor:
434
-
435
- [batch_size, sequence_length + 1]
436
-
437
- Все samples данного microbatch берутся из одного resident shard,
438
- но из случайных документов и случайных позиций внутри документов.
439
- """
440
- if batch_size < 1:
441
- raise ValueError("batch_size must be positive")
442
-
443
- self._ensure_current_shard()
444
-
445
- tokens = self._current_tokens
446
- offsets = self._current_offsets
447
- valid_docs = self._current_valid_docs
448
-
449
- num_valid_docs = valid_docs.numel()
450
-
451
- if num_valid_docs == 0:
452
- raise RuntimeError(
453
- "Internal error: current shard has no valid docs"
454
- )
455
-
456
- # Сразу выбираем batch_size документов.
457
- selected_positions = torch.randint(
458
- low=0,
459
- high=num_valid_docs,
460
- size=(batch_size,),
461
- generator=generator,
462
- )
463
-
464
- selected_docs = valid_docs[selected_positions]
465
- samples = []
466
-
467
- for doc_tensor in selected_docs:
468
- doc_index = int(doc_tensor)
469
-
470
- begin = int(offsets[doc_index])
471
- end = int(offsets[doc_index + 1])
472
- document_length = end - begin
473
-
474
- max_local_start = (
475
- document_length - self.required_tokens
476
- )
477
-
478
- if max_local_start > 0:
479
- local_start = int(
480
- torch.randint(
481
- low=0,
482
- high=max_local_start + 1,
483
- size=(1,),
484
- generator=generator,
485
- ).item()
486
- )
487
- else:
488
- local_start = 0
489
-
490
- start = begin + local_start
491
- stop = start + self.required_tokens
492
-
493
- sample = tokens[start:stop]
494
-
495
- if sample.numel() != self.required_tokens:
496
- raise RuntimeError(
497
- "Internal slicing error: "
498
- f"expected {self.required_tokens}, "
499
- f"got {sample.numel()}"
500
- )
501
-
502
- # uint16 → int64 для embedding lookup.
503
- samples.append(sample.to(torch.long))
504
-
505
- self._batches_on_current_shard += 1
506
-
507
- return torch.stack(
508
- samples,
509
- dim=0,
510
- )
511
-
512
- def state_dict(self):
513
- """
514
- Опциональное состояние sampler для точного resume.
515
- Сейчас trainer его не сохраняет, но интерфейс оставлен.
516
- """
517
- return {
518
- "epoch": self._epoch,
519
- "shard_order": list(self._shard_order),
520
- "shard_cursor": self._shard_cursor,
521
- "current_shard_index":
522
- self._current_shard_index,
523
- "batches_on_current_shard":
524
- self._batches_on_current_shard,
525
- }
526
-
527
- def load_state_dict(self, state):
528
- """
529
- Восстанавливает shard cursor. RNG generator trainer сохраняется
530
- отдельно только если вы добавите его state в checkpoint.
531
- """
532
- self._epoch = int(state["epoch"])
533
- self._shard_order = list(state["shard_order"])
534
- self._shard_cursor = int(state["shard_cursor"])
535
-
536
- current_shard_index = state.get(
537
- "current_shard_index"
538
- )
539
-
540
- if current_shard_index is None:
541
- self._current_shard_index = None
542
- self._current_tokens = None
543
- self._current_offsets = None
544
- self._current_valid_docs = None
545
- self._batches_on_current_shard = 0
546
- self._advance_to_next_usable_shard()
547
- return
548
-
549
- tokens, offsets, valid_docs = self._load_shard(
550
- int(current_shard_index)
551
- )
552
-
553
- if valid_docs.numel() == 0:
554
- raise RuntimeError(
555
- "Saved current shard is no longer usable"
556
- )
557
-
558
- self._current_shard_index = int(
559
- current_shard_index
560
- )
561
- self._current_tokens = tokens
562
- self._current_offsets = offsets
563
- self._current_valid_docs = valid_docs
564
- self._batches_on_current_shard = int(
565
- state.get(
566
- "batches_on_current_shard",
567
- 0,
568
- )
569
- )
570
-
571
-
572
- # ------------------------------------------------------------
573
- # Schedules and optimizer
574
- # ------------------------------------------------------------
575
-
576
- def get_learning_rate(
577
- step: int,
578
- max_steps: int,
579
- warmup_steps: int,
580
- learning_rate: float,
581
- min_learning_rate: float,
582
- ) -> float:
583
- if step < warmup_steps:
584
- return learning_rate * float(step + 1) / max(1, warmup_steps)
585
-
586
- if step >= max_steps:
587
- return min_learning_rate
588
-
589
- decay_ratio = (
590
- step - warmup_steps
591
- ) / max(1, max_steps - warmup_steps)
592
-
593
- coefficient = 0.5 * (
594
- 1.0 + math.cos(math.pi * decay_ratio)
595
- )
596
-
597
- return (
598
- min_learning_rate
599
- + coefficient
600
- * (learning_rate - min_learning_rate)
601
- )
602
-
603
-
604
- def configure_optimizer(
605
- model: torch.nn.Module,
606
- learning_rate: float,
607
- weight_decay: float,
608
- betas: Tuple[float, float],
609
- fused: bool,
610
- ):
611
- decay_params = []
612
- no_decay_params = []
613
-
614
- for name, parameter in model.named_parameters():
615
- if not parameter.requires_grad:
616
- continue
617
-
618
- if parameter.dim() >= 2:
619
- decay_params.append(parameter)
620
- else:
621
- no_decay_params.append(parameter)
622
-
623
- groups = [
624
- {
625
- "params": decay_params,
626
- "weight_decay": weight_decay,
627
- },
628
- {
629
- "params": no_decay_params,
630
- "weight_decay": 0.0,
631
- },
632
- ]
633
-
634
- kwargs = dict(
635
- lr=learning_rate,
636
- betas=betas,
637
- eps=1e-8,
638
- )
639
-
640
- if fused and torch.cuda.is_available():
641
- kwargs["fused"] = True
642
-
643
- return torch.optim.AdamW(groups, **kwargs)
644
-
645
-
646
- # ------------------------------------------------------------
647
- # Evaluation and generation
648
- # ------------------------------------------------------------
649
-
650
- @torch.no_grad()
651
- def estimate_loss(
652
- model,
653
- dataset: TokenShardStore,
654
- device: torch.device,
655
- batch_size: int,
656
- eval_batches: int,
657
- seed: int,
658
- ) -> Tuple[float, int]:
659
- model.eval()
660
-
661
- generator = torch.Generator(device="cpu")
662
- generator.manual_seed(seed + get_rank())
663
-
664
- loss_sum = torch.zeros(
665
- 1,
666
- device=device,
667
- dtype=torch.float64,
668
- )
669
- token_count = torch.zeros(
670
- 1,
671
- device=device,
672
- dtype=torch.float64,
673
- )
674
-
675
- for _ in range(eval_batches):
676
- batch = dataset.sample_batch(
677
- batch_size=batch_size,
678
- generator=generator,
679
- )
680
-
681
- batch = batch.to(
682
- device,
683
- non_blocking=True,
684
- )
685
-
686
- #inputs = batch[:, :-1]
687
- #labels = inputs
688
-
689
- inputs = batch[:, :-1]
690
- labels = batch[:, 1:]
691
-
692
- with torch.autocast(
693
- device_type="cuda",
694
- dtype=torch.bfloat16,
695
- enabled=device.type == "cuda",
696
- ):
697
- outputs = model(
698
- input_ids=inputs,
699
- labels=labels,
700
- return_dict=True,
701
- )
702
-
703
- #count = labels[:, 1:].numel()
704
-
705
- count = labels.numel()
706
-
707
- loss_sum += outputs.loss.double() * count
708
- token_count += count
709
-
710
- all_reduce_sum(loss_sum)
711
- all_reduce_sum(token_count)
712
-
713
- mean_loss = float((loss_sum / token_count).item())
714
- total_tokens = int(token_count.item())
715
-
716
- model.train()
717
- return mean_loss, total_tokens
718
-
719
-
720
- @torch.no_grad()
721
- def generate_diagnostics(
722
- raw_model,
723
- tokenizer,
724
- prompts,
725
- device,
726
- max_new_tokens,
727
- temperature=0.0,
728
- top_k=None,
729
- ):
730
- raw_model.eval()
731
- outputs = []
732
-
733
- for prompt in prompts:
734
- encoded = tokenizer(
735
- prompt,
736
- return_tensors="pt",
737
- add_special_tokens=False,
738
- )
739
-
740
- input_ids = encoded["input_ids"].to(device)
741
-
742
- generated = raw_model.generate_simple(
743
- input_ids=input_ids,
744
- max_new_tokens=max_new_tokens,
745
- temperature=temperature,
746
- top_k=top_k,
747
- eos_token_id=tokenizer.eos_token_id,
748
- )
749
-
750
- text = tokenizer.decode(
751
- generated[0],
752
- skip_special_tokens=True,
753
- )
754
-
755
- outputs.append(text)
756
-
757
- raw_model.train()
758
- return outputs
759
-
760
-
761
- @torch.no_grad()
762
- def log_input_representation_diagnostics(
763
- raw_model,
764
- tokenizer,
765
- logger,
766
- device: torch.device,
767
- step: int,
768
- texts=None,
769
- vector_prefix: int = 16,
770
- ):
771
- if not is_main_process():
772
- return
773
-
774
- if texts is None:
775
- texts = [
776
- "A",
777
- " A",
778
- "А", # Кириллическая A
779
- " А",
780
- "0",
781
- " 0",
782
- "1",
783
- " 1",
784
- "the",
785
- " the",
786
- ]
787
-
788
- embedding = raw_model.get_input_embeddings()
789
-
790
- embedding_parameters = list(embedding.parameters())
791
- trainable_input_parameters = sum(
792
- p.numel()
793
- for p in embedding_parameters
794
- if p.requires_grad
795
- )
796
-
797
- logger.log(
798
- f"step {step:07d}: [input_repr] "
799
- f"module={embedding.__class__.__name__}, "
800
- f"trainable_input_parameters={trainable_input_parameters:,}"
801
- )
802
-
803
- for text in texts:
804
- token_ids = tokenizer.encode(
805
- text,
806
- add_special_tokens=False,
807
- )
808
-
809
- if not token_ids:
810
- logger.log(
811
- f"step {step:07d}: [input_repr] "
812
- f"text={text!r}, token_ids=[]"
813
- )
814
- continue
815
-
816
- ids = torch.tensor(
817
- [token_ids],
818
- dtype=torch.long,
819
- device=device,
820
- )
821
-
822
- vectors = embedding(ids)[0].detach().float().cpu()
823
-
824
- pieces = tokenizer.convert_ids_to_tokens(token_ids)
825
-
826
- for local_index, token_id in enumerate(token_ids):
827
- vector = vectors[local_index]
828
- prefix = vector[:vector_prefix].tolist()
829
-
830
- logger.log(
831
- f"step {step:07d}: [input_repr] "
832
- f"text={text!r}, "
833
- f"piece_index={local_index}, "
834
- f"token_id={token_id}, "
835
- f"token_piece={pieces[local_index]!r}, "
836
- f"first_{vector_prefix}="
837
- f"{[round(x, 6) for x in prefix]}, "
838
- f"mean={vector.mean().item():.6f}, "
839
- f"std={vector.std(unbiased=False).item():.6f}, "
840
- f"norm={vector.norm().item():.6f}"
841
- )
842
-
843
-
844
- # ------------------------------------------------------------
845
- # Checkpoints
846
- # ------------------------------------------------------------
847
-
848
- def save_model_safetensors(path, raw_model):
849
- if not is_main_process():
850
- return
851
-
852
- state = {
853
- name: tensor.detach().cpu().contiguous()
854
- for name, tensor in raw_model.state_dict().items()
855
- }
856
-
857
- tmp_path = path + ".tmp"
858
- save_file(state, tmp_path)
859
- os.replace(tmp_path, path)
860
-
861
-
862
- def save_checkpoint(
863
- path: str,
864
- raw_model: ClassicForCausalLM,
865
- optimizer,
866
- step: int,
867
- tokens_seen: int,
868
- args,
869
- ):
870
- if not is_main_process():
871
- return
872
-
873
- Path(path).parent.mkdir(parents=True, exist_ok=True)
874
-
875
- tmp_path = path + ".tmp"
876
-
877
- checkpoint = {
878
- "step": step,
879
- "tokens_seen": tokens_seen,
880
- "model": raw_model.state_dict(),
881
- "optimizer": optimizer.state_dict(),
882
- "config": raw_model.config.to_dict(),
883
- "args": vars(args),
884
- "torch_rng_state": torch.get_rng_state(),
885
- "cuda_rng_state": (
886
- torch.cuda.get_rng_state_all()
887
- if torch.cuda.is_available()
888
- else None
889
- ),
890
- }
891
-
892
- torch.save(checkpoint, tmp_path)
893
- os.replace(tmp_path, path)
894
-
895
-
896
- def load_checkpoint(
897
- path: str,
898
- raw_model: ClassicForCausalLM,
899
- optimizer,
900
- device: torch.device,
901
- ):
902
- checkpoint = torch.load(
903
- path,
904
- map_location=device,
905
- weights_only=False,
906
- )
907
-
908
- raw_model.load_state_dict(checkpoint["model"])
909
- optimizer.load_state_dict(checkpoint["optimizer"])
910
-
911
- if "torch_rng_state" in checkpoint:
912
- torch.set_rng_state(checkpoint["torch_rng_state"].cpu())
913
-
914
- if (
915
- torch.cuda.is_available()
916
- and checkpoint.get("cuda_rng_state") is not None
917
- ):
918
- torch.cuda.set_rng_state_all(
919
- checkpoint["cuda_rng_state"]
920
- )
921
-
922
- return (
923
- int(checkpoint.get("step", 0)),
924
- int(checkpoint.get("tokens_seen", 0)),
925
- )
926
-
927
-
928
- # ------------------------------------------------------------
929
- # CLI
930
- # ------------------------------------------------------------
931
-
932
- def parse_args():
933
- parser = argparse.ArgumentParser()
934
-
935
- # Paths
936
- parser.add_argument("--data_dir", required=True)
937
- parser.add_argument("--output_dir", required=True)
938
- parser.add_argument(
939
- "--tokenizer",
940
- default="HuggingFaceTB/SmolLM2-135M",
941
- )
942
- parser.add_argument("--tokenizer_revision", default=None)
943
- parser.add_argument("--resume", default=None)
944
-
945
- # Architecture
946
- parser.add_argument("--d_model", type=int, default=960)
947
- parser.add_argument("--n_layer", type=int, default=24)
948
- parser.add_argument("--n_head", type=int, default=15)
949
- parser.add_argument(
950
- "--ffn_multiplier",
951
- type=float,
952
- default=8.0 / 3.0,
953
- )
954
- parser.add_argument("--multiple_of", type=int, default=256)
955
- parser.add_argument("--sequence_length", type=int, default=2048)
956
- parser.add_argument("--rope_theta", type=float, default=10000.0)
957
- parser.add_argument("--dropout", type=float, default=0.0)
958
- parser.add_argument("--rms_norm_eps", type=float, default=1e-5)
959
- parser.add_argument(
960
- "--tie_word_embeddings",
961
- action="store_true",
962
- )
963
- parser.add_argument(
964
- "--gradient_checkpointing",
965
- action="store_true",
966
- )
967
- parser.add_argument(
968
- "--compile",
969
- action="store_true",
970
- )
971
-
972
- # Training
973
- parser.add_argument(
974
- "--micro_batch_size",
975
- type=int,
976
- default=2,
977
- )
978
- parser.add_argument(
979
- "--gradient_accumulation_steps",
980
- type=int,
981
- default=16,
982
- )
983
- parser.add_argument("--max_steps", type=int, default=200_000)
984
- parser.add_argument(
985
- "--max_tokens",
986
- type=int,
987
- default=0,
988
- help="0 disables token-based stopping",
989
- )
990
- parser.add_argument(
991
- "--learning_rate",
992
- type=float,
993
- default=3e-4,
994
- )
995
- parser.add_argument(
996
- "--min_learning_rate",
997
- type=float,
998
- default=3e-5,
999
- )
1000
- parser.add_argument("--warmup_steps", type=int, default=2000)
1001
- parser.add_argument("--weight_decay", type=float, default=0.1)
1002
- parser.add_argument("--beta1", type=float, default=0.9)
1003
- parser.add_argument("--beta2", type=float, default=0.95)
1004
- parser.add_argument("--grad_clip", type=float, default=1.0)
1005
- parser.add_argument("--seed", type=int, default=42)
1006
-
1007
- # Data memory
1008
- parser.add_argument(
1009
- "--train_cache_shards",
1010
- type=int,
1011
- default=2,
1012
- )
1013
- parser.add_argument(
1014
- "--valid_cache_shards",
1015
- type=int,
1016
- default=2,
1017
- )
1018
-
1019
- # Diagnostics
1020
- parser.add_argument(
1021
- "--diagnostic_interval_seconds",
1022
- type=float,
1023
- default=3600.0,
1024
- )
1025
- parser.add_argument(
1026
- "--log_interval_steps",
1027
- type=int,
1028
- default=20,
1029
- )
1030
- parser.add_argument(
1031
- "--eval_batches",
1032
- type=int,
1033
- default=32,
1034
- )
1035
- parser.add_argument(
1036
- "--eval_batch_size",
1037
- type=int,
1038
- default=2,
1039
- )
1040
- parser.add_argument(
1041
- "--generation_tokens",
1042
- type=int,
1043
- default=24,
1044
- )
1045
- parser.add_argument(
1046
- "--save_interval_steps",
1047
- type=int,
1048
- default=2000,
1049
- )
1050
- parser.add_argument(
1051
- "--batches_per_shard",
1052
- type=int,
1053
- default=256,
1054
- help=(
1055
- "Number of microbatches sampled from one resident shard "
1056
- "before loading the next shard."
1057
- ),
1058
- )
1059
- parser.add_argument(
1060
- "--latest_save_interval_steps",
1061
- type=int,
1062
- default=2000,
1063
- )
1064
-
1065
- parser.add_argument(
1066
- "--milestone_save_interval_steps",
1067
- type=int,
1068
- default=20000,
1069
- )
1070
-
1071
- return parser.parse_args()
1072
-
1073
-
1074
- def find_nonfinite_gradients(model, max_names=20):
1075
- bad = []
1076
-
1077
- for name, parameter in model.named_parameters():
1078
- gradient = parameter.grad
1079
-
1080
- if gradient is None:
1081
- continue
1082
-
1083
- finite = torch.isfinite(gradient)
1084
-
1085
- if not bool(finite.all()):
1086
- bad.append(
1087
- {
1088
- "name": name,
1089
- "shape": tuple(gradient.shape),
1090
- "dtype": str(gradient.dtype),
1091
- "nan": int(torch.isnan(gradient).sum().item()),
1092
- "inf": int(torch.isinf(gradient).sum().item()),
1093
- }
1094
- )
1095
-
1096
- if len(bad) >= max_names:
1097
- break
1098
-
1099
- return bad
1100
-
1101
-
1102
- def grad_norm_for_named_parameters(named_parameters):
1103
- squares = []
1104
-
1105
- for _, parameter in named_parameters:
1106
- if parameter.grad is None:
1107
- continue
1108
-
1109
- gradient = parameter.grad.detach().float()
1110
- squares.append(gradient.pow(2).sum())
1111
-
1112
- if not squares:
1113
- return 0.0
1114
-
1115
- return torch.sqrt(torch.stack(squares).sum()).item()
1116
-
1117
-
1118
- # ------------------------------------------------------------
1119
- # Main
1120
- # ------------------------------------------------------------
1121
-
1122
- def main():
1123
- args = parse_args()
1124
-
1125
- distributed, local_rank, device = setup_distributed()
1126
- rank = get_rank()
1127
- world_size = get_world_size()
1128
-
1129
- if device.type != "cuda":
1130
- raise RuntimeError(
1131
- "This training script is intended for CUDA GPUs."
1132
- )
1133
-
1134
- torch.backends.cuda.matmul.allow_tf32 = True
1135
- torch.backends.cudnn.allow_tf32 = True
1136
-
1137
- # Prefer Flash SDP where available.
1138
- torch.backends.cuda.enable_flash_sdp(True)
1139
- torch.backends.cuda.enable_mem_efficient_sdp(True)
1140
- torch.backends.cuda.enable_math_sdp(True)
1141
-
1142
- seed = args.seed + rank
1143
- random.seed(seed)
1144
- np.random.seed(seed)
1145
- torch.manual_seed(seed)
1146
- torch.cuda.manual_seed_all(seed)
1147
-
1148
- Path(args.output_dir).mkdir(
1149
- parents=True,
1150
- exist_ok=True,
1151
- )
1152
-
1153
- logger = Logger(
1154
- os.path.join(args.output_dir, "train.log")
1155
- )
1156
-
1157
- tokenizer = AutoTokenizer.from_pretrained(
1158
- args.tokenizer,
1159
- revision=args.tokenizer_revision,
1160
- use_fast=True,
1161
- )
1162
-
1163
- vocab_size = len(tokenizer)
1164
-
1165
- # Не используем произвольный legacy pad ID.
1166
- # Dataset состоит из fixed-length samples и padding не нужен.
1167
- config = ClassicConfig(
1168
- vocab_size=vocab_size,
1169
- d_model=args.d_model,
1170
- n_layer=args.n_layer,
1171
- n_head=args.n_head,
1172
- ffn_multiplier=args.ffn_multiplier,
1173
- multiple_of=args.multiple_of,
1174
- block_size=args.sequence_length,
1175
- rope_theta=args.rope_theta,
1176
- dropout=args.dropout,
1177
- rms_norm_eps=args.rms_norm_eps,
1178
- pad_token_id=tokenizer.pad_token_id,
1179
- bos_token_id=tokenizer.bos_token_id,
1180
- eos_token_id=tokenizer.eos_token_id,
1181
- tie_word_embeddings=args.tie_word_embeddings,
1182
- )
1183
-
1184
- raw_model = ClassicForCausalLM(config)
1185
-
1186
- if args.gradient_checkpointing:
1187
- raw_model.gradient_checkpointing = True
1188
-
1189
- raw_model.to(device)
1190
-
1191
- if is_main_process():
1192
- first_parameter = next(raw_model.parameters())
1193
-
1194
- logger.log(
1195
- "[dtype] "
1196
- f"parameter_dtype={first_parameter.dtype}, "
1197
- f"optimizer_master_expected=float32"
1198
- )
1199
-
1200
- if next(raw_model.parameters()).dtype != torch.float32:
1201
- raise RuntimeError(
1202
- "Model parameters must remain FP32; "
1203
- "BF16 should be enabled only through autocast"
1204
- )
1205
-
1206
- parameter_counts = raw_model.count_parameters()
1207
-
1208
- if is_main_process():
1209
- logger.log(
1210
- "[model] "
1211
- + ", ".join(
1212
- f"{key}={value:,}"
1213
- for key, value in parameter_counts.items()
1214
- )
1215
- )
1216
- logger.log(
1217
- "[config] "
1218
- + json.dumps(
1219
- config.to_dict(),
1220
- ensure_ascii=False,
1221
- sort_keys=True,
1222
- )
1223
- )
1224
- logger.log(
1225
- f"[run] world_size={world_size}, "
1226
- f"micro_batch={args.micro_batch_size}, "
1227
- f"grad_accum={args.gradient_accumulation_steps}, "
1228
- f"seq={args.sequence_length}, "
1229
- f"global_tokens_per_step="
1230
- #f"{world_size * args.micro_batch_size * args.gradient_accumulation_steps * (args.sequence_length - 1):,}"
1231
- f"{world_size * args.micro_batch_size * args.gradient_accumulation_steps * args.sequence_length:,}"
1232
- )
1233
-
1234
- optimizer = configure_optimizer(
1235
- raw_model,
1236
- learning_rate=args.learning_rate,
1237
- weight_decay=args.weight_decay,
1238
- betas=(args.beta1, args.beta2),
1239
- fused=True,
1240
- )
1241
-
1242
- start_step = 0
1243
- tokens_seen = 0
1244
-
1245
- if args.resume is not None:
1246
- start_step, tokens_seen = load_checkpoint(
1247
- args.resume,
1248
- raw_model,
1249
- optimizer,
1250
- device,
1251
- )
1252
-
1253
- logger.log(
1254
- f"[resume] path={args.resume}, "
1255
- f"step={start_step}, "
1256
- f"tokens_seen={tokens_seen:,}"
1257
- )
1258
-
1259
- model = raw_model
1260
-
1261
- if args.compile:
1262
- model = torch.compile(
1263
- model,
1264
- mode="max-autotune",
1265
- dynamic=False,
1266
- )
1267
-
1268
- if distributed:
1269
- '''model = DDP(
1270
- model,
1271
- device_ids=[local_rank],
1272
- output_device=local_rank,
1273
- broadcast_buffers=False,
1274
- gradient_as_bucket_view=True,
1275
- static_graph=not args.gradient_checkpointing,
1276
- )'''
1277
- model = DDP(
1278
- model,
1279
- device_ids=[local_rank],
1280
- output_device=local_rank,
1281
- broadcast_buffers=False,
1282
- gradient_as_bucket_view=True,
1283
- static_graph=False,
1284
- find_unused_parameters=False,
1285
- )
1286
-
1287
- train_data = TokenShardStore(
1288
- data_dir=args.data_dir,
1289
- split="train",
1290
- sequence_length=args.sequence_length,
1291
- cache_shards=args.train_cache_shards,
1292
- seed=args.seed,
1293
- batches_per_shard=args.batches_per_shard,
1294
- )
1295
-
1296
- valid_data = TokenShardStore(
1297
- data_dir=args.data_dir,
1298
- split="valid",
1299
- sequence_length=args.sequence_length,
1300
- cache_shards=args.valid_cache_shards,
1301
- seed=args.seed + 10_000,
1302
- batches_per_shard=max(
1303
- args.batches_per_shard,
1304
- args.eval_batches,
1305
- ),
1306
- )
1307
-
1308
- train_generator = torch.Generator(device="cpu")
1309
- train_generator.manual_seed(args.seed + rank * 100_003)
1310
-
1311
- prompts = [
1312
- # English factual completion
1313
- "London is the capital of",
1314
- "The capital of France is",
1315
- "The largest planet in the Solar System is",
1316
- "Water freezes at",
1317
- "The chemical symbol for gold is",
1318
- "The Pacific Ocean is",
1319
- "The human heart pumps",
1320
- "The Second World War ended in",
1321
- "The author of Romeo and Juliet was",
1322
- "A triangle has",
1323
-
1324
- # English continuation and grammar
1325
- "Once upon a time, there was",
1326
- "The scientist opened the laboratory door and",
1327
- "When the rain finally stopped,",
1328
- "She went to the store because",
1329
- "If I had known about the problem,",
1330
- "The old house on the hill",
1331
- "Although the experiment failed,",
1332
- "In order to solve this problem, we need to",
1333
- "The main difference between cats and dogs is",
1334
- "This article explains how to",
1335
-
1336
- # Definitions and explanations
1337
- "Photosynthesis is the process by which",
1338
- "Gravity is a force that",
1339
- "A computer program is",
1340
- "Democracy can be defined as",
1341
- "Machine learning is used to",
1342
- "The purpose of a database is to",
1343
- "An ecosystem consists of",
1344
- "Inflation occurs when",
1345
- "The Internet allows people to",
1346
- "Energy cannot be created or destroyed, but",
1347
-
1348
- # Arithmetic and symbolic patterns
1349
- "2 + 2 =",
1350
- "10 - 3 =",
1351
- "6 * 7 =",
1352
- "12 / 4 =",
1353
- "1, 2, 3, 4,",
1354
- "2, 4, 6, 8,",
1355
- "The square root of 9 is",
1356
- "If x = 5, then x + 2 =",
1357
- "One hundred divided by ten equals",
1358
- "The next number after 99 is",
1359
- ]
1360
-
1361
- # Running training loss since the previous diagnostic.
1362
- interval_loss_sum = 0.0
1363
- interval_loss_tokens = 0
1364
- interval_start_tokens = tokens_seen
1365
- interval_start_time = time.monotonic()
1366
-
1367
- last_diagnostic_time = time.monotonic()
1368
-
1369
- # Initial diagnostic.
1370
- if is_main_process():
1371
- logger.log(f"[start] {timestamp()}")
1372
-
1373
- log_input_representation_diagnostics(
1374
- raw_model=raw_model,
1375
- tokenizer=tokenizer,
1376
- logger=logger,
1377
- device=device,
1378
- step=start_step,
1379
- )
1380
-
1381
- model.train()
1382
- optimizer.zero_grad(set_to_none=True)
1383
-
1384
- for step in range(start_step, args.max_steps):
1385
- lr = get_learning_rate(
1386
- step=step,
1387
- max_steps=args.max_steps,
1388
- warmup_steps=args.warmup_steps,
1389
- learning_rate=args.learning_rate,
1390
- min_learning_rate=args.min_learning_rate,
1391
- )
1392
-
1393
- for group in optimizer.param_groups:
1394
- group["lr"] = lr
1395
-
1396
- step_loss_sum = 0.0
1397
- step_loss_tokens = 0
1398
-
1399
- for micro_step in range(
1400
- args.gradient_accumulation_steps
1401
- ):
1402
- batch = train_data.sample_batch(
1403
- batch_size=args.micro_batch_size,
1404
- generator=train_generator,
1405
- )
1406
-
1407
- batch = batch.to(
1408
- device,
1409
- non_blocking=True,
1410
- )
1411
-
1412
- # forward() сам сдвигает logits и labels на один токен.
1413
- #input_ids = batch[:, :-1]
1414
- #labels = input_ids
1415
-
1416
- input_ids = batch[:, :-1]
1417
- labels = batch[:, 1:]
1418
-
1419
- should_sync = (
1420
- micro_step
1421
- == args.gradient_accumulation_steps - 1
1422
- )
1423
-
1424
- sync_context = contextlib.nullcontext()
1425
-
1426
- if distributed and not should_sync:
1427
- sync_context = model.no_sync()
1428
-
1429
- with sync_context:
1430
- with torch.autocast(
1431
- device_type="cuda",
1432
- dtype=torch.bfloat16,
1433
- ):
1434
- outputs = model(
1435
- input_ids=input_ids,
1436
- labels=labels,
1437
- return_dict=True,
1438
- )
1439
-
1440
- loss = (
1441
- outputs.loss
1442
- / args.gradient_accumulation_steps
1443
- )
1444
-
1445
- loss.backward()
1446
-
1447
- # Первый label не имеет предшествующего logit после внутреннего shift.
1448
-
1449
- #local_tokens = labels[:, 1:].numel()
1450
-
1451
- local_tokens = labels.numel()
1452
-
1453
- step_loss_sum += (
1454
- float(outputs.loss.detach()) * local_tokens
1455
- )
1456
- step_loss_tokens += local_tokens
1457
-
1458
-
1459
- bad_gradients = find_nonfinite_gradients(raw_model)
1460
-
1461
- local_bad = torch.tensor(
1462
- [1 if bad_gradients else 0],
1463
- device=device,
1464
- dtype=torch.int32,
1465
- )
1466
-
1467
- if distributed:
1468
- dist.all_reduce(
1469
- local_bad,
1470
- op=dist.ReduceOp.MAX,
1471
- )
1472
-
1473
- if int(local_bad.item()) != 0:
1474
- if bad_gradients:
1475
- logger.log(
1476
- "[nonfinite_gradients] "
1477
- + json.dumps(
1478
- bad_gradients,
1479
- ensure_ascii=False,
1480
- )
1481
- )
1482
-
1483
- emergency_path = os.path.join(
1484
- args.output_dir,
1485
- f"checkpoint_nonfinite_step_{step + 1:07d}.pt",
1486
- )
1487
-
1488
- save_checkpoint(
1489
- path=emergency_path,
1490
- raw_model=raw_model,
1491
- optimizer=optimizer,
1492
- step=step,
1493
- tokens_seen=tokens_seen,
1494
- args=args,
1495
- )
1496
-
1497
- raise RuntimeError(
1498
- f"Non-finite gradients at step {step + 1}"
1499
- )
1500
-
1501
- input_norm = grad_norm_for_named_parameters(
1502
- raw_model.token_embeddings.named_parameters()
1503
- )
1504
-
1505
- body_norm = grad_norm_for_named_parameters(
1506
- (
1507
- (name, parameter)
1508
- for name, parameter in raw_model.named_parameters()
1509
- if not name.startswith("token_embeddings.")
1510
- and not name.startswith("lm_head.")
1511
- )
1512
- )
1513
-
1514
- output_norm = grad_norm_for_named_parameters(
1515
- raw_model.lm_head.named_parameters()
1516
- )
1517
-
1518
- #logger.log(
1519
- # f"[grad_groups] input={input_norm:.3f}, "
1520
- # f"body={body_norm:.3f}, output={output_norm:.3f}"
1521
- #)
1522
-
1523
- if args.grad_clip > 0:
1524
- grad_norm = torch.nn.utils.clip_grad_norm_(
1525
- raw_model.parameters(),
1526
- args.grad_clip,
1527
- error_if_nonfinite=True,
1528
- )
1529
- else:
1530
- grad_norm = torch.tensor(
1531
- float("nan"),
1532
- device=device,
1533
- )
1534
-
1535
- optimizer.step()
1536
- optimizer.zero_grad(set_to_none=True)
1537
-
1538
- # Global token count.
1539
- global_step_tokens = (
1540
- step_loss_tokens * world_size
1541
- )
1542
- tokens_seen += global_step_tokens
1543
-
1544
- # Loss is averaged approximately across ranks here.
1545
- loss_stats = torch.tensor(
1546
- [step_loss_sum, step_loss_tokens],
1547
- device=device,
1548
- dtype=torch.float64,
1549
- )
1550
- all_reduce_sum(loss_stats)
1551
-
1552
- global_loss_sum = float(loss_stats[0].item())
1553
- global_loss_tokens = int(loss_stats[1].item())
1554
-
1555
- interval_loss_sum += global_loss_sum
1556
- interval_loss_tokens += global_loss_tokens
1557
-
1558
- completed_step = step + 1
1559
-
1560
- if (
1561
- completed_step % args.log_interval_steps == 0
1562
- and is_main_process()
1563
- ):
1564
- mean_step_loss = (
1565
- global_loss_sum / global_loss_tokens
1566
- )
1567
-
1568
- logger.log(
1569
- f"step {completed_step:07d}: "
1570
- f"loss {mean_step_loss:.4f}, "
1571
- f"lr {lr:.8f}, "
1572
- f"grad_norm {float(grad_norm):.4f}, "
1573
- f"tokens_seen={human_tokens(tokens_seen)}, "
1574
- f"{timestamp()}"
1575
- )
1576
-
1577
- now = time.monotonic()
1578
-
1579
- '''diagnostic_due = (
1580
- now - last_diagnostic_time
1581
- >= args.diagnostic_interval_seconds
1582
- )'''
1583
-
1584
- diagnostic_due = (
1585
- completed_step % 700 == 0
1586
- )
1587
-
1588
- final_step = completed_step == args.max_steps
1589
-
1590
- token_budget_reached = (
1591
- args.max_tokens > 0
1592
- and tokens_seen >= args.max_tokens
1593
- )
1594
-
1595
- if diagnostic_due or final_step or token_budget_reached:
1596
- if distributed:
1597
- dist.barrier()
1598
-
1599
- elapsed = max(
1600
- now - interval_start_time,
1601
- 1e-9,
1602
- )
1603
- interval_tokens = (
1604
- tokens_seen - interval_start_tokens
1605
- )
1606
- tokens_per_second = interval_tokens / elapsed
1607
-
1608
- train_loss = (
1609
- interval_loss_sum
1610
- / max(interval_loss_tokens, 1)
1611
- )
1612
-
1613
- val_loss, val_tokens = estimate_loss(
1614
- model=model,
1615
- dataset=valid_data,
1616
- device=device,
1617
- batch_size=args.eval_batch_size,
1618
- eval_batches=args.eval_batches,
1619
- seed=args.seed + 123_456,
1620
- )
1621
-
1622
- if is_main_process():
1623
- train_ppl = math.exp(min(train_loss, 20.0))
1624
- val_ppl = math.exp(min(val_loss, 20.0))
1625
-
1626
- logger.log(
1627
- f"step {completed_step:07d}: "
1628
- f"train loss {train_loss:.4f}, "
1629
- f"val loss {val_loss:.4f}, "
1630
- f"Train PPL {train_ppl:.3f}, "
1631
- f"Val PPL {val_ppl:.3f}, "
1632
- f"tokens_seen={human_tokens(tokens_seen)}, "
1633
- f"val_tokens={human_tokens(val_tokens)}, "
1634
- f"toks/s={tokens_per_second:,.0f}, "
1635
- f"{timestamp()}"
1636
- )
1637
-
1638
- should_log_input_repr = (
1639
- completed_step == 1
1640
- or completed_step % args.milestone_save_interval_steps == 0
1641
- or final_step
1642
- or token_budget_reached
1643
- )
1644
-
1645
- if should_log_input_repr:
1646
- log_input_representation_diagnostics(
1647
- raw_model=raw_model,
1648
- tokenizer=tokenizer,
1649
- logger=logger,
1650
- device=device,
1651
- step=completed_step,
1652
- )
1653
-
1654
- greedy_prompts = prompts
1655
- sample_prompts = prompts
1656
-
1657
- greedy_generations = generate_diagnostics(
1658
- raw_model=raw_model,
1659
- tokenizer=tokenizer,
1660
- prompts=greedy_prompts,
1661
- device=device,
1662
- max_new_tokens=args.generation_tokens,
1663
- temperature=0.0,
1664
- top_k=None,
1665
- )
1666
-
1667
- sample_generations = generate_diagnostics(
1668
- raw_model=raw_model,
1669
- tokenizer=tokenizer,
1670
- prompts=sample_prompts,
1671
- device=device,
1672
- max_new_tokens=args.generation_tokens,
1673
- temperature=0.8,
1674
- top_k=50,
1675
- )
1676
-
1677
- for prompt, text in zip(greedy_prompts, greedy_generations):
1678
- logger.log(
1679
- f"step {completed_step:07d}: "
1680
- f"[greedy] prompt={prompt!r} => {text}"
1681
- )
1682
-
1683
- for prompt, text in zip(sample_prompts, sample_generations):
1684
- logger.log(
1685
- f"step {completed_step:07d}: "
1686
- f"[sample t=0.8 k=50] prompt={prompt!r} => {text}"
1687
- )
1688
-
1689
- logger.log(
1690
- f"step {completed_step:07d}: "
1691
- f"LR: {lr:.8f}, "
1692
- f"opt_step: {completed_step}, "
1693
- f"{timestamp()}"
1694
- )
1695
-
1696
- if distributed:
1697
- dist.barrier()
1698
-
1699
- last_diagnostic_time = now
1700
- interval_start_time = now
1701
- interval_start_tokens = tokens_seen
1702
- interval_loss_sum = 0.0
1703
- interval_loss_tokens = 0
1704
-
1705
- save_latest = (
1706
- completed_step % args.latest_save_interval_steps == 0
1707
- or final_step
1708
- or token_budget_reached
1709
- )
1710
-
1711
- save_milestone = (
1712
- completed_step % args.milestone_save_interval_steps == 0
1713
- or final_step
1714
- or token_budget_reached
1715
- )
1716
-
1717
- if save_latest or save_milestone:
1718
- if distributed:
1719
- dist.barrier()
1720
-
1721
- if save_latest:
1722
- save_checkpoint(
1723
- path=os.path.join(
1724
- args.output_dir,
1725
- "checkpoint_latest.pt",
1726
- ),
1727
- raw_model=raw_model,
1728
- optimizer=optimizer,
1729
- step=completed_step,
1730
- tokens_seen=tokens_seen,
1731
- args=args,
1732
- )
1733
-
1734
- save_model_safetensors(
1735
- path=os.path.join(
1736
- args.output_dir,
1737
- "model_latest.safetensors",
1738
- ),
1739
- raw_model=raw_model,
1740
- )
1741
-
1742
- if save_milestone:
1743
- save_checkpoint(
1744
- path=os.path.join(
1745
- args.output_dir,
1746
- f"checkpoint_{completed_step:07d}.pt",
1747
- ),
1748
- raw_model=raw_model,
1749
- optimizer=optimizer,
1750
- step=completed_step,
1751
- tokens_seen=tokens_seen,
1752
- args=args,
1753
- )
1754
-
1755
- save_model_safetensors(
1756
- path=os.path.join(
1757
- args.output_dir,
1758
- f"model_{completed_step:07d}.safetensors",
1759
- ),
1760
- raw_model=raw_model,
1761
- )
1762
-
1763
- if save_latest or save_milestone:
1764
- if distributed:
1765
- dist.barrier()
1766
-
1767
- if token_budget_reached:
1768
- if is_main_process():
1769
- logger.log(
1770
- f"[stop] max_tokens reached: "
1771
- f"{tokens_seen:,}"
1772
- )
1773
- break
1774
-
1775
- cleanup_distributed()
1776
-
1777
-
1778
- if __name__ == "__main__":
1779
- main()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
training_original/model_16_dim_bin.py DELETED
@@ -1,208 +0,0 @@
1
- #!/usr/bin/env python3
2
-
3
- from typing import Optional
4
-
5
- import torch
6
- import torch.nn as nn
7
-
8
- from classic_model import ClassicConfig, ClassicForCausalLM
9
-
10
-
11
- class Binary16Config(ClassicConfig):
12
- model_type = "binary16_causal_lm"
13
-
14
- def __init__(
15
- self,
16
- binary_dim: int = 16,
17
- binary_encoding: str = "zero_one",
18
- binary_scale: float = 1.0,
19
- binary_permutation_seed: Optional[int] = None,
20
- **kwargs,
21
- ):
22
- super().__init__(**kwargs)
23
-
24
- if binary_dim != 16:
25
- raise ValueError(
26
- "This implementation requires binary_dim=16"
27
- )
28
-
29
- if self.vocab_size > (1 << binary_dim):
30
- raise ValueError(
31
- f"vocab_size={self.vocab_size} exceeds "
32
- f"2**{binary_dim}"
33
- )
34
-
35
- if self.d_model % binary_dim != 0:
36
- raise ValueError(
37
- f"d_model={self.d_model} must be divisible "
38
- f"by binary_dim={binary_dim}"
39
- )
40
-
41
- if binary_encoding not in ("zero_one", "bipolar"):
42
- raise ValueError(
43
- "binary_encoding must be zero_one or bipolar"
44
- )
45
-
46
- if self.tie_word_embeddings:
47
- raise ValueError(
48
- "tie_word_embeddings must be False for "
49
- "fixed binary input"
50
- )
51
-
52
- self.binary_dim = binary_dim
53
- self.binary_encoding = binary_encoding
54
- self.binary_scale = float(binary_scale)
55
- self.binary_permutation_seed = binary_permutation_seed
56
- self.binary_repeat = self.d_model // binary_dim
57
-
58
-
59
- def build_binary_codebook(
60
- vocab_size: int,
61
- binary_dim: int,
62
- encoding: str,
63
- permutation_seed: Optional[int],
64
- ) -> torch.Tensor:
65
- if vocab_size > (1 << binary_dim):
66
- raise ValueError(
67
- f"vocab_size={vocab_size} does not fit "
68
- f"in {binary_dim} bits"
69
- )
70
-
71
- if permutation_seed is None:
72
- code_ids = torch.arange(
73
- vocab_size,
74
- dtype=torch.int64,
75
- )
76
- else:
77
- generator = torch.Generator(device="cpu")
78
- generator.manual_seed(permutation_seed)
79
-
80
- code_ids = torch.randperm(
81
- 1 << binary_dim,
82
- generator=generator,
83
- dtype=torch.int64,
84
- )[:vocab_size]
85
-
86
- shifts = torch.arange(
87
- binary_dim,
88
- dtype=torch.int64,
89
- )
90
-
91
- codebook = (
92
- (code_ids[:, None] >> shifts[None, :]) & 1
93
- ).to(torch.float32)
94
-
95
- if encoding == "bipolar":
96
- codebook = codebook.mul(2.0).sub(1.0)
97
- elif encoding != "zero_one":
98
- raise ValueError(f"Unknown encoding: {encoding}")
99
-
100
- return codebook.contiguous()
101
-
102
-
103
- class FixedBinary16Embedding(nn.Module):
104
- def __init__(self, config: Binary16Config):
105
- super().__init__()
106
-
107
- codebook = build_binary_codebook(
108
- vocab_size=config.vocab_size,
109
- binary_dim=config.binary_dim,
110
- encoding=config.binary_encoding,
111
- permutation_seed=config.binary_permutation_seed,
112
- )
113
-
114
- # Buffer, не Parameter.
115
- self.register_buffer(
116
- "codebook",
117
- codebook,
118
- persistent=True,
119
- )
120
-
121
- self.vocab_size = config.vocab_size
122
- self.binary_dim = config.binary_dim
123
- self.d_model = config.d_model
124
- self.repeat = config.binary_repeat
125
- self.binary_scale = config.binary_scale
126
-
127
- @property
128
- def weight(self):
129
- # Совместимость с get_input_embeddings().
130
- # Возвращается buffer, не trainable Parameter.
131
- return self.codebook
132
-
133
- def forward(self, input_ids: torch.Tensor) -> torch.Tensor:
134
- code = self.codebook[input_ids.long()]
135
- x = code.repeat(
136
- *([1] * (code.ndim - 1)),
137
- self.repeat,
138
- )
139
-
140
- if self.binary_scale != 1.0:
141
- x = x * self.binary_scale
142
-
143
- return x
144
-
145
- class Binary16ForCausalLM(ClassicForCausalLM):
146
- config_class = Binary16Config
147
-
148
- def __init__(self, config: Binary16Config):
149
- # Временно создаётся стандартная embedding-таблица,
150
- # затем немедленно удаляется до optimizer construction.
151
- super().__init__(config)
152
-
153
- self.token_embeddings = FixedBinary16Embedding(
154
- config
155
- )
156
-
157
- def get_input_embeddings(self):
158
- return self.token_embeddings
159
-
160
- def set_input_embeddings(self, value):
161
- raise RuntimeError(
162
- "Binary16ForCausalLM has a fixed input interface"
163
- )
164
-
165
- def tie_weights(self):
166
- if getattr(
167
- self.config,
168
- "tie_word_embeddings",
169
- False,
170
- ):
171
- raise ValueError(
172
- "Fixed binary input cannot be tied "
173
- "to the output projection"
174
- )
175
-
176
- def count_parameters(self):
177
- total = sum(
178
- parameter.numel()
179
- for parameter in self.parameters()
180
- )
181
-
182
- trainable = sum(
183
- parameter.numel()
184
- for parameter in self.parameters()
185
- if parameter.requires_grad
186
- )
187
-
188
- frozen_parameters = sum(
189
- parameter.numel()
190
- for parameter in self.parameters()
191
- if not parameter.requires_grad
192
- )
193
-
194
- output_parameters = (
195
- self.lm_head.weight.numel()
196
- )
197
-
198
- return {
199
- "total_parameters": total,
200
- "trainable_parameters": trainable,
201
- "frozen_parameters": frozen_parameters,
202
- "input_trainable_parameters": 0,
203
- "fixed_codebook_values":
204
- self.token_embeddings.codebook.numel(),
205
- "output_parameters": output_parameters,
206
- "body_parameters":
207
- total - output_parameters,
208
- }