Views
No views yet
x * (1 + scale[lang_id]) + shift[lang_id]wne suffix positional embedding for FIM1from modeling_fimmy import load_fimmy_model
2from transformers import GPT2Tokenizer
3import torch
4
5model = load_fimmy_model("di-zhang-fdu/fimmy-small-java")
6tok = GPT2Tokenizer.from_pretrained("gpt2")
7LANG_ID = 2
8
9# Greedy completion
10ids = tok.encode("def hello():
11 ", return_tensors="pt")
12out = model.generate(ids, lang_id=LANG_ID, max_new_tokens=15)
13print(tok.decode(out[0, ids.shape[1]:]))
14
15# Beam search
16out = model.generate(ids, lang_id=LANG_ID, beam_width=3, max_new_tokens=15)
17
18# Temperature sampling
19out = model.generate(ids, lang_id=LANG_ID, temperature=0.7, top_k=5, top_p=0.9)
20
21# Extend to 16k context (zero quality loss for original 1024)
22model.transformer.extend_positional_embeddings(scale=16, method="yarn")before (prefix) and after (suffix) context.1before = "def add(a, b):
2 return "
3after = "a + b
4
5result = add(1, 2)"
6
7tok = GPT2Tokenizer.from_pretrained("gpt2")
8
9# Method 1: Suffix-first FIM (recommended)
10# Put suffix before prefix so the model sees both contexts via causal attention
11after_ids = tok.encode(after)
12before_ids = tok.encode(before)
13combined = torch.tensor([after_ids + before_ids])
14out = model.generate(combined, lang_id=LANG_ID, max_new_tokens=10)
15print(tok.decode(out[0, combined.shape[1]:]))
16
17# Method 2: Prefix-only completion (no suffix context)
18ids = tok.encode(before, return_tensors="pt")
19out = model.generate(ids, lang_id=LANG_ID, max_new_tokens=10)
20print(tok.decode(out[0, ids.shape[1]:]))
21
22# Method 3: wne-based FIM (experimental)
23# Uses wne (suffix positional embedding) to encode suffix position
24def fim_forward(model, before_ids, after_ids, lang_id):
25 all_ids = torch.cat([before_ids, after_ids], dim=1)
26 B, T = all_ids.shape
27 with torch.no_grad():
28 pos = torch.arange(T).unsqueeze(0)
29 x = model.transformer.wte(all_ids) + model.transformer.wpe(pos)
30 # Add wne positional embedding to suffix tokens
31 after_len = after_ids.shape[1]
32 if after_len > 0:
33 suffix_pos = torch.arange(after_len).unsqueeze(0)
34 x[0, before_ids.shape[1]:] += model.transformer.wne(suffix_pos)[0]
35 for block in model.transformer.h:
36 x, _ = block(x, lang_id=lang_id)
37 x = model.transformer.ln_f(x, lang_id)
38 logits = model.lm_head(x)
39 return logits[0, before_ids.shape[1] - 1]
40
41next_logits = fim_forward(model, before_ids, after_ids, lang_id=LANG_ID)
42print(tok.decode([next_logits.argmax().item()]))| Field | Value |
|---|---|
| n_layer | 12 |
| n_head | 12 |
| n_embd | 768 |
| intermediate_size | 3072 |
| vocab_size | 50257 |
| n_positions | 1024 (extendable to 16k+) |
| n_lang | 30 |
| language | java |
| language_id | 2 |