| 544 | return self.poolers_by_task[task].get_pooling_updates(task) |
| 545 | |
| 546 | def forward( |
| 547 | self, |
| 548 | hidden_states: Union[paddle.Tensor, list[paddle.Tensor]], |
| 549 | pooling_metadata: PoolingMetadata, |
| 550 | ) -> PoolerOutput: |
| 551 | poolers_by_task = self.poolers_by_task |
| 552 | |
| 553 | outputs = list[PoolingSequenceGroupOutput]() |
| 554 | offset = 0 |
| 555 | for task, group in groupby(get_tasks(pooling_metadata)): |
| 556 | if not (pooler := poolers_by_task.get(task)): |
| 557 | raise ValueError(f"Unsupported task: {task} " f"Supported tasks: {self.get_supported_tasks()}") |
| 558 | |
| 559 | num_items = len(list(group)) |
| 560 | group_output: PoolerOutput = pooler( |
| 561 | hidden_states, |
| 562 | pooling_metadata[offset : offset + num_items], |
| 563 | ) |
| 564 | outputs.extend(group_output) |
| 565 | offset += num_items |
| 566 | |
| 567 | return PoolerOutput(outputs) |