Skip to content

Commit 02bfc0e

Browse files
committed
add gpt oss preset
1 parent db23ef4 commit 02bfc0e

5 files changed

Lines changed: 490 additions & 5 deletions

File tree

‎roundpipe/models/__init__.py‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -40,6 +40,7 @@
4040

4141
SUPPORTED_MODELS = {
4242
"function": ".function",
43+
"GptOssForCausalLM": ".gpt_oss",
4344
"LlamaForCausalLM": ".llama",
4445
"Qwen3MoeForCausalLM": ".qwen3_moe",
4546
"Qwen3ForCausalLM": ".qwen3",

‎roundpipe/models/gpt_oss.py‎

Lines changed: 364 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,364 @@
1+
from typing_extensions import *
2+
import warnings
3+
4+
import torch
5+
import torch.nn as nn
6+
from transformers.masking_utils import (
7+
create_causal_mask,
8+
create_sliding_window_causal_mask,
9+
)
10+
from transformers.modeling_outputs import MoeCausalLMOutputWithPast
11+
from transformers.models.gpt_oss.modeling_gpt_oss import (
12+
GptOssForCausalLM,
13+
GptOssDecoderLayer,
14+
GptOssExperts,
15+
load_balancing_loss_func,
16+
)
17+
18+
from ..context import doing_recompute, save_for_recompute, get_recompute_data
19+
from ..roundpipe import RoundPipe
20+
from .function import CompileForCausalLMLoss, ChunkedCompileLinearForCausalLMLoss
21+
22+
23+
class GptOssOptExperts(nn.Module):
24+
def __init__(self, mod: GptOssExperts) -> None:
25+
super().__init__()
26+
self.num_experts = mod.num_experts
27+
self.hidden_size = mod.hidden_size
28+
self.alpha = mod.alpha
29+
self.limit = mod.limit
30+
31+
self.gate_up_proj = mod.gate_up_proj
32+
self.gate_up_proj_bias = mod.gate_up_proj_bias
33+
self.down_proj = mod.down_proj
34+
self.down_proj_bias = mod.down_proj_bias
35+
36+
def forward(
37+
self,
38+
hidden_states: torch.Tensor,
39+
router_indices: torch.Tensor,
40+
routing_weights: torch.Tensor,
41+
) -> torch.Tensor:
42+
batch_size, sequence_length, hidden_dim = hidden_states.shape
43+
hidden_states = hidden_states.view(-1, hidden_dim)
44+
45+
top_k = router_indices.shape[-1]
46+
selected_experts = router_indices.view(-1)
47+
routing_weights = torch.gather(routing_weights, 1, router_indices)
48+
routing_weights = routing_weights.view(-1)
49+
50+
_, sort_idx = torch.sort(selected_experts)
51+
permute_weight = routing_weights[sort_idx]
52+
batch_idx = sort_idx.div(top_k, rounding_mode="floor")
53+
54+
if doing_recompute():
55+
(token_per_expert_cpu,) = get_recompute_data()
56+
else:
57+
token_per_expert = torch.zeros(
58+
self.num_experts, dtype=torch.long, device=hidden_states.device
59+
)
60+
token_per_expert.index_add_(
61+
0,
62+
selected_experts,
63+
torch.ones_like(selected_experts, dtype=torch.long),
64+
)
65+
token_per_expert_cpu = token_per_expert.cpu().numpy()
66+
save_for_recompute(token_per_expert_cpu)
67+
68+
final_hidden_states = torch.zeros(
69+
(batch_size * sequence_length, hidden_dim),
70+
dtype=hidden_states.dtype,
71+
device=hidden_states.device,
72+
)
73+
start_idx = 0
74+
for expert_id in range(self.num_experts):
75+
num_tokens = token_per_expert_cpu[expert_id]
76+
if num_tokens == 0:
77+
continue
78+
expert_tokens = batch_idx[start_idx : start_idx + num_tokens]
79+
expert_input = hidden_states[expert_tokens]
80+
81+
gate_up = (
82+
expert_input @ self.gate_up_proj[expert_id]
83+
+ self.gate_up_proj_bias[expert_id]
84+
)
85+
gate, up = gate_up[..., ::2], gate_up[..., 1::2]
86+
gate = gate.clamp(min=None, max=self.limit)
87+
up = up.clamp(min=-self.limit, max=self.limit)
88+
glu = gate * torch.sigmoid(gate * self.alpha)
89+
gated_output = (up + 1) * glu
90+
expert_output = (
91+
gated_output @ self.down_proj[expert_id]
92+
+ self.down_proj_bias[expert_id]
93+
)
94+
expert_output *= permute_weight[
95+
start_idx : start_idx + num_tokens
96+
].unsqueeze(-1)
97+
final_hidden_states.index_add_(0, expert_tokens, expert_output)
98+
start_idx += num_tokens
99+
100+
return final_hidden_states.view(batch_size, sequence_length, hidden_dim)
101+
102+
103+
class GptOssForCausalLMPrefix(nn.Module):
104+
def __init__(self, model: GptOssForCausalLM) -> None:
105+
super().__init__()
106+
self.embed_tokens = model.model.embed_tokens
107+
self.rotary_emb = model.model.rotary_emb
108+
self.config = model.model.config
109+
110+
def forward(
111+
self,
112+
input_ids: Optional[torch.Tensor] = None,
113+
attention_mask: Optional[torch.Tensor] = None,
114+
position_ids: Optional[torch.Tensor] = None,
115+
past_key_values: Optional[Any] = None,
116+
inputs_embeds: Optional[torch.Tensor] = None,
117+
labels: Optional[torch.Tensor] = None,
118+
use_cache: Optional[bool] = None,
119+
output_router_logits: Optional[bool] = None,
120+
cache_position: Optional[torch.Tensor] = None,
121+
logits_to_keep: Union[int, torch.Tensor] = 0,
122+
**kwargs: Any,
123+
):
124+
if (input_ids is None) ^ (inputs_embeds is not None):
125+
raise ValueError(
126+
"You must specify exactly one of input_ids or inputs_embeds"
127+
)
128+
129+
if inputs_embeds is None:
130+
inputs_embeds = cast(torch.Tensor, self.embed_tokens(input_ids))
131+
132+
if output_router_logits is None:
133+
output_router_logits = self.config.output_router_logits
134+
135+
if doing_recompute():
136+
causal_mask_mapping, position_ids, position_embeddings = (
137+
get_recompute_data()
138+
)
139+
return (
140+
inputs_embeds,
141+
causal_mask_mapping,
142+
position_ids,
143+
position_embeddings,
144+
kwargs,
145+
labels,
146+
logits_to_keep,
147+
output_router_logits,
148+
attention_mask,
149+
[], # router_logits
150+
)
151+
152+
if use_cache:
153+
warnings.warn(
154+
"`use_cache` will set to False. Caching behavior is not supported in RoundPipe."
155+
)
156+
use_cache = False
157+
if past_key_values is not None:
158+
warnings.warn(
159+
"`past_key_values` will be ignored. Caching behavior is not supported in RoundPipe."
160+
)
161+
past_key_values = None
162+
163+
if cache_position is None:
164+
past_seen_tokens = (
165+
past_key_values.get_seq_length() if past_key_values is not None else 0
166+
)
167+
cache_position = torch.arange(
168+
past_seen_tokens,
169+
past_seen_tokens + inputs_embeds.shape[1],
170+
device=inputs_embeds.device,
171+
)
172+
if position_ids is None:
173+
position_ids = cache_position.unsqueeze(0)
174+
175+
if not isinstance(causal_mask_mapping := attention_mask, dict):
176+
mask_kwargs = {
177+
"config": self.config,
178+
"input_embeds": inputs_embeds,
179+
"attention_mask": attention_mask,
180+
"cache_position": cache_position,
181+
"past_key_values": past_key_values,
182+
}
183+
causal_mask_mapping = {
184+
"full_attention": create_causal_mask(**mask_kwargs),
185+
"sliding_attention": create_sliding_window_causal_mask(**mask_kwargs),
186+
}
187+
188+
hidden_states = inputs_embeds
189+
position_embeddings = self.rotary_emb(hidden_states, position_ids)
190+
191+
save_for_recompute(causal_mask_mapping, position_ids, position_embeddings)
192+
return (
193+
hidden_states,
194+
causal_mask_mapping,
195+
position_ids,
196+
position_embeddings,
197+
kwargs,
198+
labels,
199+
logits_to_keep,
200+
output_router_logits,
201+
attention_mask,
202+
[], # router_logits
203+
)
204+
205+
206+
class GptOssForCausalLMWrappedLayer(nn.Module):
207+
def __init__(self, layer: GptOssDecoderLayer) -> None:
208+
super().__init__()
209+
self.hidden_size = layer.hidden_size
210+
self.self_attn = layer.self_attn
211+
self.mlp = layer.mlp
212+
self.input_layernorm = layer.input_layernorm
213+
self.post_attention_layernorm = layer.post_attention_layernorm
214+
self.attention_type = layer.attention_type
215+
216+
def forward(self, input):
217+
(
218+
hidden_states,
219+
causal_mask_mapping,
220+
position_ids,
221+
position_embeddings,
222+
kwargs,
223+
labels,
224+
logits_to_keep,
225+
output_router_logits,
226+
attention_mask,
227+
router_logits,
228+
) = input
229+
230+
residual = hidden_states
231+
232+
hidden_states = self.input_layernorm(hidden_states)
233+
234+
# Self Attention
235+
hidden_states, _ = self.self_attn(
236+
hidden_states=hidden_states,
237+
attention_mask=causal_mask_mapping[self.attention_type],
238+
position_ids=position_ids,
239+
past_key_values=None,
240+
use_cache=False,
241+
cache_position=None,
242+
position_embeddings=position_embeddings,
243+
**kwargs,
244+
)
245+
hidden_states = residual + hidden_states
246+
247+
# Fully Connected
248+
residual = hidden_states
249+
hidden_states = self.post_attention_layernorm(hidden_states)
250+
hidden_states, router_logit = self.mlp(hidden_states)
251+
if output_router_logits:
252+
router_logits.append(router_logit)
253+
hidden_states = residual + hidden_states
254+
255+
return (
256+
hidden_states,
257+
causal_mask_mapping,
258+
position_ids,
259+
position_embeddings,
260+
kwargs,
261+
labels,
262+
logits_to_keep,
263+
output_router_logits,
264+
attention_mask,
265+
router_logits,
266+
)
267+
268+
269+
class GptOssForCausalLMPostfix(nn.Module):
270+
def __init__(self, model: GptOssForCausalLM) -> None:
271+
super().__init__()
272+
self.norm = model.model.norm
273+
self.vocab_size = model.config.vocab_size
274+
self.lm_head = model.lm_head
275+
self.loss_function = model.loss_function
276+
277+
self.num_experts: int = model.num_experts
278+
self.num_experts_per_tok: int = model.num_experts_per_tok
279+
self.router_aux_loss_coef: float = model.router_aux_loss_coef
280+
281+
def forward(self, input) -> MoeCausalLMOutputWithPast:
282+
(
283+
hidden_states,
284+
causal_mask_mapping,
285+
position_ids,
286+
position_embeddings,
287+
kwargs,
288+
labels,
289+
logits_to_keep,
290+
output_router_logits,
291+
attention_mask,
292+
router_logits,
293+
) = input
294+
hidden_states = self.norm(hidden_states)
295+
296+
# Only compute necessary logits, and do not upcast them to float if we are not computing the loss
297+
slice_indices = (
298+
slice(-logits_to_keep, None)
299+
if isinstance(logits_to_keep, int)
300+
else logits_to_keep
301+
)
302+
logits = None
303+
if kwargs.get("return_logits", True):
304+
logits = self.lm_head(hidden_states[:, slice_indices, :])
305+
306+
loss = None
307+
if labels is not None:
308+
if logits is None:
309+
loss = ChunkedCompileLinearForCausalLMLoss(
310+
hidden_states[:, slice_indices, :],
311+
self.lm_head,
312+
labels,
313+
**kwargs,
314+
)
315+
else:
316+
loss = self.loss_function(
317+
logits=logits, labels=labels, vocab_size=self.vocab_size, **kwargs
318+
)
319+
320+
aux_loss = None
321+
if output_router_logits:
322+
aux_loss = cast(
323+
torch.FloatTensor,
324+
load_balancing_loss_func(
325+
tuple(t.float() for t in router_logits),
326+
self.num_experts,
327+
self.num_experts_per_tok,
328+
attention_mask,
329+
),
330+
)
331+
if loss is not None:
332+
loss += self.router_aux_loss_coef * aux_loss.to(
333+
loss.device
334+
) # make sure to reside in the same device
335+
336+
return MoeCausalLMOutputWithPast(
337+
loss=cast(Optional[torch.FloatTensor], loss),
338+
aux_loss=aux_loss,
339+
logits=logits,
340+
router_logits=router_logits,
341+
)
342+
343+
344+
EXPECTED_MODEL_CLASS = GptOssForCausalLM
345+
346+
347+
def wrap_model(model: GptOssForCausalLM, **roundpipe_kwargs: Any) -> RoundPipe:
348+
model.loss_function = CompileForCausalLMLoss
349+
350+
for layer in model.model.layers:
351+
layer = cast(GptOssDecoderLayer, layer)
352+
layer.mlp.experts = cast(GptOssExperts, GptOssOptExperts(layer.mlp.experts))
353+
354+
prefix = GptOssForCausalLMPrefix(model)
355+
layers = [
356+
GptOssForCausalLMWrappedLayer(cast(GptOssDecoderLayer, layer))
357+
for layer in model.model.layers
358+
]
359+
postfix = GptOssForCausalLMPostfix(model)
360+
wrapped_model = RoundPipe(
361+
nn.Sequential(prefix, *layers, postfix), **roundpipe_kwargs
362+
)
363+
wrapped_model.set_original_model(model)
364+
return wrapped_model

‎roundpipe/models/qwen3_moe.py‎

Lines changed: 3 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -201,14 +201,13 @@ def forward(self, hidden_states: torch.Tensor) -> Tuple[torch.Tensor, torch.Tens
201201
num_tokens = token_per_expert_cpu[expert_id]
202202
if num_tokens == 0:
203203
continue
204-
expert_input = hidden_states[batch_idx[start_idx : start_idx + num_tokens]]
204+
expert_tokens = batch_idx[start_idx : start_idx + num_tokens]
205+
expert_input = hidden_states[expert_tokens]
205206
expert_output = self.experts[expert_id](expert_input)
206207
expert_output *= permute_weight[
207208
start_idx : start_idx + num_tokens
208209
].unsqueeze(-1)
209-
final_hidden_states.index_add_(
210-
0, batch_idx[start_idx : start_idx + num_tokens], expert_output
211-
)
210+
final_hidden_states.index_add_(0, expert_tokens, expert_output)
212211
start_idx += num_tokens
213212

214213
final_hidden_states = final_hidden_states.reshape(

‎setup.cfg‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,6 @@
11
[metadata]
22
name = RoundPipe
3-
version = 0.1.0
3+
version = 0.1.1
44
author = ITcarrot
55
author_email = luo-yb25@mails.tsinghua.edu.cn
66
description = Large DNNs training framework for consumer GPUs

0 commit comments

Comments
 (0)