Initial release: code for WCL2026-1544 (context-aware embedding masking via DRL)
This commit is contained in:
@@ -0,0 +1,128 @@
|
||||
# =========================================================
|
||||
# extract_bert_embeddings.py
|
||||
# Pre-compute mean-pooled BERT sentence embeddings for AG News
|
||||
# and save them as a .pt tensor.
|
||||
#
|
||||
# The output file is consumed by drl_mask_policy.py via
|
||||
# --embed-file path/to/bert_agnews_8000.pt
|
||||
#
|
||||
# The consumer script applies dataset-level mean centering and
|
||||
# L2 normalization, so the saved tensor here contains the RAW
|
||||
# mean-pooled embeddings (no centering, no normalization).
|
||||
#
|
||||
# Usage:
|
||||
# python extract_bert_embeddings.py \
|
||||
# --out bert_agnews_8000.pt \
|
||||
# --max-sentences 8000
|
||||
# =========================================================
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import random
|
||||
|
||||
import torch
|
||||
from transformers import BertModel, BertTokenizer
|
||||
|
||||
|
||||
def load_agnews_sentences(max_sentences, min_len=5, max_len=30):
|
||||
"""Load first-sentence headlines from AG News. Falls back to
|
||||
synthetic templated text if the dataset cannot be fetched."""
|
||||
sentences = []
|
||||
try:
|
||||
from datasets import load_dataset
|
||||
ds = load_dataset("ag_news", split="train")
|
||||
for example in ds:
|
||||
first = example["text"].split(".")[0].strip()
|
||||
words = first.split()
|
||||
if min_len <= len(words) <= max_len:
|
||||
sentences.append(first)
|
||||
if len(sentences) >= max_sentences:
|
||||
break
|
||||
except Exception as e:
|
||||
print(f"[WARN] AG News unavailable ({e}). Using synthetic.")
|
||||
|
||||
if len(sentences) < max_sentences:
|
||||
print(f"[INFO] Padding with synthetic sentences "
|
||||
f"(have {len(sentences)}, need {max_sentences}).")
|
||||
templates = [
|
||||
"The {} {} the {} in the {}.",
|
||||
"A {} {} quickly {} the {}.",
|
||||
"Several {} {} near the {} {}.",
|
||||
]
|
||||
words = ["system", "signal", "network", "channel", "user",
|
||||
"device", "antenna", "receiver", "transmitter",
|
||||
"processes", "transmits", "receives", "encodes",
|
||||
"wireless", "digital", "robust", "adaptive"]
|
||||
while len(sentences) < max_sentences:
|
||||
t = random.choice(templates)
|
||||
n = t.count("{}")
|
||||
sentences.append(t.format(*random.choices(words, k=n)))
|
||||
|
||||
random.shuffle(sentences)
|
||||
return sentences[:max_sentences]
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def encode_batch(model, tokenizer, texts, device, max_length=64):
|
||||
"""Mean-pool token embeddings over non-padding positions."""
|
||||
inputs = tokenizer(texts, padding=True, truncation=True,
|
||||
max_length=max_length,
|
||||
return_tensors="pt").to(device)
|
||||
out = model(**inputs)
|
||||
hidden = out.last_hidden_state # (B, T, d)
|
||||
mask = inputs["attention_mask"].unsqueeze(-1).float()
|
||||
summed = (hidden * mask).sum(dim=1) # (B, d)
|
||||
count = mask.sum(dim=1).clamp(min=1.0) # (B, 1)
|
||||
return (summed / count).cpu() # (B, d)
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--model", type=str, default="bert-base-uncased")
|
||||
parser.add_argument("--out", type=str, default="bert_agnews_8000.pt")
|
||||
parser.add_argument("--max-sentences", type=int, default=8000)
|
||||
parser.add_argument("--batch-size", type=int, default=64)
|
||||
parser.add_argument("--cuda", action="store_true")
|
||||
parser.add_argument("--seed", type=int, default=0)
|
||||
args = parser.parse_args()
|
||||
|
||||
random.seed(args.seed)
|
||||
torch.manual_seed(args.seed)
|
||||
|
||||
if args.cuda and torch.cuda.is_available():
|
||||
device = torch.device("cuda")
|
||||
elif torch.backends.mps.is_available():
|
||||
device = torch.device("mps")
|
||||
else:
|
||||
device = torch.device("cpu")
|
||||
print(f"[INFO] Device: {device}")
|
||||
|
||||
print(f"[INFO] Loading {args.model} ...")
|
||||
tokenizer = BertTokenizer.from_pretrained(args.model)
|
||||
model = BertModel.from_pretrained(args.model).to(device)
|
||||
model.eval()
|
||||
d_bert = model.config.hidden_size
|
||||
print(f"[INFO] BERT hidden dim = {d_bert}")
|
||||
|
||||
print(f"[INFO] Loading {args.max_sentences} AG News sentences ...")
|
||||
sents = load_agnews_sentences(args.max_sentences)
|
||||
print(f"[INFO] Got {len(sents)} sentences.")
|
||||
|
||||
all_emb = []
|
||||
for i in range(0, len(sents), args.batch_size):
|
||||
batch = sents[i:i + args.batch_size]
|
||||
emb = encode_batch(model, tokenizer, batch, device)
|
||||
all_emb.append(emb)
|
||||
if (i // args.batch_size) % 20 == 0:
|
||||
print(f"[INFO] Processed {i + len(batch)}/{len(sents)}")
|
||||
emb = torch.cat(all_emb, dim=0)
|
||||
print(f"[INFO] Final tensor shape: {tuple(emb.shape)}")
|
||||
|
||||
os.makedirs(os.path.dirname(args.out) or ".", exist_ok=True)
|
||||
torch.save(emb, args.out)
|
||||
print(f"[OK] Saved raw mean-pooled embeddings to {args.out}")
|
||||
print(f" Feed into drl_mask_policy.py via --embed-file {args.out}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user