@@ -307,13 +307,26 @@ def get_eos_mask(response_id: torch.Tensor, eos_token: int = 2, dtype=torch.int6
307307 return eos_mask
308308
309309
310- def get_pad_mask (response_id : torch .Tensor , pad_token : int = 0 , dtype = torch .int64 ):
310+ def get_pad_mask (response_id : torch .Tensor , pad_token : int = 0 , eos_token : int = 1 , dtype = torch .int64 ):
311311 """
312312 e.g. pad token=0
313313 response_id: [1, 2, 2, 42, 3, 5, 1, 0, 0]
314314 pad_mask: [1, 1, 1, 1, 1, 1, 1, 0, 0]
315+
316+ If eos_token == pad_token, the first pad token (which is the eos token) should be kept.
317+ e.g. pad_token=0, eos_token=0
318+ response_id: [1, 2, 2, 42, 3, 5, 0, 0, 0]
319+ pad_mask: [1, 1, 1, 1, 1, 1, 1, 0, 0] (first pad token/eos token is kept)
315320 """
316321 pad_mask = response_id .not_equal (pad_token ).to (dtype )
322+
323+ # eos_token == pad_token, 需要保留第一个pad token否则会误将eos token mask掉
324+ if eos_token == pad_token :
325+ pad_positions = response_id .eq (pad_token ).to (dtype )
326+ cumsum_pad = torch .cumsum (pad_positions , dim = - 1 )
327+ first_pad_token = (cumsum_pad == 1 ).to (dtype )
328+ pad_mask = pad_mask | first_pad_token
329+
317330 assert (
318331 not (pad_mask [:, 0 ] == 0 ).logical_and (pad_mask .sum (- 1 ) != 0 ).any ()
319332 ), f"response_id is not valid: { response_id } , pad_token is { pad_token } "
@@ -812,7 +825,7 @@ def postprocess_generate(
812825 attention_mask = (
813826 attention_mask .unsqueeze (1 ).repeat (1 , num_return_sequences , 1 ).view (output_batch_size , prompt_length )
814827 )
815- response_mask = get_pad_mask (response_id = response , pad_token = pad_token_id , dtype = attention_mask .dtype )
828+ response_mask = get_pad_mask (response_id = response , pad_token = pad_token_id , eos_token = eos_token_id , dtype = attention_mask .dtype )
816829 attention_mask = torch .cat ((attention_mask , response_mask ), dim = - 1 )
817830
818831 position_ids = prompts .batch ["position_ids" ]
0 commit comments