Text Classification
Transformers
JAX
pcmlp
feature-extraction
predictive-coding
local-loss
flax
tiny-model
custom-architecture
custom_code
Eval Results (legacy)
Instructions to use zeechimp/pc-mlp-tiny with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use zeechimp/pc-mlp-tiny with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-classification", model="zeechimp/pc-mlp-tiny", trust_remote_code=True)# pip install -U transformers accelerate # Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("zeechimp/pc-mlp-tiny", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download modeling_pcmlp.py from zeechimp/pc-mlp-tiny: direct link, hf CLI and curl.
- Browser
- Download file 1.69 kB
-
https://huggingface.co/zeechimp/pc-mlp-tiny/resolve/main/modeling_pcmlp.py
- Command line
-
hf download hf://zeechimp/pc-mlp-tiny/modeling_pcmlp.py
-
curl -L -o modeling_pcmlp.py https://huggingface.co/zeechimp/pc-mlp-tiny/resolve/main/modeling_pcmlp.py
1.69 kB
| """PC-MLP model for HuggingFace Hub with auto_map.""" | |
| import math | |
| import torch | |
| import torch.nn as nn | |
| from transformers import PreTrainedModel | |
| from .configuration_pcmlp import PCMLPConfig | |
| class PCMLP(PreTrainedModel): | |
| config_class = PCMLPConfig | |
| def __init__(self, config): | |
| super().__init__(config) | |
| d = config.vocab_size | |
| h = config.hidden_dim | |
| scale = 1.0 / math.sqrt(h) | |
| self.tok = nn.Embedding(config.vocab_size, d) | |
| self.pos = nn.Embedding(config.seq_len, d) | |
| self.W1 = nn.Linear(d, h, bias=False) | |
| self.W2 = nn.Linear(h, h, bias=False) | |
| self.head = nn.Linear(h, config.num_classes) | |
| # match JAX init scale | |
| nn.init.normal_(self.tok.weight, std=0.05) | |
| nn.init.normal_(self.pos.weight, std=0.05) | |
| nn.init.normal_(self.W1.weight, std=scale) | |
| nn.init.normal_(self.W2.weight, std=scale) | |
| nn.init.normal_(self.head.weight, std=0.02) | |
| def forward(self, input_ids, labels=None): | |
| B, T = input_ids.shape | |
| pos_ids = torch.arange(T, device=input_ids.device) | |
| h = self.tok(input_ids) + self.pos(pos_ids) | |
| h = h.reshape(B, -1)[:, :self.W1.in_features] | |
| h0 = h | |
| h1 = torch.nn.functional.gelu(self.W1(h0)) | |
| h2 = torch.nn.functional.gelu(self.W2(h1)) | |
| logits = self.head(h2) | |
| loss = None | |
| if labels is not None: | |
| ce = torch.nn.functional.cross_entropy(logits, labels) | |
| local = ((h1.mean(1) - h0.mean(1)) ** 2).mean() + \ | |
| ((h2.mean(1) - h1.mean(1)) ** 2).mean() | |
| loss = ce + self.config.local_loss_weight * local | |
| return {"loss": loss, "logits": logits} |