对文档进行重排序并且分 bucket和block,分配到一个 GPU 上的文档被称为 bucket, bucket 和 bucket 之间的 chunk 数量尽可能均衡 Args: docs: 文档数据列表 Returns: List[List[Document]]: 每个 bucket 分配到的文档列表
(self, docs: List[str])
| 1476 | return bucket_docs |
| 1477 | |
| 1478 | def _sort_reference(self, docs: List[str]) -> Tuple[List[str], List[int], List[int], List[List[int]]]: |
| 1479 | """ |
| 1480 | 对文档进行重排序并且分 bucket和block,分配到一个 GPU 上的文档被称为 bucket, |
| 1481 | bucket 和 bucket 之间的 chunk 数量尽可能均衡 |
| 1482 | Args: |
| 1483 | docs: 文档数据列表 |
| 1484 | |
| 1485 | Returns: |
| 1486 | List[List[Document]]: 每个 bucket 分配到的文档列表 |
| 1487 | """ |
| 1488 | |
| 1489 | |
| 1490 | documents: List[Document] = [] |
| 1491 | kernel_sz = self.model_config.pooling_kernel_size |
| 1492 | for idx, doc in enumerate(docs): # idx 是 block 内部的 doc 索引 |
| 1493 | new_doc, doc_inputs = compose_input(doc, idx, self.tokenizer) |
| 1494 | length = len(doc_inputs["input_ids"]) |
| 1495 | num_chunks = (length + kernel_sz - 1) // kernel_sz |
| 1496 | # print(f"doc {idx} str {len(doc)} id length: {length}, num_chunks: {num_chunks}") |
| 1497 | |
| 1498 | documents.append(Document(doc=doc, doc_id=idx, num_chunks=num_chunks)) |
| 1499 | |
| 1500 | return MSAEngine.balanced_bucket_partition(documents, self.generate_config.world) |
| 1501 | |
| 1502 | def _load_memory_file(self): |
| 1503 | """加载memory文件""" |
no test coverage detected