MCPcopy Create free account
hub / github.com/Francis-Rings/FlashPortrait / split_tensor_with_padding

Function split_tensor_with_padding

wan/utils/utils.py:494–525  ·  view source on GitHub ↗
(input_tensor, pos_idx_range, expand_length=0)

Source from the content-addressed store, hash-verified

492
493
494def split_tensor_with_padding(input_tensor, pos_idx_range, expand_length=0):
495 pos_idx_range = [
496 [idx[0] - expand_length, idx[1] + expand_length] for idx in pos_idx_range
497 ]
498 sub_sequences = []
499 seq_len = input_tensor.size(1)
500 max_valid_idx = seq_len - 1
501 k_lens_list = []
502 for start, end in pos_idx_range:
503 pad_front = max(-start, 0)
504 pad_back = max(end - max_valid_idx, 0)
505
506 valid_start = max(start, 0)
507 valid_end = min(end, max_valid_idx)
508
509 if valid_start <= valid_end:
510 valid_part = input_tensor[:, valid_start: valid_end + 1, :]
511 else:
512 valid_part = input_tensor.new_zeros((1, 0, input_tensor.size(2)))
513
514 padded_subseq = F.pad(
515 valid_part,
516 (0, 0, 0, pad_back + pad_front, 0, 0),
517 mode="constant",
518 value=0,
519 )
520 k_lens_list.append(padded_subseq.size(-2) - pad_back - pad_front)
521
522 sub_sequences.append(padded_subseq)
523 return torch.stack(sub_sequences, dim=1), torch.tensor(
524 k_lens_list, dtype=torch.long
525 )

Callers 2

mainFunction · 0.90
__call__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected