File size: 1,406 Bytes
95456ed
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
import torch

from ScanDL2 import ScanDL2


class EndpointHandler:
    def __init__(self, path: str = ""):

        self.models = {
            "sentence": ScanDL2(
                text_type="sentence",
                bsz=2,
                save=None,
                filename=None,
            ),
            "paragraph": ScanDL2(
                text_type="paragraph",
                bsz=2,
                save=None,
                filename=None,
            ),
        }

        for m in self.models.values():
            # m.to(self.device)
            m.eval()

    def __call__(self, data):

        inputs = data.get("inputs", data)

        parameters = data.get("parameters", {})

        text_type = parameters.get("text_type", "sentence")
        model = self.models[text_type]
        bsz = parameters.get("bsz", 2)

        if model.scandl_module.args.batch_size != bsz:
            model.scandl_module.args.batch_size = bsz
            model.fixdur_module.bsz = bsz
            model.fixdur_module.args["bsz"] = bsz

        if isinstance(inputs, str):
            texts = [inputs]
        elif isinstance(inputs, list):
            texts = inputs
        else:
            raise ValueError("'inputs' must be a string or list of strings.")

        with torch.no_grad():
            output = model(texts=texts)

        return output