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)
| 91 | |
| 92 | |
| 93 | def 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, |
no test coverage detected