fimmy-small-java
Java-specialized code completion model.
Architecture
GPT-2 variant with language-conditional LayerNorm (12 layers, 12 heads, dim=768, 126M params).
- Per-language affine after LayerNorm:
x * (1 + scale[lang_id]) + shift[lang_id] - KV cache for efficient autoregressive generation
- Beam search + temperature/top-k/top-p sampling
- Context extension (YaRN/periodic/linear) from 1024 to 16k+
wnesuffix positional embedding for FIM- GPT-2 BPE tokenizer (vocab_size=50257)
Quick Start
from modeling_fimmy import load_fimmy_model
from transformers import GPT2Tokenizer
import torch
model = load_fimmy_model("di-zhang-fdu/fimmy-small-java")
tok = GPT2Tokenizer.from_pretrained("gpt2")
LANG_ID = 2
# Greedy completion
ids = tok.encode("def hello():
", return_tensors="pt")
out = model.generate(ids, lang_id=LANG_ID, max_new_tokens=15)
print(tok.decode(out[0, ids.shape[1]:]))
# Beam search
out = model.generate(ids, lang_id=LANG_ID, beam_width=3, max_new_tokens=15)
# Temperature sampling
out = model.generate(ids, lang_id=LANG_ID, temperature=0.7, top_k=5, top_p=0.9)
# Extend to 16k context (zero quality loss for original 1024)
model.transformer.extend_positional_embeddings(scale=16, method="yarn")
Fill-In-the-Middle (FIM)
FIM completes code between a before (prefix) and after (suffix) context.
before = "def add(a, b):
return "
after = "a + b
result = add(1, 2)"
tok = GPT2Tokenizer.from_pretrained("gpt2")
# Method 1: Suffix-first FIM (recommended)
# Put suffix before prefix so the model sees both contexts via causal attention
after_ids = tok.encode(after)
before_ids = tok.encode(before)
combined = torch.tensor([after_ids + before_ids])
out = model.generate(combined, lang_id=LANG_ID, max_new_tokens=10)
print(tok.decode(out[0, combined.shape[1]:]))
# Method 2: Prefix-only completion (no suffix context)
ids = tok.encode(before, return_tensors="pt")
out = model.generate(ids, lang_id=LANG_ID, max_new_tokens=10)
print(tok.decode(out[0, ids.shape[1]:]))
# Method 3: wne-based FIM (experimental)
# Uses wne (suffix positional embedding) to encode suffix position
def fim_forward(model, before_ids, after_ids, lang_id):
all_ids = torch.cat([before_ids, after_ids], dim=1)
B, T = all_ids.shape
with torch.no_grad():
pos = torch.arange(T).unsqueeze(0)
x = model.transformer.wte(all_ids) + model.transformer.wpe(pos)
# Add wne positional embedding to suffix tokens
after_len = after_ids.shape[1]
if after_len > 0:
suffix_pos = torch.arange(after_len).unsqueeze(0)
x[0, before_ids.shape[1]:] += model.transformer.wne(suffix_pos)[0]
for block in model.transformer.h:
x, _ = block(x, lang_id=lang_id)
x = model.transformer.ln_f(x, lang_id)
logits = model.lm_head(x)
return logits[0, before_ids.shape[1] - 1]
next_logits = fim_forward(model, before_ids, after_ids, lang_id=LANG_ID)
print(tok.decode([next_logits.argmax().item()]))
Configuration
| 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 |
All Models
- fimmy-nano β 6L, 83M, generic
- fimmy-small-python β 12L, 126M, python
- fimmy-small-java β 12L, 126M, java (this model)
- fimmy-small-go β 12L, 126M, go
- fimmy-small-ruby β 12L, 126M, ruby
- fimmy-small-rust β 12L, 126M, rust
- fimmy-small-cpp β 12L, 126M, cpp
- fimmy-small-dart β 12L, 126M, dart
- fimmy-small-julia β 12L, 126M, julia
- fimmy-small-hcl β 12L, 126M, hcl
- fimmy-small-generic β 12L, 126M, generic
- fimmy-medium-generic β 24L, 359M, generic
License
MIT
- Downloads last month
- -