Skip to content

Commit f3831bc

Browse files
committed
fix causal forward, prefetch, and remote code
1 parent 6c6dfad commit f3831bc

7 files changed

Lines changed: 106 additions & 53 deletions

File tree

‎xtuner/_lite/accelerate/dispatches/huggingface/internlm2.py‎

Lines changed: 28 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -622,33 +622,46 @@ def internlm2_causal_forward(
622622

623623
hidden_states = outputs[0]
624624

625+
loss = None
625626
if labels is None:
626627
logits = self.output(hidden_states)
627628
else:
628629

629-
if label_shifted:
630-
shift_hidden_states = hidden_states
631-
shift_labels = labels
632-
else:
633-
shift_hidden_states = hidden_states[..., :-1, :].contiguous()
634-
shift_labels = labels[..., 1:].contiguous()
630+
if liger_kernel_is_available():
631+
# unable to return logits when using Liger Kernel
632+
logits = None
635633

636-
shift_hidden_states = shift_hidden_states.view(-1, self.config.hidden_size)
637-
shift_labels = shift_labels.view(-1)
638-
shift_labels = shift_labels.to(shift_hidden_states.device)
634+
if label_shifted:
635+
shift_hidden_states = hidden_states
636+
shift_labels = labels
637+
else:
638+
shift_hidden_states = hidden_states[..., :-1, :].contiguous()
639+
shift_labels = labels[..., 1:].contiguous()
640+
641+
shift_hidden_states = shift_hidden_states.view(-1, self.config.hidden_size)
642+
shift_labels = shift_labels.view(-1)
643+
shift_labels = shift_labels.to(shift_hidden_states.device)
639644

640-
if liger_kernel_is_available():
641645
from liger_kernel.transformers.fused_linear_cross_entropy import LigerFusedLinearCrossEntropyLoss
642646

643647
loss_fct = LigerFusedLinearCrossEntropyLoss()
644648
loss = loss_fct(self.output.weight, shift_hidden_states, shift_labels, self.output.bias)
645649

646-
# unable to return logits when using Liger Kernel
647-
logits = None
648650
else:
649-
shift_logits = self.output(shift_hidden_states)
650-
651-
loss_fct = CrossEntropyLoss()
651+
logits = self.output(hidden_states)
652+
653+
if label_shifted:
654+
shift_logits = logits
655+
shift_labels = labels
656+
else:
657+
shift_logits = logits[..., :-1, :].contiguous()
658+
shift_labels = labels[..., 1:].contiguous()
659+
660+
shift_logits = shift_logits.view(-1, self.config.vocab_size)
661+
shift_labels = shift_labels.view(-1)
662+
shift_labels = shift_labels.to(shift_logits.device)
663+
664+
loss_fct = torch.nn.CrossEntropyLoss()
652665
loss = loss_fct(shift_logits, shift_labels)
653666

654667
if not return_dict:

‎xtuner/_lite/accelerate/dispatches/huggingface/llama.py‎

Lines changed: 27 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -209,32 +209,44 @@ def llama_casual_forward(
209209
hidden_states = outputs[0]
210210

211211
if labels is None:
212+
loss = None
212213
logits = self.lm_head(hidden_states)
213214
else:
215+
if liger_kernel_is_available():
216+
# unable to return logits when using Liger Kernel
217+
logits = None
214218

215-
if label_shifted:
216-
shift_hidden_states = hidden_states
217-
shift_labels = labels
218-
else:
219-
shift_hidden_states = hidden_states[..., :-1, :].contiguous()
220-
shift_labels = labels[..., 1:].contiguous()
219+
if label_shifted:
220+
shift_hidden_states = hidden_states
221+
shift_labels = labels
222+
else:
223+
shift_hidden_states = hidden_states[..., :-1, :].contiguous()
224+
shift_labels = labels[..., 1:].contiguous()
221225

222-
shift_hidden_states = shift_hidden_states.view(-1, self.config.hidden_size)
223-
shift_labels = shift_labels.view(-1)
224-
shift_labels = shift_labels.to(shift_hidden_states.device)
226+
shift_hidden_states = shift_hidden_states.view(-1, self.config.hidden_size)
227+
shift_labels = shift_labels.view(-1)
228+
shift_labels = shift_labels.to(shift_hidden_states.device)
225229

226-
if liger_kernel_is_available():
227230
from liger_kernel.transformers.fused_linear_cross_entropy import LigerFusedLinearCrossEntropyLoss
228231

229232
loss_fct = LigerFusedLinearCrossEntropyLoss()
230233
loss = loss_fct(self.lm_head.weight, shift_hidden_states, shift_labels, self.lm_head.bias)
231234

232-
# unable to return logits when using Liger Kernel
233-
logits = None
234235
else:
235-
shift_logits = self.lm_head(shift_hidden_states)
236-
237-
loss_fct = CrossEntropyLoss()
236+
logits = self.lm_head(hidden_states)
237+
238+
if label_shifted:
239+
shift_logits = logits
240+
shift_labels = labels
241+
else:
242+
shift_logits = logits[..., :-1, :].contiguous()
243+
shift_labels = labels[..., 1:].contiguous()
244+
245+
shift_logits = shift_logits.view(-1, self.config.vocab_size)
246+
shift_labels = shift_labels.view(-1)
247+
shift_labels = shift_labels.to(shift_logits.device)
248+
249+
loss_fct = torch.nn.CrossEntropyLoss()
238250
loss = loss_fct(shift_logits, shift_labels)
239251

240252
if not return_dict:

‎xtuner/_lite/accelerate/dispatches/huggingface/qwen2.py‎

Lines changed: 27 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -402,32 +402,44 @@ def qwen2_casual_forward(
402402
hidden_states = outputs[0]
403403

404404
if labels is None:
405+
loss = None
405406
logits = self.lm_head(hidden_states)
406407
else:
408+
if liger_kernel_is_available():
409+
# unable to return logits when using Liger Kernel
410+
logits = None
407411

408-
if label_shifted:
409-
shift_hidden_states = hidden_states
410-
shift_labels = labels
411-
else:
412-
shift_hidden_states = hidden_states[..., :-1, :].contiguous()
413-
shift_labels = labels[..., 1:].contiguous()
412+
if label_shifted:
413+
shift_hidden_states = hidden_states
414+
shift_labels = labels
415+
else:
416+
shift_hidden_states = hidden_states[..., :-1, :].contiguous()
417+
shift_labels = labels[..., 1:].contiguous()
414418

415-
shift_hidden_states = shift_hidden_states.view(-1, self.config.hidden_size)
416-
shift_labels = shift_labels.view(-1)
417-
shift_labels = shift_labels.to(shift_hidden_states.device)
419+
shift_hidden_states = shift_hidden_states.view(-1, self.config.hidden_size)
420+
shift_labels = shift_labels.view(-1)
421+
shift_labels = shift_labels.to(shift_hidden_states.device)
418422

419-
if liger_kernel_is_available():
420423
from liger_kernel.transformers.fused_linear_cross_entropy import LigerFusedLinearCrossEntropyLoss
421424

422425
loss_fct = LigerFusedLinearCrossEntropyLoss()
423426
loss = loss_fct(self.lm_head.weight, shift_hidden_states, shift_labels, self.lm_head.bias)
424427

425-
# unable to return logits when using Liger Kernel
426-
logits = None
427428
else:
428-
shift_logits = self.lm_head(shift_hidden_states)
429-
430-
loss_fct = CrossEntropyLoss()
429+
logits = self.lm_head(hidden_states)
430+
431+
if label_shifted:
432+
shift_logits = logits
433+
shift_labels = labels
434+
else:
435+
shift_logits = logits[..., :-1, :].contiguous()
436+
shift_labels = labels[..., 1:].contiguous()
437+
438+
shift_logits = shift_logits.view(-1, self.config.vocab_size)
439+
shift_labels = shift_labels.view(-1)
440+
shift_labels = shift_labels.to(shift_logits.device)
441+
442+
loss_fct = torch.nn.CrossEntropyLoss()
431443
loss = loss_fct(shift_logits, shift_labels)
432444

433445
if not return_dict:

‎xtuner/_lite/algorithms/ppo/model.py‎

Lines changed: 5 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -6,28 +6,25 @@
66
from xtuner._lite.accelerate import LoadWoInit
77

88

9-
def build_actor_model(model_path, dtype=torch.float32):
9+
def build_actor_model(model_path, dtype=torch.float32, trust_remote_code=True):
1010

1111
config = AutoConfig.from_pretrained(model_path, trust_remote_code=True)
1212
if is_flash_attn_2_available():
1313
config.attn_implementation = 'flash_attention_2'
1414
elif is_torch_sdpa_available():
1515
config.attn_implementation = 'sdpa'
1616

17-
# config.use_cache = False
18-
# config.torch_dtype = dtype
1917
with LoadWoInit():
2018
policy = AutoModelForCausalLM.from_pretrained(
2119
model_path,
2220
attn_implementation='flash_attention_2',
2321
torch_dtype=dtype,
24-
trust_remote_code=True)
22+
trust_remote_code=trust_remote_code)
2523

26-
# policy.to(dtype)
2724
return policy
2825

2926

30-
def build_reward_model(model_path, dtype=torch.float32):
27+
def build_reward_model(model_path, dtype=torch.float32, trust_remote_code=True):
3128

3229
config = AutoConfig.from_pretrained(model_path, trust_remote_code=True)
3330
if is_flash_attn_2_available():
@@ -42,8 +39,8 @@ def build_reward_model(model_path, dtype=torch.float32):
4239
model_path,
4340
attn_implementation='flash_attention_2',
4441
torch_dtype=dtype,
45-
trust_remote_code=True)
42+
trust_remote_code=trust_remote_code)
4643

