-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathconfig_registry.py
More file actions
115 lines (102 loc) · 4.44 KB
/
Copy pathconfig_registry.py
File metadata and controls
115 lines (102 loc) · 4.44 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
# Copyright (c) 2026 Advanced Micro Devices, Inc.
#
# SPDX-License-Identifier: MIT
from torchtitan.trainer import Trainer
from torchtitan.protocols.model_converter import ModelConvertersContainer
from torchtitan.models.gpt_oss.config_registry import (
gpt_oss_debugmodel as gpt_oss_debugmodel_orig,
gpt_oss_20b as gpt_oss_20b_orig,
)
from alto.components.converter import ModelOptConverter
__all__ = [
"gpt_oss_debugmodel",
"gpt_oss_debugmodel_lpt",
"gpt_oss_20b",
"gpt_oss_20b_pretrain",
"gpt_oss_20b_lpt",
]
def gpt_oss_debugmodel() -> Trainer.Config:
config = gpt_oss_debugmodel_orig()
config.profiling.enable_profiling = False
config.training.steps = 10
config.training.local_batch_size = 4
config.training.global_batch_size = 16
config.training.seq_len = 2048
config.activation_checkpoint.mode = "none"
config.debug.seed = 1234
return config
def gpt_oss_debugmodel_lpt() -> Trainer.Config:
config = gpt_oss_debugmodel()
config.model_converters = ModelConvertersContainer.Config(converters=[
ModelOptConverter.Config(recipe="./alto/models/gpt_oss/configs/lpt_recipe.yaml",),
],)
return config
def gpt_oss_20b() -> Trainer.Config:
config = gpt_oss_20b_orig()
config.hf_assets_path = "/huggingface/hub/models--openai--gpt-oss-20b/snapshots/6cee5e81ee83917806bbde320786a8fb61efebee/"
config.dump_folder = "gpt_oss_20b-outputs"
config.profiling.enable_profiling = False
config.training.steps = 0
config.training.local_batch_size = 1
config.training.seq_len = 8192
config.metrics.log_freq = 1
config.metrics.enable_tensorboard = True
config.dataloader.dataset = "c4_test"
config.parallelism.expert_parallel_degree = 1
config.parallelism.expert_tensor_parallel_degree = 1
config.parallelism.tensor_parallel_degree = 1
config.checkpoint.enable = True
config.checkpoint.initial_load_path = "/huggingface/hub/models--openai--gpt-oss-20b/snapshots/6cee5e81ee83917806bbde320786a8fb61efebee/"
config.checkpoint.initial_load_in_hf = True
config.checkpoint.initial_load_in_hf_quantized = True
config.checkpoint.interval = 100
config.validator.enable = True
config.validator.dataloader.dataset = "wikitext_test"
config.validator.freq = 10
config.validator.steps = 10
config.activation_checkpoint.mode = "none"
config.debug.seed = 1234
return config
def gpt_oss_20b_pretrain() -> Trainer.Config:
config = gpt_oss_20b_orig()
config.hf_assets_path = "/huggingface/hub/models--openai--gpt-oss-20b/snapshots/6cee5e81ee83917806bbde320786a8fb61efebee/"
config.dump_folder = "gpt_oss_20b-pretrain-subset-lr4e-4-outputs"
config.profiling.enable_profiling = False
config.training.steps = 1200000
config.training.local_batch_size = 1
config.training.global_batch_size = 16
config.training.seq_len = 8192
config.optimizer.lr = 4e-4
config.optimizer.weight_decay = 0.1
config.optimizer.beta1 = 0.9
config.optimizer.beta2 = 0.95
config.optimizer.eps = 1e-5
config.lr_scheduler.min_lr_factor = 0.1
config.lr_scheduler.warmup_steps = 128
config.lr_scheduler.decay_ratio = 1 - 128 / config.training.steps
config.lr_scheduler.decay_type = "cosine"
config.metrics.log_freq = 1
config.metrics.enable_tensorboard = True
config.dataloader.dataset = "megatron"
config.dataloader.dataset_path = "/workspace/workspace/megatron_dataset/data/c4-train.en_6_text_document.idx"
config.parallelism.expert_parallel_degree = 8
config.parallelism.expert_tensor_parallel_degree = 1
config.parallelism.tensor_parallel_degree = 1
config.checkpoint.enable = True
config.checkpoint.interval = 1000
config.checkpoint.keep_latest_k = 2
config.validator.enable = True
config.validator.dataloader.dataset = "megatron"
config.validator.dataloader.dataset_path = "/workspace/workspace/megatron_dataset/data/c4-validation-91205-samples.en_text_document.idx"
config.validator.freq = 768
config.validator.steps = 64
config.activation_checkpoint.mode = "none"
config.debug.seed = 1234
return config
def gpt_oss_20b_lpt() -> Trainer.Config:
config = gpt_oss_20b_pretrain()
config.dump_folder = "gpt_oss_20b-pretrain-subset-mxfp4gemm_1d2d-hadamard-sr-lr4e-4-outputs"
config.model_converters = ModelConvertersContainer.Config(converters=[
ModelOptConverter.Config(recipe="./alto/models/gpt_oss/configs/lpt_recipe.yaml",),
],)
return config