ducanhdinh/jepa_proof_barlow_twins_replace
0
ducanhdinh/jepaproofbarlowtwinsreplace
BERT encoder pretrained from scratch với Barlow Twins + Lexical Substitution augmentation.
Augmentation strategy
Thay vì span masking, model này dùng lexical substitution để tạo view 2:
Quá trình tạo view 2 được thực hiện offline 1 lần trước khi train.
Kiến trúc Barlow Twins
View 1 ──► Encoder (θ) ──► Projector (θ) ──► z1 ──┐
├──► Cross-correlation C = Z1ᵀZ2 / N ──► Loss
View 2 ──► Encoder (θ) ──► Projector (θ) ──► z2 ──┘
Loss = Σ(C_ii - 1)² + λ · Σ_{i≠j} C_ij²
on-diagonal off-diagonal (redundancy reduction)Encoder và Projector dùng shared weights (không có target network).
Thông số huấn luyện
Cách dùng — BERT encoder (feature extraction)
from transformers import BertModel, BertTokenizerFast
import torch
tokenizer = BertTokenizerFast.from_pretrained("ducanhdinh/jepa_proof_barlow_twins_replace")
bert = BertModel.from_pretrained("ducanhdinh/jepa_proof_barlow_twins_replace/encoder")
encoded = tokenizer(
["Hello world!", "Barlow Twins with lexical substitution."],
return_tensors="pt",
padding=True,
truncation=True,
)
with torch.no_grad():
out = bert(**encoded)
cls_emb = out.last_hidden_state[:, 0, :] # [CLS] token → (B, 768)Cách dùng — Full model
import torch
from text_barlow_twins_replace import TextBarlowTwinsReplace, BarlowTwinsReplacePretrainConfig
cfg = BarlowTwinsReplacePretrainConfig()
model = TextBarlowTwinsReplace(cfg)
state = torch.load(
hf_hub_download("ducanhdinh/jepa_proof_barlow_twins_replace", "pytorch_model.bin"),
map_location="cpu",
)
model.load_state_dict(state)
model.eval()