CoolFace
Modelpublic

Laeyoung/BTS-comments-generator

sourceHugging Faceupdated 5y agoView on Hugging Face
0likes25downloads
handler.py77 linesDownload Raw Back to root
1import torch2import gc3from ts.torch_handler.base_handler import BaseHandler4from transformers import GPT2LMHeadModel5 6import logging7 8logger = logging.getLogger(__name__)9 10 11class SampleTransformerModel(BaseHandler):12    def __init__(self):13        super(SampleTransformerModel, self).__init__()14        self.model = None15        self.device = None16        self.initialized = False17 18    def load_model(self, model_dir):19        self.model = GPT2LMHeadModel.from_pretrained(model_dir, return_dict=True)20        self.model.to(self.device)21 22    def initialize(self, ctx):23        # self.manifest = ctx.manifest24        properties = ctx.system_properties25        model_dir = properties.get("model_dir")26        self.device = torch.device("cuda:" + str(properties.get("gpu_id")) if torch.cuda.is_available() else "cpu")27 28        self.load_model(model_dir)29 30        self.model.eval()31        self.initialized = True32 33    def preprocess(self, requests):34        input_batch = {}35        for idx, data in enumerate(requests):36            input_ids = torch.tensor([data.get("body").get("text")]).to(self.device)37            input_batch["input_ids"] = input_ids38            input_batch["num_samples"] = data.get("body").get("num_samples")39            input_batch["length"] = data.get("body").get("length") + len(data.get("body").get("text"))40        del requests41        gc.collect()42        return input_batch43 44    def inference(self, input_batch):45        input_ids = input_batch["input_ids"]46        length = input_batch["length"]47 48        inference_output = self.model.generate(input_ids,49                                               bos_token_id=self.model.config.bos_token_id,50                                               eos_token_id=self.model.config.eos_token_id,51                                               pad_token_id=self.model.config.eos_token_id,52                                               do_sample=True,53                                               max_length=length,54                                               top_k=50,55                                               top_p=0.95,56                                               no_repeat_ngram_size=2,57                                               num_return_sequences=input_batch["num_samples"])58 59        if torch.cuda.is_available():60            torch.cuda.empty_cache()61        del input_batch62        gc.collect()63        return inference_output64 65    def postprocess(self, inference_output):66        output = inference_output.cpu().numpy().tolist()67        del inference_output68        gc.collect()69        return [output]70 71    def handle(self, data, context):72        # self.context = context73        data = self.preprocess(data)74        data = self.inference(data)75        data = self.postprocess(data)76        return data77