File size: 1,254 Bytes
70af406
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
# Copyright (c) Meta Platforms, Inc. and affiliates. All Rights Reserved\n

from abc import ABCMeta, abstractmethod
from typing import List

import torch


class Ranker(torch.nn.Module, metaclass=ABCMeta):
    @abstractmethod
    def forward(self, audio: list[torch.Tensor], **kwargs) -> torch.Tensor:
        """

        Args:

            audio: (list[torch.Tensor]) where each element in the list corresponds to

                the candidates for the i'th generation (num_candidates, num_frames)

        Returns:

            (torch.Tensor) of shape (batch_size, num_candidates) correspoding to the ranking scores

        """
        pass


class EnsembleRanker(Ranker):
    def __init__(self, rankers: List[Ranker], weights: List[float]):
        super().__init__()
        assert len(rankers) == len(weights)
        self.rankers = torch.nn.ModuleList(rankers)
        self.weights = weights

    def forward(self, **kwargs) -> torch.Tensor:
        result = None
        for weight, ranker in zip(self.weights, self.rankers, strict=False):
            if result is None:
                result = weight * ranker(**kwargs)
            else:
                result += weight * ranker(**kwargs)
        return result