Team Ai
Apppublic

ishantvivek/codegen

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
model_generator.py48 linesDownload Raw Back to services
1import torch2from transformers import AutoModelForCausalLM, AutoTokenizer3from utils.config import Config4from utils.logger import Logger5 6logger = Logger.get_logger(__name__)7 8 9class ModelGenerator:10    """11    Singleton class responsible for generating text using a specified language model.12 13    This class initializes a language model and tokenizer, and provides methods 14    to generate text and extract code blocks from generated text.15 16    Attributes:17        device (torch.device): Device to run the model on (CPU or GPU).18        model (AutoModelForCausalLM): Language model for text generation.19        tokenizer (AutoTokenizer): Tokenizer corresponding to the language model.20 21    Methods:22        acceptTextGenerator(self, visitor, *args, **kwargs):23            Accepts a visitor to generates text based on the input provided with the model generator.24        acceptExtractCodeBlock(self, visitor, *args, **kwargs):25            Accepts a visitor to extract code blocks from the output text.26    """27    _instance = None28    _format_data_time = "%Y-%m-%d %H:%M:%S"29 30    def __new__(cls, model_name=Config.read('app', 'model')):31        if cls._instance is None:32            cls._instance = super(ModelGenerator, cls).__new__(cls)33            cls._instance._initialize(model_name)34        return cls._instance35 36    def _initialize(self, model_name):37        self.device = torch.device(38            "cuda" if torch.cuda.is_available() else "cpu")39        self.model = AutoModelForCausalLM.from_pretrained(40            model_name).to(self.device)41        self.tokenizer = AutoTokenizer.from_pretrained(model_name)42 43    def acceptTextGenerator(self, visitor, *args, **kwargs):44        return visitor.visit(self, *args, **kwargs)45 46    def acceptExtractCodeBlock(self, visitor, *args, **kwargs):47        return visitor.visit(self, *args, **kwargs)48