MCPcopy Create free account
hub / github.com/DataArcTech/DataArc-SynData-Toolkit / postprocess_batch_func

Function postprocess_batch_func

verl/workers/engine/utils.py:93–154  ·  view source on GitHub ↗

postprocess the output of a forward_backward_batch. output_lst is a list of dict containing outputs for each micro-batch reorder entropy and outputs. Return None for other pp ranks only on last rank. It should be on every tp rank each losses_reduced contains 1. model_output, 2. loss

(output_lst, indices, data: TensorDict)

Source from the content-addressed store, hash-verified

91
92
93def postprocess_batch_func(output_lst, indices, data: TensorDict):
94 """postprocess the output of a forward_backward_batch.
95 output_lst is a list of dict containing outputs for each micro-batch
96 reorder entropy and outputs. Return None for other pp ranks
97 only on last rank. It should be on every tp rank
98
99 each losses_reduced contains 1. model_output, 2. loss, 3. metrics.
100 """
101
102 use_dynamic_bsz = tu.get_non_tensor_data(data=data, key="use_dynamic_bsz", default=True)
103 pad_mode = tu.get_non_tensor_data(data=data, key="pad_mode", default=DatasetPadMode.NO_PADDING)
104 assert pad_mode == DatasetPadMode.NO_PADDING, "postprocess_batch_func only support NO_PADDING pad_mode"
105
106 # losses_reduced is a list of dict containing outputs for each micro-batch
107 # reorder entropy and outputs. Return None for other pp ranks
108 # only on last rank. It should be on every tp rank
109
110 # losses_reduced contains 1. model_output, 2. loss, 3. metrics.
111 # We perform reverse
112
113 model_output = {}
114 losses = []
115 aggregated_metrics = {}
116
117 # model output
118 for o in output_lst:
119 if "model_output" in o:
120 for key, val in o["model_output"].items():
121 if key not in model_output:
122 model_output[key] = []
123 model_output[key].append(val)
124
125 # concat results from micro batches
126 for key, val in model_output.items():
127 if pad_mode == DatasetPadMode.NO_PADDING:
128 tensors = [tensor for nt in model_output[key] for tensor in nt.unbind()]
129 model_output[key] = torch.nested.as_nested_tensor(tensors, layout=torch.jagged)
130 else:
131 raise NotImplementedError(f"pad_mode {pad_mode} not implemented")
132
133 # reverse with dynamic bsz
134 if use_dynamic_bsz:
135 model_output[key] = restore_dynamic_batch(model_output[key], indices)
136
137 # loss
138 for o in output_lst:
139 if "loss" in o:
140 losses.append(o["loss"])
141
142 # metrics
143 for o in output_lst:
144 if "metrics" in o:
145 metrics = o["metrics"]
146 append_to_dict(aggregated_metrics, metrics)
147
148 output = {
149 "model_output": model_output,
150 "loss": losses,

Callers 2

Calls 2

restore_dynamic_batchFunction · 0.90
append_to_dictFunction · 0.90

Tested by

no test coverage detected