MCPcopy Create free account
hub / github.com/zai-org/CodeGeeX / get_batch_pipe

Function get_batch_pipe

codegeex/megatron/tools/pretrain_codegeex.py:123–149  ·  view source on GitHub ↗

Modification of `get_batch` to work on `next(data_iterator)` instead of `data_iterator`

(data)

Source from the content-addressed store, hash-verified

121
122
123def get_batch_pipe(data):
124 """Modification of `get_batch` to work on `next(data_iterator)` instead of `data_iterator`"""
125 args = get_args()
126 tokenizer = get_tokenizer()
127
128 # Items and their type.
129 keys = ["input_ids"]
130 datatype = torch.int64
131
132 # Broadcast data.
133 data_b = mpu.broadcast_data(keys, data, datatype)
134
135 # Unpack.
136 tokens_ = data_b["input_ids"].long()
137 labels = tokens_[:, 1:].contiguous()
138 tokens = tokens_[:, :-1].contiguous()
139
140 # Get the masks and postition ids.
141 attention_mask, loss_mask, position_ids = get_ltor_masks_and_position_ids(
142 tokens,
143 tokenizer.eod,
144 args.reset_position_ids,
145 args.reset_attention_mask,
146 args.eod_mask_loss,
147 )
148
149 return (tokens, position_ids, attention_mask), (labels, loss_mask)
150
151
152def loss_func(loss_mask, output_tensor):

Callers

nothing calls this directly

Calls 3

get_argsFunction · 0.90
get_tokenizerFunction · 0.90

Tested by

no test coverage detected