MCPcopy Create free account
hub / github.com/OpenBMB/AgentCPM-GUI / _create_inputs

Function _create_inputs

rft/trainer/utils/process.py:48–75  ·  view source on GitHub ↗
(
    processing_class,
    prompt_inputs,
    completions,
)

Source from the content-addressed store, hash-verified

46 }
47
48def _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
77def _process_inputs(
78 inputs,

Callers 2

sample_stepMethod · 0.85
_process_inputsFunction · 0.85

Calls 1

convert_tokens_to_idsMethod · 0.80

Tested by

no test coverage detected