(
processing_class,
prompt_inputs,
completions,
)
| 46 | } |
| 47 | |
| 48 | def _create_inputs( |
| 49 | processing_class, |
| 50 | prompt_inputs, |
| 51 | completions, |
| 52 | ): |
| 53 | # now handle completion_ids and completion_mask |
| 54 | pad_token_id = getattr(processing_class,"pad_token_id", getattr(processing_class.tokenizer,"pad_token_id",None)) |
| 55 | if pad_token_id is None: |
| 56 | pad_token_id = 0 |
| 57 | completion_ids = torch.full((len(prompt_inputs["input_ids"]),max(map(len,completions))), pad_token_id , dtype=prompt_inputs["input_ids"].dtype,device=prompt_inputs["input_ids"].device) |
| 58 | for idx,completion in enumerate(completions): |
| 59 | completion_ids[idx,:len(completion)] = completion |
| 60 | |
| 61 | # Mask everything after the first EOS token |
| 62 | im_eos = completion_ids == processing_class.tokenizer.convert_tokens_to_ids('<|im_end|>') |
| 63 | s_eos = completion_ids == processing_class.tokenizer.convert_tokens_to_ids('</s>') |
| 64 | is_eos = im_eos | s_eos |
| 65 | |
| 66 | eos_idx = torch.full((is_eos.size(0),), is_eos.size(1), dtype=torch.long,device=completion_ids.device) |
| 67 | eos_idx[is_eos.any(dim=1)] = is_eos.int().argmax(dim=1)[is_eos.any(dim=1)] |
| 68 | sequence_indices = torch.arange(is_eos.size(1)).expand(is_eos.size(0), -1).to(device=eos_idx.device) |
| 69 | completion_mask = (sequence_indices <= eos_idx.unsqueeze(1)).int() |
| 70 | |
| 71 | |
| 72 | prompt_inputs["input_ids"] = torch.cat([prompt_inputs["input_ids"],completion_ids],dim=-1).to(dtype=torch.int64) |
| 73 | prompt_inputs["attention_mask"] = torch.cat([prompt_inputs["attention_mask"], completion_mask], dim=1) # (B, P+C) |
| 74 | |
| 75 | return prompt_inputs,completion_mask |
| 76 | |
| 77 | def _process_inputs( |
| 78 | inputs, |
no test coverage detected