4744
reward.model.use_cache = False
48-
# policy.to(dtype)
45+
4946
return reward

‎xtuner/_lite/parallel/megatron/internlm2.py‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,7 @@
11
from functools import partial
2+
from packaging import version
23

4+
import torch
35
from torch import nn
46
from torch.distributed._tensor import Replicate, distribute_tensor
57
from torch.distributed.tensor.parallel import (ColwiseParallel,
@@ -146,6 +148,10 @@ def megatron_internlm2(model,
146148
if i < num_recompute_layers:
147149
checkpoint(block)
148150

151+
if version.parse(torch.__version__) >= version.parse("2.5.0"):
152+
for layer_cur, layer_next in zip(model.layers[:-1], model.layers[1:]):
153+
layer_cur.set_modules_to_forward_prefetch([layer_next])
154+
149155
model.tok_embeddings.apply(param_init_fn)
150156
model.norm.apply(param_init_fn)
151157

‎xtuner/_lite/parallel/megatron/llama.py‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,7 @@
11
from functools import partial
2+
from packaging import version
23

4+
import torch
35
from torch import nn
46
from torch.distributed._tensor import Replicate, distribute_tensor
57
from torch.distributed.tensor.parallel import (ColwiseParallel,
@@ -147,6 +149,11 @@ def megatron_llama(model,
147149
if i < num_recompute_layers:
148150
checkpoint(block)
149151

152+
if version.parse(torch.__version__) >= version.parse("2.5.0"):
153+
for layer_cur, layer_next in zip(model.layers[:-1], model.layers[1:]):
154+
layer_cur.set_modules_to_forward_prefetch([layer_next])
155+
156+
150157
model.embed_tokens.apply(param_init_fn)
151158
model.norm.apply(param_init_fn)
152159
if hasattr(model, 'rotary_emb'):

‎xtuner/_lite/parallel/megatron/qwen2.py‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,7 @@
11
from functools import partial
2+
from packaging import version
23

4+
import torch
35
from torch import nn
46
from torch.distributed._tensor import Replicate, distribute_tensor
57
from torch.distributed.tensor.parallel import (ColwiseParallel,
@@ -155,6 +157,10 @@ def megatron_qwen2(model,
155157
if i < num_recompute_layers:
156158
checkpoint(block)
157159

160+
if version.parse(torch.__version__) >= version.parse("2.5.0"):
161+
for layer_cur, layer_next in zip(model.layers[:-1], model.layers[1:]):
162+
layer_cur.set_modules_to_forward_prefetch([layer_next])
163+
158164
model.embed_tokens.apply(param_init_fn)
159165
model.norm.apply(param_init_fn)
160166

0 commit comments

Comments
 (0)