| from transformers import PretrainedConfig | |
| class BrainOCRVisionConfig(PretrainedConfig): | |
| model_type = "brain_ocr_vision" | |
| def __init__(self, image_size=560, patch_size=14, hidden_size=1280, | |
| intermediate_size=5120, num_hidden_layers=32, | |
| num_attention_heads=16, channels=3, layer_norm_eps=1e-6, | |
| attention_dropout=0.0, **kwargs): | |
| super().__init__(**kwargs) | |
| self.image_size = image_size | |
| self.patch_size = patch_size | |
| self.hidden_size = hidden_size | |
| self.intermediate_size = intermediate_size | |
| self.num_hidden_layers = num_hidden_layers | |
| self.num_attention_heads = num_attention_heads | |
| self.channels = channels | |
| self.layer_norm_eps = layer_norm_eps | |
| self.attention_dropout = attention_dropout | |
| class BrainOCRConfig(PretrainedConfig): | |
| model_type = "brain_ocr" | |
| sub_configs = {"vision_config": BrainOCRVisionConfig} | |
| def __init__(self, **kwargs): | |
| vision_config = kwargs.pop("vision_config", None) | |
| if isinstance(vision_config, dict): | |
| self.vision_config = BrainOCRVisionConfig(**vision_config) | |
| elif isinstance(vision_config, BrainOCRVisionConfig): | |
| self.vision_config = vision_config | |
| else: | |
| self.vision_config = BrainOCRVisionConfig() | |
| super().__init__(**kwargs) | |