-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathdiffusion_builders.py
More file actions
137 lines (116 loc) · 5.28 KB
/
Copy pathdiffusion_builders.py
File metadata and controls
137 lines (116 loc) · 5.28 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
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
from megatron.core.models.diffusion.diffusion_model import DiffusionLanguageModel
from megatron.core.models.diffusion.editflow_model import EditFlowLanguageModel
from megatron.core.models.diffusion.diffusion_layer_specs import (
get_diffusion_layer_local_spec,
get_diffusion_layer_with_transformer_engine_spec,
)
try:
from megatron.core.models.diffusion.diffusion_layer_specs import get_bd3lm_flex_layer_local_spec
_HAS_FLEX_SPEC = True
except ImportError:
_HAS_FLEX_SPEC = False
from megatron.training import print_rank_0
from megatron.training.arguments import core_transformer_config_from_args
def diffusion_builder(args, pre_process, post_process, vp_stage=None, config=None, pg_collection=None):
print_rank_0('building Diffusion Language Model ...')
if config is None:
config = core_transformer_config_from_args(args)
use_te = args.transformer_impl == "transformer_engine"
if use_te:
transformer_layer_spec = get_diffusion_layer_with_transformer_engine_spec()
else:
transformer_layer_spec = get_diffusion_layer_local_spec(
qk_layernorm=config.qk_layernorm,
)
model = DiffusionLanguageModel(
config=config,
transformer_layer_spec=transformer_layer_spec,
vocab_size=args.padded_vocab_size,
max_sequence_length=args.max_position_embeddings,
pre_process=pre_process,
post_process=post_process,
fp16_lm_cross_entropy=args.fp16_lm_cross_entropy,
parallel_output=True,
share_embeddings_and_output_weights=not args.untie_embeddings_and_output_weights,
position_embedding_type=args.position_embedding_type,
rotary_percent=args.rotary_percent,
rotary_base=args.rotary_base,
vp_stage=vp_stage,
pg_collection=pg_collection,
)
return model
def bd3lm_builder(args, pre_process, post_process, vp_stage=None, config=None, pg_collection=None):
"""Build a DiffusionLanguageModel for BD3LM training.
Selects flex_attention layer spec when --diffusion-attn-backend flex is set,
otherwise uses the standard dense attention spec.
"""
print_rank_0('building BD3LM Diffusion Language Model ...')
if config is None:
config = core_transformer_config_from_args(args)
attn_backend = getattr(args, 'diffusion_attn_backend', 'dense')
use_te = args.transformer_impl == "transformer_engine"
if attn_backend == 'flex' and not use_te and _HAS_FLEX_SPEC:
print_rank_0(' Using flex_attention backend for BD3LM block-sparse attention.')
transformer_layer_spec = get_bd3lm_flex_layer_local_spec(
qk_layernorm=config.qk_layernorm,
)
elif use_te:
transformer_layer_spec = get_diffusion_layer_with_transformer_engine_spec()
else:
transformer_layer_spec = get_diffusion_layer_local_spec(
qk_layernorm=config.qk_layernorm,
)
model = DiffusionLanguageModel(
config=config,
transformer_layer_spec=transformer_layer_spec,
vocab_size=args.padded_vocab_size,
max_sequence_length=args.max_position_embeddings,
pre_process=pre_process,
post_process=post_process,
fp16_lm_cross_entropy=args.fp16_lm_cross_entropy,
parallel_output=True,
share_embeddings_and_output_weights=not args.untie_embeddings_and_output_weights,
position_embedding_type=args.position_embedding_type,
rotary_percent=args.rotary_percent,
rotary_base=args.rotary_base,
vp_stage=vp_stage,
pg_collection=pg_collection,
)
return model
def editflow_builder(args, pre_process, post_process, vp_stage=None, config=None, pg_collection=None):
"""Build an EditFlowLanguageModel.
Identical architecture to DiffusionLanguageModel plus three extra output heads:
sub_logits : token distribution for substitutions [b, s, v]
ins_logits : token distribution for insertions [b, s, v]
rate_heads : per-position rates (sub/del/ins) via Softplus [b, s, 3]
The EditFlow loss (survival + positive edit terms) is computed externally
in pretrain_editflow.py, not inside the model.
"""
print_rank_0('building EditFlow Language Model ...')
if config is None:
config = core_transformer_config_from_args(args)
use_te = args.transformer_impl == "transformer_engine"
if use_te:
transformer_layer_spec = get_diffusion_layer_with_transformer_engine_spec()
else:
transformer_layer_spec = get_diffusion_layer_local_spec(
qk_layernorm=config.qk_layernorm,
)
model = EditFlowLanguageModel(
config=config,
transformer_layer_spec=transformer_layer_spec,
vocab_size=args.padded_vocab_size,
max_sequence_length=args.max_position_embeddings,
pre_process=pre_process,
post_process=post_process,
fp16_lm_cross_entropy=args.fp16_lm_cross_entropy,
parallel_output=True,
share_embeddings_and_output_weights=not args.untie_embeddings_and_output_weights,
position_embedding_type=args.position_embedding_type,
rotary_percent=args.rotary_percent,
rotary_base=args.rotary_base,
vp_stage=vp_stage,
pg_collection=pg_collection,
)
return model