(forward_meta: ForwardMeta, fd_config)
| 64 | |
| 65 | |
| 66 | def split_batch_decoder_layers(forward_meta: ForwardMeta, fd_config): |
| 67 | split_num = 2 |
| 68 | res = [creat_empty_forward_meta(forward_meta), forward_meta] |
| 69 | res[0].tbo_microbatch_id = 0 |
| 70 | res[1].tbo_microbatch_id = 1 |
| 71 | total_token_num = forward_meta.ids_remove_padding.shape[0] |
| 72 | |
| 73 | if total_token_num < 1024: |
| 74 | return res |
| 75 | |
| 76 | chunk_token_num = (total_token_num + split_num - 1) // split_num |
| 77 | |
| 78 | split_sections = [] |
| 79 | for i in range(0, split_num): |
| 80 | start_token_id = i * chunk_token_num |
| 81 | end_token_id = start_token_id + chunk_token_num |
| 82 | end_token_id = min(total_token_num, end_token_id) |
| 83 | split_sections.append(end_token_id) |
| 84 | |
| 85 | # 由于多模的图片理解,需要将多模拟的token聚集在一起! |
| 86 | # 所以需要将split_sections[0]适当的偏移一下! |
| 87 | |
| 88 | special_tokens = [ |
| 89 | fd_config.model_config.image_patch_id, |
| 90 | ] |
| 91 | |
| 92 | ids_remove_padding_cpu = forward_meta.ids_remove_padding.numpy().tolist() |
| 93 | detect_pos = split_sections[0] |
| 94 | while ids_remove_padding_cpu[detect_pos] in special_tokens: |
| 95 | detect_pos += 1 |
| 96 | if detect_pos >= len(ids_remove_padding_cpu): |
| 97 | return res |
| 98 | split_sections[0] = detect_pos |
| 99 | |
| 100 | for i in range(0, split_num): |
| 101 | start_token_id = 0 if i == 0 else split_sections[i - 1] |
| 102 | end_token_id = split_sections[i] |
| 103 | |
| 104 | res[i] = ForwardMeta( |
| 105 | ids_remove_padding=None, |
| 106 | rotary_embs=forward_meta.rotary_embs, |
| 107 | attn_backend=forward_meta.attn_backend, |
| 108 | caches=forward_meta.caches, |
| 109 | ) |
| 110 | |
| 111 | # 我们需要处理的这一段token位于[start_bs, end_bs)里面! |
| 112 | start_bs = forward_meta.batch_id_per_token[start_token_id] |
| 113 | end_bs = forward_meta.batch_id_per_token[end_token_id - 1] |
| 114 | end_bs += 1 |
| 115 | |
| 116 | if len(forward_meta.rotary_embs.shape) == 6: |
| 117 | max_bs = forward_meta.rotary_embs.shape[0] |
| 118 | assert max_bs == forward_meta.block_tables.shape[0] |
| 119 | assert forward_meta.rotary_embs.shape[1:3] == [2, 1] |
| 120 | assert forward_meta.rotary_embs.shape[4] == 1 |
| 121 | res[i].rotary_embs = forward_meta.rotary_embs[start_bs:end_bs] |
| 122 | res[i].block_tables = forward_meta.block_tables[start_bs:end_bs] |
| 123 | res[i].ids_remove_padding = forward_meta.ids_remove_padding[start_token_id:end_token_id] |
nothing calls this directly
no test coverage detected