| 466 | return torch.float16 |
| 467 | |
| 468 | def split_audio_adapter_sequence(adapter_proj_length, num_frames=80): |
| 469 | tokens_pre_frame = adapter_proj_length / num_frames |
| 470 | tokens_pre_latents_frame = tokens_pre_frame * 4 |
| 471 | half_tokens_pre_latents_frame = tokens_pre_latents_frame / 2 |
| 472 | pos_idx = [] |
| 473 | for i in range(int((num_frames - 1) / 4) + 1): |
| 474 | if i == 0: |
| 475 | pos_idx.append(0) |
| 476 | else: |
| 477 | begin_token_id = tokens_pre_frame * ((i - 1) * 4 + 1) |
| 478 | end_token_id = tokens_pre_frame * (i * 4 + 1) |
| 479 | pos_idx.append(int((sum([begin_token_id, end_token_id]) / 2)) - 1) |
| 480 | pos_idx_range = [ |
| 481 | [ |
| 482 | idx - int(half_tokens_pre_latents_frame), |
| 483 | idx + int(half_tokens_pre_latents_frame), |
| 484 | ] |
| 485 | for idx in pos_idx |
| 486 | ] |
| 487 | pos_idx_range[0] = [ |
| 488 | -(int(half_tokens_pre_latents_frame) * 2 - pos_idx_range[1][0]), |
| 489 | pos_idx_range[1][0], |
| 490 | ] |
| 491 | return pos_idx_range |
| 492 | |
| 493 | |
| 494 | def split_tensor_with_padding(input_tensor, pos_idx_range, expand_length=0): |