Views
No views yet
meta-llama/Llama-3.2-1B.1import torch
2import torch.nn.functional as F
3from transformers import AutoTokenizer, AutoModelForCausalLM
4from peft import PeftModel
5
6model_id = "yixuantt/DREAM-1B"
7base_id = "meta-llama/Llama-3.2-1B"
8
9tokenizer = AutoTokenizer.from_pretrained(model_id)
10base = AutoModelForCausalLM.from_pretrained(
11 base_id,
12 torch_dtype=torch.bfloat16,
13 device_map="auto",
14)
15model = PeftModel.from_pretrained(base, model_id)
16model.eval()
17
18@torch.no_grad()
19def encode(texts, max_length=512):
20 inputs = tokenizer(
21 texts,
22 padding=True,
23 truncation=True,
24 max_length=max_length,
25 return_tensors="pt",
26 ).to(model.device)
27 outputs = model(**inputs, output_hidden_states=True, use_cache=False)
28 hidden = outputs.hidden_states[-1]
29 # Pool the last non-padding token. This works for both left and right padding.
30 last_idx = inputs["attention_mask"].size(1) - 1 - inputs["attention_mask"].flip(dims=[1]).argmax(dim=1)
31 emb = hidden[torch.arange(hidden.size(0), device=hidden.device), last_idx]
32 return F.normalize(emb.float(), p=2, dim=-1)
33
34queries = encode(["What is DREAM?"])
35docs = encode(["DREAM trains dense retrievers with autoregressive supervision."])
36scores = queries @ docs.T
37print(scores)