yocto-demo / app.py
reinforceai-labs's picture
Update app.py
63e3beb verified
Raw
History Blame Contribute Delete
15.2 kB
"""
YOCTO — Hugging Face Space
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
import gradio as gr
from huggingface_hub import hf_hub_download
from tokenizers import Tokenizer as HFTokenizer
from typing import Optional, Tuple, List
from dataclasses import dataclass, field
import time
# ==============================================================================
# CONFIG
# ==============================================================================
@dataclass
class ModelConfig:
vocab_size: int = 4000
embed_dim: int = 72
num_heads: int = 3
num_layers: int = 4
ffn_dim: int = 288
max_seq_len: int = 512
dropout: float = 0.1
head_dim: int = field(init=False)
third_dim: int = field(init=False)
component_head_dim: int = field(init=False)
def __post_init__(self):
self.head_dim = self.embed_dim // self.num_heads
self.third_dim = self.embed_dim // 3
self.component_head_dim = self.third_dim // self.num_heads
# ==============================================================================
# MODEL
# ==============================================================================
class RotaryPositionEmbedding(nn.Module):
def __init__(self, dim: int, max_seq_len: int = 2048, base: float = 10000.0):
super().__init__()
self.dim = dim
self.max_seq_len = max_seq_len
inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2).float() / dim))
self.register_buffer("inv_freq", inv_freq, persistent=False)
self._build_cache(max_seq_len)
def _build_cache(self, seq_len: int):
positions = torch.arange(seq_len, dtype=torch.float32, device=self.inv_freq.device)
freqs = torch.outer(positions, self.inv_freq)
emb = torch.cat([freqs, freqs], dim=-1)
self.register_buffer("cos_cached", emb.cos(), persistent=False)
self.register_buffer("sin_cached", emb.sin(), persistent=False)
self.max_seq_len = seq_len
def forward(self, x, seq_len, offset=0):
if offset + seq_len > self.max_seq_len:
self._build_cache(offset + seq_len)
return (
self.cos_cached[offset:offset + seq_len].to(x.dtype),
self.sin_cached[offset:offset + seq_len].to(x.dtype)
)
def rotate_half(x):
x1, x2 = x[..., :x.shape[-1]//2], x[..., x.shape[-1]//2:]
return torch.cat([-x2, x1], dim=-1)
def apply_rotary_pos_emb(s, o, cos, sin):
cos, sin = cos.unsqueeze(0).unsqueeze(0), sin.unsqueeze(0).unsqueeze(0)
return (s * cos) + (rotate_half(s) * sin), (o * cos) + (rotate_half(o) * sin)
class KVCache:
def __init__(self):
self.cache = []
def reset(self, num_layers):
self.cache = [None] * num_layers
def get(self, idx):
return self.cache[idx]
def update(self, idx, k, v):
if self.cache[idx] is None:
self.cache[idx] = (k, v)
else:
ck, cv = self.cache[idx]
self.cache[idx] = (torch.cat([ck, k], dim=2), torch.cat([cv, v], dim=2))
@property
def seq_len(self):
return self.cache[0][0].shape[2] if self.cache and self.cache[0] else 0
class UnifiedAttention(nn.Module):
def __init__(self, config, layer_idx=0):
super().__init__()
self.layer_idx = layer_idx
self.num_heads = config.num_heads
self.third_dim = config.third_dim
self.component_head_dim = config.component_head_dim
self.W_unified = nn.Linear(config.embed_dim, config.third_dim * 3, bias=False)
self.W_output = nn.Linear(config.third_dim, config.embed_dim, bias=False)
self.rope = RotaryPositionEmbedding(config.component_head_dim, config.max_seq_len)
def forward(self, x, kv_cache=None, use_cache=False):
B, S, _ = x.shape
unified = self.W_unified(x)
seeking, offering, content = unified.split(self.third_dim, dim=-1)
seeking = seeking.view(B, S, self.num_heads, self.component_head_dim).transpose(1, 2)
offering = offering.view(B, S, self.num_heads, self.component_head_dim).transpose(1, 2)
content = content.view(B, S, self.num_heads, self.component_head_dim).transpose(1, 2)
offset = kv_cache.seq_len if kv_cache else 0
cos, sin = self.rope(seeking, S, offset)
seeking, offering = apply_rotary_pos_emb(seeking, offering, cos, sin)
if kv_cache:
cached = kv_cache.get(self.layer_idx)
if cached:
offering = torch.cat([cached[0], offering], dim=2)
content = torch.cat([cached[1], content], dim=2)
if use_cache:
new_k = unified[..., self.third_dim:2*self.third_dim].view(B, S, self.num_heads, self.component_head_dim).transpose(1, 2)
new_k, _ = apply_rotary_pos_emb(new_k, new_k, cos, sin)
new_v = unified[..., 2*self.third_dim:].view(B, S, self.num_heads, self.component_head_dim).transpose(1, 2)
kv_cache.update(self.layer_idx, new_k, new_v)
out = F.scaled_dot_product_attention(seeking, offering, content, is_causal=(not kv_cache or kv_cache.seq_len == 0))
return self.W_output(out.transpose(1, 2).contiguous().view(B, S, self.third_dim))
class FeedForward(nn.Module):
def __init__(self, config):
super().__init__()
self.fc1 = nn.Linear(config.embed_dim, config.ffn_dim)
self.fc2 = nn.Linear(config.ffn_dim, config.embed_dim)
def forward(self, x):
return self.fc2(F.gelu(self.fc1(x)))
class TransformerBlock(nn.Module):
def __init__(self, config, layer_idx=0):
super().__init__()
self.attn = UnifiedAttention(config, layer_idx)
self.ffn = FeedForward(config)
self.norm1 = nn.LayerNorm(config.embed_dim)
self.norm2 = nn.LayerNorm(config.embed_dim)
def forward(self, x, kv_cache=None, use_cache=False):
x = x + self.attn(self.norm1(x), kv_cache, use_cache)
x = x + self.ffn(self.norm2(x))
return x
class Yocto(nn.Module):
def __init__(self, config):
super().__init__()
self.config = config
self.token_embedding = nn.Embedding(config.vocab_size, config.embed_dim)
self.blocks = nn.ModuleList([TransformerBlock(config, i) for i in range(config.num_layers)])
self.norm = nn.LayerNorm(config.embed_dim)
self.output = nn.Linear(config.embed_dim, config.vocab_size, bias=False)
self.output.weight = self.token_embedding.weight
def forward(self, input_ids, kv_cache=None, use_cache=False):
x = self.token_embedding(input_ids)
for block in self.blocks:
x = block(x, kv_cache, use_cache)
return self.output(self.norm(x))
# ==============================================================================
# LOAD MODEL
# ==============================================================================
def load_model():
model_path = hf_hub_download(repo_id="Reinforce-ai/yocto", filename="model.pt")
tokenizer_path = hf_hub_download(repo_id="Reinforce-ai/yocto", filename="tokenizer.json")
checkpoint = torch.load(model_path, map_location="cpu", weights_only=False)
cfg = checkpoint['config']['model']
config = ModelConfig(
vocab_size=cfg['vocab_size'],
embed_dim=cfg['embed_dim'],
num_heads=cfg['num_heads'],
num_layers=cfg['num_layers'],
ffn_dim=cfg['ffn_dim'],
max_seq_len=cfg.get('max_seq_len', 512),
dropout=cfg.get('dropout', 0.1),
)
model = Yocto(config)
model.load_state_dict(checkpoint['model_state_dict'])
model.eval()
tokenizer = HFTokenizer.from_file(tokenizer_path)
return model, tokenizer, config
MODEL, TOKENIZER, CONFIG = load_model()
# ==============================================================================
# STREAMING GENERATION
# ==============================================================================
@torch.no_grad()
def generate_stream(prompt, max_tokens=150, temperature=0.8, top_k=50, top_p=0.95):
input_ids = torch.tensor([TOKENIZER.encode(prompt).ids])
kv_cache = KVCache()
kv_cache.reset(len(MODEL.blocks))
logits = MODEL(input_ids, kv_cache=kv_cache, use_cache=True)
output_text = f"{prompt}"
yield output_text
token_count = 0
start_time = time.perf_counter()
for _ in range(max_tokens):
next_logits = logits[:, -1, :] / temperature
if top_k > 0:
v, _ = torch.topk(next_logits, min(top_k, next_logits.size(-1)))
next_logits[next_logits < v[:, -1:]] = float('-inf')
if top_p < 1.0:
sorted_logits, sorted_idx = torch.sort(next_logits, descending=True)
cumsum = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)
mask = cumsum > top_p
mask[..., 1:] = mask[..., :-1].clone()
mask[..., 0] = False
next_logits[mask.scatter(1, sorted_idx, mask)] = float('-inf')
probs = F.softmax(next_logits, dim=-1)
next_token = torch.multinomial(probs, 1)
if next_token.item() == 3:
break
token_text = TOKENIZER.decode([next_token.item()])
output_text += token_text
token_count += 1
yield output_text
logits = MODEL(next_token, kv_cache=kv_cache, use_cache=True)
elapsed = time.perf_counter() - start_time
tokens_per_sec = token_count / elapsed if elapsed > 0 else 0
output_text += f"\n\n─────────────────────────────────────────\n⚡ {token_count} tokens • {elapsed:.2f}s • {tokens_per_sec:.1f} tok/s"
yield output_text
# ==============================================================================
# GRADIO INTERFACE
# ==============================================================================
CSS = """
@import url('https://fonts.googleapis.com/css2?family=VT323&family=Space+Mono&display=swap');
/* Global */
.gradio-container {
background: #000 !important;
max-width: 100% !important;
}
footer {display: none !important;}
/* All textareas and inputs */
textarea, input[type="text"] {
background: #111 !important;
color: #0f0 !important;
border: 2px solid #0f0 !important;
border-radius: 0 !important;
font-family: 'VT323', monospace !important;
font-size: 1.3rem !important;
caret-color: #0f0 !important;
}
textarea:focus, input[type="text"]:focus {
outline: none !important;
box-shadow: 0 0 20px rgba(0,255,0,0.3) !important;
}
textarea::placeholder, input::placeholder {
color: #060 !important;
}
/* Labels */
label {
color: #0f0 !important;
font-family: 'Space Mono', monospace !important;
}
/* Buttons */
button {
font-family: 'Space Mono', monospace !important;
}
#generate-btn {
background: #0f0 !important;
color: #000 !important;
border: none !important;
border-radius: 0 !important;
font-size: 1.1rem !important;
font-weight: bold !important;
padding: 12px 30px !important;
text-transform: uppercase !important;
}
#generate-btn:hover {
background: #0c0 !important;
box-shadow: 0 0 25px rgba(0,255,0,0.5) !important;
}
.chip-btn {
background: transparent !important;
color: #0f0 !important;
border: 1px solid #080 !important;
border-radius: 20px !important;
font-size: 0.9rem !important;
padding: 8px 18px !important;
margin: 4px !important;
}
.chip-btn:hover {
background: rgba(0,255,0,0.1) !important;
border-color: #0f0 !important;
box-shadow: 0 0 15px rgba(0,255,0,0.3) !important;
}
"""
def generate_story(prompt):
if not prompt.strip():
prompt = "Once upon a time"
for output in generate_stream(prompt, max_tokens=150, temperature=0.8, top_k=50, top_p=0.95):
yield output
with gr.Blocks(title="YOCTO") as demo:
# Header
gr.HTML("""
<div style="text-align:center; padding:30px 0 20px 0;">
<h1 style="font-family:'VT323',monospace; font-size:5rem; color:#0f0;
text-shadow:0 0 20px #0f0, 0 0 40px #0f0; letter-spacing:0.4em; margin:0;">
YOCTO
</h1>
<p style="font-family:'Space Mono',monospace; color:#0a0; font-size:0.9rem; letter-spacing:0.15em; margin-top:10px;">
WORLD'S SMALLEST STORY TELLER • 484K PARAMS • 946 KB
</p>
</div>
""")
# Output
output = gr.Textbox(
label="Output",
value="Ready to generate stories...\n\nType a prompt below or click a suggestion.",
lines=14,
interactive=False
)
# Input + Button
with gr.Row():
prompt = gr.Textbox(
label="Your prompt",
placeholder="Start your story here...",
lines=1,
scale=5
)
generate_btn = gr.Button("GENERATE", scale=1, elem_id="generate-btn")
# Prompt chips
gr.HTML("<div style='text-align:center; margin:15px 0 5px 0; color:#0a0; font-family:Space Mono,monospace; font-size:0.8rem;'>TRY THESE:</div>")
with gr.Row():
btn1 = gr.Button("Once upon a time", elem_classes="chip-btn")
btn2 = gr.Button("There was a little girl named Lily", elem_classes="chip-btn")
btn3 = gr.Button("One day, a curious boy found", elem_classes="chip-btn")
btn4 = gr.Button("The little rabbit was very", elem_classes="chip-btn")
btn5 = gr.Button("Mommy and Timmy went to the", elem_classes="chip-btn")
btn6 = gr.Button("There was a big red", elem_classes="chip-btn")
# Footer
gr.HTML("""
<div style="text-align:center; padding:20px 0; margin-top:20px; border-top:1px solid #030;">
<span style="font-family:'Space Mono',monospace; font-size:0.75rem; color:#080;">
REINFORCE AI © 2025 •
<a href="https://www.reinforceai.com/yocto" style="color:#0f0; text-decoration:none;">WEBSITE</a> •
<a href="https://github.com/reinforceai/yocto" style="color:#0f0; text-decoration:none;">GITHUB</a> •
<a href="https://huggingface.co/Reinforce-ai/yocto" style="color:#0f0; text-decoration:none;">MODEL</a>
</span>
</div>
""")
# Events
btn1.click(lambda: "Once upon a time", outputs=prompt)
btn2.click(lambda: "There was a little girl named Lily", outputs=prompt)
btn3.click(lambda: "One day, a curious boy found", outputs=prompt)
btn4.click(lambda: "The little rabbit was very", outputs=prompt)
btn5.click(lambda: "Mommy and Timmy went to the", outputs=prompt)
btn6.click(lambda: "There was a big red", outputs=prompt)
generate_btn.click(generate_story, inputs=[prompt], outputs=[output])
prompt.submit(generate_story, inputs=[prompt], outputs=[output])
if __name__ == "__main__":
demo.launch(server_name="0.0.0.0", server_port=7860, css=CSS)