1import torch
2import torch.nn.functional as F
3from transformers import AutoModelForCausalLM, AutoTokenizer
4
5
6def laser_encode(model, tokenizer, texts, max_length=512, num_thinking_steps=3):
7 """Encode texts using LaSER's latent thinking mechanism."""
8 device = next(model.parameters()).device
9 batch = tokenizer(texts, padding=True, truncation=True, max_length=max_length, return_tensors="pt").to(device)
10 input_ids, attention_mask = batch["input_ids"], batch["attention_mask"]
11
12 batch_size = input_ids.size(0)
13 thinking_slots = num_thinking_steps - 1
14 eos_id = tokenizer.eos_token_id
15
16 if thinking_slots > 0:
17 eos_padding = torch.full((batch_size, thinking_slots), eos_id, dtype=input_ids.dtype, device=device)
18 mask_padding = torch.ones((batch_size, thinking_slots), dtype=attention_mask.dtype, device=device)
19 input_ids = torch.cat([input_ids, eos_padding], dim=1)
20 attention_mask = torch.cat([attention_mask, mask_padding], dim=1)
21
22 input_embeds = model.get_input_embeddings()(input_ids)
23 embedding_table = model.get_input_embeddings().weight
24 base_seq_len = input_embeds.size(1) - thinking_slots
25
26 past_key_values = None
27 hidden_steps = []
28
29 for step_idx in range(thinking_slots):
30 pos = base_seq_len + step_idx
31 step_embeds = input_embeds[:, :pos, :] if past_key_values is None else input_embeds[:, pos-1:pos, :]
32 step_mask = attention_mask[:, :pos]
33
34 outputs = model(inputs_embeds=step_embeds, attention_mask=step_mask,
35 output_hidden_states=True, past_key_values=past_key_values,
36 use_cache=True, return_dict=True)
37 hidden_steps.append(outputs.hidden_states[-1][:, -1, :])
38 token_probs = torch.softmax(outputs.logits[:, -1, :], dim=-1)
39 new_embed = token_probs @ embedding_table
40 past_key_values = outputs.past_key_values
41 pre = input_embeds[:, :pos, :]
42 post = input_embeds[:, pos+1:, :]
43 input_embeds = torch.cat([pre, new_embed.unsqueeze(1), post], dim=1)
44
45 final_embeds = input_embeds[:, -1:, :] if past_key_values else input_embeds
46 outputs = model(inputs_embeds=final_embeds, attention_mask=attention_mask,
47 output_hidden_states=True, past_key_values=past_key_values,
48 use_cache=True, return_dict=True)
49 hidden_steps.append(outputs.hidden_states[-1][:, -1, :])
50
51 embeddings = torch.stack(hidden_steps, dim=1).mean(dim=1)
52 return F.normalize(embeddings, p=2, dim=-1)
53
54
55# Load model
56model_name = "Alibaba-NLP/LaSER-Qwen3-4B"
57tokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True)
58tokenizer.padding_side = "left"
59if tokenizer.pad_token_id is None:
60 tokenizer.pad_token = tokenizer.eos_token
61
62model = AutoModelForCausalLM.from_pretrained(
63 model_name, torch_dtype=torch.float16, trust_remote_code=True
64).cuda().eval()
65
66# Encode queries and documents
67with torch.inference_mode():
68 query_emb = laser_encode(model, tokenizer, ["why is the sky blue"], num_thinking_steps=3)
69 doc_emb = laser_encode(model, tokenizer, ["Rayleigh scattering makes short wavelengths scatter more strongly"], num_thinking_steps=3)
70
71# Compute similarity
72similarity = (query_emb @ doc_emb.T).item()
73print(f"Cosine similarity: {similarity:.4f}")