File size: 1,143 Bytes
963a458
 
787da7f
93906d0
787da7f
 
 
93906d0
787da7f
 
 
 
 
 
 
 
 
963a458
 
8cbc403
963a458
93906d0
 
 
787da7f
93906d0
963a458
 
 
41fdbec
963a458
 
93906d0
787da7f
963a458
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
import gradio as gr
import torch
from transformers import AutoTokenizer, GenerationConfig, GPT2Config, GPT2LMHeadModel

# 1. Load the standard GPT-2 tokenizer
tokenizer = AutoTokenizer.from_pretrained("gpt2")
tokenizer.pad_token = tokenizer.eos_token

# 3. Load the model 
config = GPT2Config.from_pretrained(
    pretrained_model_name_or_path = "gpt2",
    vocab_size = len(tokenizer),
    n_ctx = 256,
    bos_token_id = tokenizer.bos_token_id,
    eos_token_id = tokenizer.bos_token_id
)
model = GPT2LMHeadModel(config)

def generate_code(prompt):
    inputs = tokenizer(prompt, return_tensors="pt").to(model.device)
    
    with torch.no_grad():
        outputs = model.generate(
            **inputs,
            pad_token_id=tokenizer.eos_token_id
        )
    
    return tokenizer.decode(outputs[0], skip_special_tokens=True)

# 5. Gradio Interface
demo = gr.Interface(
    fn=generate_code,
    inputs=gr.Textbox(placeholder="Write a function to...", label="Input Prompt"),
    outputs=gr.Code(label="GPT-2 Generated Code", language="python"),
    title="Small Code Snippet Generator"
)

if __name__ == "__main__":
    demo.launch()