MCPcopy Create free account
hub / github.com/OpenBitSys/BitDistiller / padding

Function padding

test/gsm8k/test.py:139–151  ·  view source on GitHub ↗
(inputs, padding_token, cutoff = None)

Source from the content-addressed store, hash-verified

137 return dict(input_ids=self.input_ids[i], labels=self.labels[i], id=i)
138
139def padding(inputs, padding_token, cutoff = None):
140 num_elems = len(inputs)
141 if cutoff is None:
142 cutoff = max([len(item) for item in inputs])
143 else:
144 cutoff = min(max([len(item) for item in inputs]), cutoff)
145
146 tokens = torch.ones(num_elems, cutoff).long().to(inputs[0].device) * padding_token
147 for i in range(num_elems):
148 toks = inputs[i]
149 length = min(cutoff, len(toks))
150 tokens[i, -length:] = toks[-length:]
151 return tokens
152
153def sequence_gather(s, world_size, pad_tok_id):
154 local_size = torch.tensor(s.size(), device=s.device)

Callers 1

__call__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected