Convert NestedTensor from no-padding to right padding format. Args: nested_tensor: NestedTensor with no-padding format data: TensorDict with - Tensor includes NestedTensors like "input_ids", "loss_mask", "position_ids" - NonTensorData includes "max_seq_len",
(nested_tensor: torch.Tensor, data: TensorDict)
| 73 | |
| 74 | |
| 75 | def no_padding_2_padding(nested_tensor: torch.Tensor, data: TensorDict) -> torch.Tensor: |
| 76 | """ |
| 77 | Convert NestedTensor from no-padding to right padding format. |
| 78 | |
| 79 | Args: |
| 80 | nested_tensor: NestedTensor with no-padding format |
| 81 | data: TensorDict with |
| 82 | - Tensor includes NestedTensors like "input_ids", "loss_mask", "position_ids" |
| 83 | - NonTensorData includes "max_seq_len", "max_response_len", "indices" |
| 84 | |
| 85 | Returns: |
| 86 | values: regular tensor right padded to max_response_len |
| 87 | """ |
| 88 | assert "indices" in data, "indices is required in left-right padding data" |
| 89 | assert "max_seq_len" in data, "max_seq_len is required in left-right padding data" |
| 90 | assert "max_response_len" in data, "max_response_len is required in left-right padding data" |
| 91 | |
| 92 | indices = tu.get_non_tensor_data(data=data, key="indices", default=None) |
| 93 | max_seq_len = tu.get_non_tensor_data(data=data, key="max_seq_len", default=2048) |
| 94 | max_response_len = tu.get_non_tensor_data(data=data, key="max_response_len", default=1024) |
| 95 | batch_size = nested_tensor.size(0) |
| 96 | |
| 97 | values = nested_tensor.values() |
| 98 | full_values = pad_input( |
| 99 | hidden_states=values.unsqueeze(-1), |
| 100 | indices=indices, |
| 101 | batch=batch_size, |
| 102 | seqlen=max_seq_len, |
| 103 | ) |
| 104 | values = full_values.squeeze(-1)[:, -max_response_len - 1 : -1] # (bsz, response_length) |
| 105 | |
| 106 | return values |