MCPcopy Create free account
hub / github.com/InternLM/InternBootcamp / no_padding_2_padding

Function no_padding_2_padding

verl/verl/workers/utils/padding.py:75–106  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

73
74
75def 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

Callers 5

test_actor_engineFunction · 0.90
test_critic_engineFunction · 0.90
_compute_valuesMethod · 0.90
_compute_ref_log_probMethod · 0.90
_compute_old_log_probMethod · 0.90

Calls 2

pad_inputFunction · 0.90
valuesMethod · 0.80

Tested by 2

test_actor_engineFunction · 0.72
test_critic_engineFunction · 0.72