Download source/src/bgc_retrieval/model.py from rustambekurokov/bgc-setnet: direct link, hf CLI and curl.
- Browser
- Download file 22.7 kB
-
https://huggingface.co/rustambekurokov/bgc-setnet/resolve/main/source/src/bgc_retrieval/model.py
- Command line
-
hf download hf://rustambekurokov/bgc-setnet/source/src/bgc_retrieval/model.py
-
curl -L -o model.py https://huggingface.co/rustambekurokov/bgc-setnet/resolve/main/source/src/bgc_retrieval/model.py
22.7 kB
| """Set Transformer using only sequence embeddings and relative position.""" | |
| from __future__ import annotations | |
| from dataclasses import asdict, dataclass | |
| import torch | |
| from torch import nn | |
| from torch.nn import functional as F | |
| class ModelConfig: | |
| architecture: str = "setnet" | |
| esm_dimension: int = 1280 | |
| pfam_dimension: int = 64 | |
| pfam_vocab_size: int = 0 | |
| hidden_dimension: int = 256 | |
| output_dimension: int = 256 | |
| attention_heads: int = 8 | |
| inducing_points: int = 32 | |
| feedforward_dimension: int = 512 | |
| attention_blocks: int = 2 | |
| position_bins: int = 64 | |
| dropout: float = 0.1 | |
| def from_dict(cls, values: dict[str, object]) -> "ModelConfig": | |
| return cls(**{key: values[key] for key in asdict(cls()) if key in values}) | |
| class GeneProjection(nn.Module): | |
| def __init__(self, config: ModelConfig) -> None: | |
| super().__init__() | |
| self.projection = nn.Linear(config.esm_dimension, config.hidden_dimension) | |
| self.normalization = nn.LayerNorm(config.hidden_dimension) | |
| self.position = nn.Embedding(config.position_bins, config.hidden_dimension) | |
| self.dropout = nn.Dropout(config.dropout) | |
| self.position_bins = config.position_bins | |
| def forward(self, embeddings: torch.Tensor, positions: torch.Tensor) -> torch.Tensor: | |
| bins = (positions * (self.position_bins - 1)).long().clamp(0, self.position_bins - 1) | |
| projected = F.gelu(self.normalization(self.projection(embeddings))) | |
| return self.dropout(projected + self.position(bins)) | |
| class PfamGeneProjection(nn.Module): | |
| """Project ESM genes together with a BGC-level Pfam inventory embedding.""" | |
| def __init__(self, config: ModelConfig) -> None: | |
| super().__init__() | |
| if config.pfam_vocab_size < 3: | |
| raise ValueError("Pfam vocabulary must contain padding, unknown, and one domain token") | |
| self.pfam_embedding = nn.Embedding( | |
| config.pfam_vocab_size, config.pfam_dimension, padding_idx=0 | |
| ) | |
| self.projection = nn.Linear( | |
| config.esm_dimension + config.pfam_dimension, config.hidden_dimension | |
| ) | |
| self.normalization = nn.LayerNorm(config.hidden_dimension) | |
| self.position = nn.Embedding(config.position_bins, config.hidden_dimension) | |
| self.dropout = nn.Dropout(config.dropout) | |
| self.position_bins = config.position_bins | |
| def forward( | |
| self, | |
| embeddings: torch.Tensor, | |
| positions: torch.Tensor, | |
| pfam_tokens: torch.Tensor, | |
| ) -> torch.Tensor: | |
| bins = (positions * (self.position_bins - 1)).long().clamp(0, self.position_bins - 1) | |
| token_mask = pfam_tokens.ne(0).unsqueeze(-1) | |
| token_values = self.pfam_embedding(pfam_tokens).masked_fill(~token_mask, 0.0) | |
| counts = token_mask.sum(dim=1).clamp_min(1) | |
| pfam_summary = token_values.sum(dim=1) / counts | |
| pfam_summary = pfam_summary.unsqueeze(1).expand(-1, embeddings.shape[1], -1) | |
| merged = torch.cat((embeddings, pfam_summary), dim=-1) | |
| projected = F.gelu(self.normalization(self.projection(merged))) | |
| return self.dropout(projected + self.position(bins)) | |
| class InducedSelfAttentionBlock(nn.Module): | |
| def __init__(self, config: ModelConfig) -> None: | |
| super().__init__() | |
| dimension = config.hidden_dimension | |
| self.inducing = nn.Parameter(torch.empty(config.inducing_points, dimension)) | |
| nn.init.xavier_uniform_(self.inducing) | |
| self.inducing_norm = nn.LayerNorm(dimension) | |
| self.input_norm = nn.LayerNorm(dimension) | |
| self.summary_attention = nn.MultiheadAttention( | |
| dimension, config.attention_heads, config.dropout, batch_first=True | |
| ) | |
| self.output_attention = nn.MultiheadAttention( | |
| dimension, config.attention_heads, config.dropout, batch_first=True | |
| ) | |
| self.output_norm = nn.LayerNorm(dimension) | |
| self.feedforward = nn.Sequential( | |
| nn.Linear(dimension, config.feedforward_dimension), | |
| nn.GELU(), | |
| nn.Dropout(config.dropout), | |
| nn.Linear(config.feedforward_dimension, dimension), | |
| nn.Dropout(config.dropout), | |
| ) | |
| def forward(self, values: torch.Tensor, padding_mask: torch.Tensor | None) -> torch.Tensor: | |
| batch_size = values.shape[0] | |
| inducing = self.inducing.unsqueeze(0).expand(batch_size, -1, -1) | |
| normalized_values = self.input_norm(values) | |
| summary, _ = self.summary_attention( | |
| inducing, normalized_values, normalized_values, key_padding_mask=padding_mask | |
| ) | |
| normalized_summary = self.inducing_norm(summary) | |
| update, _ = self.output_attention( | |
| normalized_values, normalized_summary, normalized_summary | |
| ) | |
| output = values + update | |
| return output + self.feedforward(self.output_norm(output)) | |
| class PoolingByAttention(nn.Module): | |
| def __init__(self, config: ModelConfig) -> None: | |
| super().__init__() | |
| self.seed = nn.Parameter(torch.empty(1, config.hidden_dimension)) | |
| nn.init.xavier_uniform_(self.seed) | |
| self.normalization = nn.LayerNorm(config.hidden_dimension) | |
| self.attention = nn.MultiheadAttention( | |
| config.hidden_dimension, config.attention_heads, config.dropout, batch_first=True | |
| ) | |
| def forward(self, values: torch.Tensor, padding_mask: torch.Tensor | None) -> torch.Tensor: | |
| seed = self.seed.unsqueeze(0).expand(values.shape[0], -1, -1) | |
| output, _ = self.attention( | |
| seed, | |
| self.normalization(values), | |
| self.normalization(values), | |
| key_padding_mask=padding_mask, | |
| ) | |
| return output.squeeze(1) | |
| class LeakageFreeBGCSetNet(nn.Module): | |
| """Map a variable-length BGC gene set to a normalized embedding.""" | |
| input_names = ("gene_embeddings", "relative_positions", "padding_mask") | |
| def __init__(self, config: ModelConfig) -> None: | |
| super().__init__() | |
| self.config = config | |
| self.gene_projection = GeneProjection(config) | |
| self.blocks = nn.ModuleList( | |
| InducedSelfAttentionBlock(config) for _ in range(config.attention_blocks) | |
| ) | |
| self.pooling = PoolingByAttention(config) | |
| self.output = nn.Linear(config.hidden_dimension, config.output_dimension) | |
| def encode_genes( | |
| self, | |
| gene_embeddings: torch.Tensor, | |
| relative_positions: torch.Tensor, | |
| padding_mask: torch.Tensor | None = None, | |
| ) -> torch.Tensor: | |
| values = self.gene_projection(gene_embeddings, relative_positions) | |
| for block in self.blocks: | |
| values = block(values, padding_mask) | |
| return values | |
| def forward( | |
| self, | |
| gene_embeddings: torch.Tensor, | |
| relative_positions: torch.Tensor, | |
| padding_mask: torch.Tensor | None = None, | |
| ) -> torch.Tensor: | |
| genes = self.encode_genes(gene_embeddings, relative_positions, padding_mask) | |
| pooled = self.pooling(genes, padding_mask) | |
| return F.normalize(self.output(pooled), p=2, dim=-1) | |
| class MaskedGenePredictionHead(nn.Module): | |
| def __init__(self, config: ModelConfig) -> None: | |
| super().__init__() | |
| self.layers = nn.Sequential( | |
| nn.Linear(config.hidden_dimension, config.feedforward_dimension), | |
| nn.GELU(), | |
| nn.Linear(config.feedforward_dimension, config.esm_dimension), | |
| ) | |
| def forward(self, contextual_embeddings: torch.Tensor) -> torch.Tensor: | |
| return self.layers(contextual_embeddings) | |
| class GatedDeepSets(nn.Module): | |
| """Permutation-invariant learned pooling without gene-gene interactions.""" | |
| input_names = ("gene_embeddings", "relative_positions", "padding_mask") | |
| def __init__(self, config: ModelConfig) -> None: | |
| super().__init__() | |
| self.config = config | |
| self.gene_projection = GeneProjection(config) | |
| self.gate = nn.Sequential( | |
| nn.Linear(config.hidden_dimension, config.hidden_dimension // 2), | |
| nn.GELU(), | |
| nn.Linear(config.hidden_dimension // 2, 1), | |
| ) | |
| self.output = nn.Linear(config.hidden_dimension, config.output_dimension) | |
| def encode_genes(self, gene_embeddings, relative_positions, padding_mask=None): | |
| return self.gene_projection(gene_embeddings, relative_positions) | |
| def forward(self, gene_embeddings, relative_positions, padding_mask=None): | |
| genes = self.encode_genes(gene_embeddings, relative_positions, padding_mask) | |
| logits = self.gate(genes).squeeze(-1) | |
| if padding_mask is not None: | |
| logits = logits.masked_fill(padding_mask, -torch.finfo(logits.dtype).max) | |
| weights = torch.softmax(logits, dim=-1) | |
| if padding_mask is not None: | |
| weights = weights.masked_fill(padding_mask, 0.0) | |
| pooled = (genes * weights.unsqueeze(-1)).sum(dim=1) | |
| return F.normalize(self.output(pooled), p=2, dim=-1) | |
| class RotarySelfAttentionBlock(nn.Module): | |
| """Bidirectional self-attention with rotary position phases.""" | |
| def __init__(self, config: ModelConfig) -> None: | |
| super().__init__() | |
| d = config.hidden_dimension | |
| self.heads = config.attention_heads | |
| self.head_dim = d // self.heads | |
| if self.head_dim % 2: | |
| raise ValueError("Rotary attention requires an even per-head dimension") | |
| self.norm = nn.LayerNorm(d) | |
| self.qkv = nn.Linear(d, 3 * d) | |
| self.output = nn.Linear(d, d) | |
| self.dropout = nn.Dropout(config.dropout) | |
| self.ffn_norm = nn.LayerNorm(d) | |
| self.ffn = nn.Sequential( | |
| nn.Linear(d, config.feedforward_dimension), nn.GELU(), | |
| nn.Dropout(config.dropout), nn.Linear(config.feedforward_dimension, d), | |
| nn.Dropout(config.dropout), | |
| ) | |
| def _rotate(self, values: torch.Tensor, positions: torch.Tensor) -> torch.Tensor: | |
| half = self.head_dim // 2 | |
| inv = 1.0 / (10000.0 ** ( | |
| torch.arange(half, device=values.device, dtype=values.dtype) / half | |
| )) | |
| angles = positions[:, None, :, None] * 128.0 * inv[None, None, None, :] | |
| cos, sin = angles.cos(), angles.sin() | |
| first, second = values[..., :half], values[..., half:] | |
| return torch.cat((first * cos - second * sin, first * sin + second * cos), dim=-1) | |
| def forward(self, values: torch.Tensor, positions: torch.Tensor, padding_mask=None) -> torch.Tensor: | |
| normalized = self.norm(values) | |
| batch, length, dimension = normalized.shape | |
| qkv = self.qkv(normalized).view(batch, length, 3, self.heads, self.head_dim) | |
| q, k, v = qkv.unbind(dim=2) | |
| q, k = q.transpose(1, 2), k.transpose(1, 2) | |
| v = v.transpose(1, 2) | |
| q, k = self._rotate(q, positions), self._rotate(k, positions) | |
| scores = torch.matmul(q, k.transpose(-2, -1)) / (self.head_dim ** 0.5) | |
| if padding_mask is not None: | |
| scores = scores.masked_fill( | |
| padding_mask[:, None, None, :], -torch.finfo(scores.dtype).max | |
| ) | |
| attention = self.dropout(torch.softmax(scores, dim=-1)) | |
| contextual = torch.matmul(attention, v).transpose(1, 2).reshape(batch, length, dimension) | |
| output = values + self.output(contextual) | |
| output = output + self.ffn(self.ffn_norm(output)) | |
| if padding_mask is not None: | |
| output = output.masked_fill(padding_mask.unsqueeze(-1), 0.0) | |
| return output | |
| class RoPETransformer(nn.Module): | |
| """BGC-MAP-inspired bidirectional RoPE encoder with attentive pooling.""" | |
| input_names = ("gene_embeddings", "relative_positions", "padding_mask") | |
| def __init__(self, config: ModelConfig) -> None: | |
| super().__init__() | |
| self.config = config | |
| self.gene_projection = GeneProjection(config) | |
| self.blocks = nn.ModuleList( | |
| RotarySelfAttentionBlock(config) for _ in range(config.attention_blocks) | |
| ) | |
| self.pooling = PoolingByAttention(config) | |
| self.output = nn.Linear(config.hidden_dimension, config.output_dimension) | |
| def encode_genes(self, gene_embeddings, relative_positions, padding_mask=None): | |
| values = self.gene_projection(gene_embeddings, relative_positions) | |
| for block in self.blocks: | |
| values = block(values, relative_positions, padding_mask) | |
| return values | |
| def forward(self, gene_embeddings, relative_positions, padding_mask=None): | |
| genes = self.encode_genes(gene_embeddings, relative_positions, padding_mask) | |
| return F.normalize(self.output(self.pooling(genes, padding_mask)), p=2, dim=-1) | |
| class CrossAttentionPool(nn.Module): | |
| """Learned context queries cross-attend to a self-attended BGC sequence.""" | |
| input_names = ("gene_embeddings", "relative_positions", "padding_mask") | |
| def __init__(self, config: ModelConfig) -> None: | |
| super().__init__() | |
| self.config = config | |
| d = config.hidden_dimension | |
| self.gene_projection = GeneProjection(config) | |
| layer = nn.TransformerEncoderLayer( | |
| d, config.attention_heads, config.feedforward_dimension, config.dropout, | |
| batch_first=True, norm_first=True, activation="gelu", | |
| ) | |
| self.encoder = nn.TransformerEncoder(layer, config.attention_blocks) | |
| self.queries = nn.Parameter(torch.empty(2, d)) | |
| nn.init.xavier_uniform_(self.queries) | |
| self.cross = nn.MultiheadAttention(d, config.attention_heads, config.dropout, batch_first=True) | |
| self.output = nn.Linear(d, config.output_dimension) | |
| def encode_genes(self, gene_embeddings, relative_positions, padding_mask=None): | |
| values = self.gene_projection(gene_embeddings, relative_positions) | |
| return self.encoder(values, src_key_padding_mask=padding_mask) | |
| def forward(self, gene_embeddings, relative_positions, padding_mask=None): | |
| genes = self.encode_genes(gene_embeddings, relative_positions, padding_mask) | |
| queries = self.queries.unsqueeze(0).expand(genes.shape[0], -1, -1) | |
| attended, _ = self.cross(queries, genes, genes, key_padding_mask=padding_mask) | |
| return F.normalize(self.output(attended.mean(dim=1)), p=2, dim=-1) | |
| class LocalGlobalEncoder(nn.Module): | |
| """PST-inspired local adjacency message passing followed by global pooling.""" | |
| input_names = ("gene_embeddings", "relative_positions", "padding_mask") | |
| def __init__(self, config: ModelConfig) -> None: | |
| super().__init__() | |
| self.config = config | |
| d = config.hidden_dimension | |
| self.gene_projection = GeneProjection(config) | |
| self.local = nn.Conv1d(d, d, kernel_size=3, padding=1, groups=1) | |
| self.local_gate = nn.Sequential(nn.Linear(d, d), nn.Sigmoid()) | |
| self.global_block = InducedSelfAttentionBlock(config) | |
| self.pooling = PoolingByAttention(config) | |
| self.output = nn.Linear(d, config.output_dimension) | |
| def encode_genes(self, gene_embeddings, relative_positions, padding_mask=None): | |
| values = self.gene_projection(gene_embeddings, relative_positions) | |
| local = self.local(values.transpose(1, 2)).transpose(1, 2) | |
| values = values + local * self.local_gate(values) | |
| return self.global_block(values, padding_mask) | |
| def forward(self, gene_embeddings, relative_positions, padding_mask=None): | |
| genes = self.encode_genes(gene_embeddings, relative_positions, padding_mask) | |
| return F.normalize(self.output(self.pooling(genes, padding_mask)), p=2, dim=-1) | |
| class DilatedCNN(nn.Module): | |
| """BiGCARP-inspired residual dilated convolutional encoder.""" | |
| input_names = ("gene_embeddings", "relative_positions", "padding_mask") | |
| def __init__(self, config: ModelConfig) -> None: | |
| super().__init__() | |
| self.config = config | |
| d = config.hidden_dimension | |
| self.gene_projection = GeneProjection(config) | |
| self.blocks = nn.ModuleList( | |
| nn.Sequential( | |
| nn.Conv1d(d, d, 3, padding=dilation, dilation=dilation), | |
| nn.GELU(), nn.Dropout(config.dropout), nn.Conv1d(d, d, 1), | |
| ) for dilation in (1, 2, 4, 8, 16) | |
| ) | |
| self.pooling = PoolingByAttention(config) | |
| self.output = nn.Linear(d, config.output_dimension) | |
| def encode_genes(self, gene_embeddings, relative_positions, padding_mask=None): | |
| values = self.gene_projection(gene_embeddings, relative_positions) | |
| for block in self.blocks: | |
| updated = block(values.transpose(1, 2)).transpose(1, 2) | |
| values = values + updated | |
| if padding_mask is not None: | |
| values = values.masked_fill(padding_mask.unsqueeze(-1), 0.0) | |
| return values | |
| def forward(self, gene_embeddings, relative_positions, padding_mask=None): | |
| genes = self.encode_genes(gene_embeddings, relative_positions, padding_mask) | |
| return F.normalize(self.output(self.pooling(genes, padding_mask)), p=2, dim=-1) | |
| class HierarchicalLocalGlobal(nn.Module): | |
| """Multi-scale local convolutions followed by a small global Transformer.""" | |
| input_names = ("gene_embeddings", "relative_positions", "padding_mask") | |
| def __init__(self, config: ModelConfig) -> None: | |
| super().__init__() | |
| self.config = config | |
| d = config.hidden_dimension | |
| self.gene_projection = GeneProjection(config) | |
| self.local3 = nn.Conv1d(d, d // 2, 3, padding=1) | |
| self.local5 = nn.Conv1d(d, d // 2, 5, padding=2) | |
| self.merge = nn.Linear(d, d) | |
| layer = nn.TransformerEncoderLayer( | |
| d, config.attention_heads, config.feedforward_dimension, config.dropout, | |
| batch_first=True, norm_first=True, activation="gelu", | |
| ) | |
| self.global_encoder = nn.TransformerEncoder( | |
| layer, max(1, config.attention_blocks // 2) | |
| ) | |
| self.pooling = PoolingByAttention(config) | |
| self.output = nn.Linear(d, config.output_dimension) | |
| def encode_genes(self, gene_embeddings, relative_positions, padding_mask=None): | |
| values = self.gene_projection(gene_embeddings, relative_positions) | |
| transposed = values.transpose(1, 2) | |
| local = torch.cat((self.local3(transposed), self.local5(transposed)), dim=1).transpose(1, 2) | |
| values = values + self.merge(local) | |
| return self.global_encoder(values, src_key_padding_mask=padding_mask) | |
| def forward(self, gene_embeddings, relative_positions, padding_mask=None): | |
| genes = self.encode_genes(gene_embeddings, relative_positions, padding_mask) | |
| return F.normalize(self.output(self.pooling(genes, padding_mask)), p=2, dim=-1) | |
| class PfamAugmentedSetNet(nn.Module): | |
| """SetNet whose gene projection is conditioned on the BGC Pfam inventory.""" | |
| input_names = ("gene_embeddings", "relative_positions", "padding_mask", "pfam_tokens") | |
| uses_pfam = True | |
| def __init__(self, config: ModelConfig) -> None: | |
| super().__init__() | |
| self.config = config | |
| self.gene_projection = PfamGeneProjection(config) | |
| self.blocks = nn.ModuleList( | |
| InducedSelfAttentionBlock(config) for _ in range(config.attention_blocks) | |
| ) | |
| self.pooling = PoolingByAttention(config) | |
| self.output = nn.Linear(config.hidden_dimension, config.output_dimension) | |
| def encode_genes( | |
| self, | |
| gene_embeddings: torch.Tensor, | |
| relative_positions: torch.Tensor, | |
| padding_mask: torch.Tensor | None = None, | |
| pfam_tokens: torch.Tensor | None = None, | |
| ) -> torch.Tensor: | |
| if pfam_tokens is None: | |
| raise ValueError("Pfam-augmented SetNet requires pfam_tokens") | |
| values = self.gene_projection(gene_embeddings, relative_positions, pfam_tokens) | |
| for block in self.blocks: | |
| values = block(values, padding_mask) | |
| return values | |
| def forward( | |
| self, | |
| gene_embeddings: torch.Tensor, | |
| relative_positions: torch.Tensor, | |
| padding_mask: torch.Tensor | None = None, | |
| pfam_tokens: torch.Tensor | None = None, | |
| ) -> torch.Tensor: | |
| genes = self.encode_genes( | |
| gene_embeddings, relative_positions, padding_mask, pfam_tokens | |
| ) | |
| pooled = self.pooling(genes, padding_mask) | |
| return F.normalize(self.output(pooled), p=2, dim=-1) | |
| class WeightedPfamJaccard(nn.Module): | |
| """Learn nonnegative Pfam importance weights for differentiable set Jaccard.""" | |
| input_names = ("pfam_tokens",) | |
| uses_pfam = True | |
| def __init__(self, config: ModelConfig) -> None: | |
| super().__init__() | |
| if config.pfam_vocab_size < 3: | |
| raise ValueError("Pfam vocabulary must contain padding, unknown, and one domain token") | |
| initial = float(torch.log(torch.expm1(torch.tensor(1.0)))) | |
| self.raw_weights = nn.Parameter( | |
| torch.full((config.pfam_vocab_size,), initial, dtype=torch.float32) | |
| ) | |
| with torch.no_grad(): | |
| self.raw_weights[0] = -20.0 | |
| def domain_weights(self) -> torch.Tensor: | |
| positive = F.softplus(self.raw_weights) | |
| return torch.cat((positive[:1] * 0.0, positive[1:])) | |
| def pairwise_jaccard(self, pfam_tokens: torch.Tensor) -> torch.Tensor: | |
| if pfam_tokens.ndim != 2: | |
| raise ValueError("Pfam tokens must have shape [batch, domains]") | |
| batch_size = pfam_tokens.shape[0] | |
| vocabulary = self.raw_weights.shape[0] | |
| presence = torch.zeros( | |
| batch_size, vocabulary, device=pfam_tokens.device, dtype=torch.float32 | |
| ) | |
| presence.scatter_(1, pfam_tokens.clamp_min(0), 1.0) | |
| presence[:, 0] = 0.0 | |
| weighted = presence * self.domain_weights().to(pfam_tokens.device) | |
| totals = weighted.sum(dim=1) | |
| intersection = weighted @ presence.T | |
| union = totals[:, None] + totals[None, :] - intersection | |
| return intersection / union.clamp_min(1e-8) | |
| def forward(self, pfam_tokens: torch.Tensor) -> torch.Tensor: | |
| return self.domain_weights() | |
| ARCHITECTURES = { | |
| "setnet": LeakageFreeBGCSetNet, | |
| "weighted_pfam_jaccard": WeightedPfamJaccard, | |
| "pfam_setnet": PfamAugmentedSetNet, | |
| "gated_deepsets": GatedDeepSets, | |
| "rope_transformer": RoPETransformer, | |
| "cross_attention": CrossAttentionPool, | |
| "local_global": LocalGlobalEncoder, | |
| "dilated_cnn": DilatedCNN, | |
| "hierarchical": HierarchicalLocalGlobal, | |
| } | |
| def build_model(config: ModelConfig) -> nn.Module: | |
| try: | |
| model_class = ARCHITECTURES[config.architecture] | |
| except KeyError as exc: | |
| raise ValueError(f"Unknown architecture: {config.architecture}") from exc | |
| return model_class(config) | |