CoolFace
Modelpublic

ssaha3/memory-model

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes16downloads
configuration_custom_model.py30 linesDownload Raw Back to root
1# from transformers import PretrainedConfig2# class CustomModelConfig(PretrainedConfig):3#     model_type = "custom_model"4#     def __init__(self, parameter_1=128, parameter_2="default_value", **kwargs):5#         self.parameter_1 = parameter_16#         self.parameter_2 = parameter_27#         super().__init__(**kwargs)8# CustomModelConfig.register_for_auto_class()9# 10 11from transformers import AutoConfig, PretrainedConfig12 13class MemoryModelConfig(PretrainedConfig):14    model_type = "memory_model"15    16    def __init__(17        self,18        base_model_name: str = "google/gemma-2b",19        use_gradient_checkpointing: bool = False,20        deterministic = False,21        **kwargs22    ):23        super().__init__(**kwargs)24        self.base_model_name = base_model_name25        self.base_model_config = AutoConfig.from_pretrained(base_model_name)26        self.use_gradient_checkpointing = use_gradient_checkpointing27        self.deterministic = deterministic28        29        30MemoryModelConfig.register_for_auto_class()