Spaces:
Sleeping
Sleeping
Update app.py
Browse files
app.py
CHANGED
|
@@ -26,22 +26,41 @@ class StringStoppingCriteria(StoppingCriteria):
|
|
| 26 |
def create_demo(api_url, model_id=None):
|
| 27 |
# Standalone mode check
|
| 28 |
hf_model_id = model_id or os.environ.get("HF_MODEL_ID", "ethicalabs/Echo-DSRN-Small-Instruct")
|
| 29 |
-
|
| 30 |
local_model = None
|
| 31 |
local_tokenizer = None
|
| 32 |
-
|
|
|
|
| 33 |
if hf_model_id:
|
| 34 |
print(f"📦 Standalone Mode: Loading {hf_model_id} via Transformers...")
|
| 35 |
try:
|
| 36 |
local_tokenizer = AutoTokenizer.from_pretrained(hf_model_id, trust_remote_code=True)
|
| 37 |
local_model = AutoModelForCausalLM.from_pretrained(
|
| 38 |
-
hf_model_id,
|
| 39 |
trust_remote_code=True,
|
| 40 |
torch_dtype=torch.bfloat16 if torch.cuda.is_available() else torch.float32,
|
| 41 |
-
device_map="auto"
|
| 42 |
)
|
| 43 |
local_model.eval()
|
| 44 |
print("✅ Model loaded successfully.")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 45 |
except Exception as e:
|
| 46 |
print(f"❌ Failed to load model: {e}")
|
| 47 |
traceback.print_exc()
|
|
@@ -51,43 +70,49 @@ def create_demo(api_url, model_id=None):
|
|
| 51 |
font-family: 'JetBrains Mono', 'Courier New', monospace;
|
| 52 |
line-height: 1.8;
|
| 53 |
padding: 24px;
|
| 54 |
-
background: #
|
| 55 |
border-radius: 12px;
|
| 56 |
color: #e2e8f0;
|
| 57 |
min-height: 250px;
|
| 58 |
white-space: pre-wrap;
|
| 59 |
overflow-wrap: break-word;
|
| 60 |
border: 1px solid #1e293b;
|
|
|
|
|
|
|
| 61 |
}
|
| 62 |
.token {
|
| 63 |
-
display: inline;
|
| 64 |
-
padding:
|
| 65 |
-
border-radius:
|
| 66 |
position: relative;
|
| 67 |
cursor: help;
|
| 68 |
-
transition:
|
| 69 |
}
|
| 70 |
.token:hover {
|
| 71 |
-
background: rgba(255, 255, 255, 0.
|
|
|
|
| 72 |
}
|
| 73 |
.token-tooltip {
|
| 74 |
-
|
| 75 |
position: absolute;
|
| 76 |
-
bottom:
|
| 77 |
left: 50%;
|
| 78 |
transform: translateX(-50%);
|
| 79 |
background: #1e293b;
|
| 80 |
color: white;
|
| 81 |
-
padding:
|
| 82 |
-
border-radius:
|
| 83 |
-
font-size: 0.
|
| 84 |
-
z-index:
|
| 85 |
white-space: nowrap;
|
| 86 |
-
box-shadow: 0
|
| 87 |
-
border: 1px solid
|
|
|
|
|
|
|
|
|
|
| 88 |
}
|
| 89 |
.token:hover .token-tooltip {
|
| 90 |
-
|
| 91 |
}
|
| 92 |
.mode-badge {
|
| 93 |
display: inline-block;
|
|
@@ -99,21 +124,46 @@ def create_demo(api_url, model_id=None):
|
|
| 99 |
}
|
| 100 |
.mode-standalone { background: #059669; color: white; }
|
| 101 |
.mode-api { background: #4f46e5; color: white; }
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 102 |
"""
|
| 103 |
|
| 104 |
-
def generate_html(prompt, max_tokens, temperature, top_k):
|
| 105 |
try:
|
| 106 |
if local_model and local_tokenizer:
|
| 107 |
# Standalone Mode Logic
|
| 108 |
input_tokens = local_tokenizer(prompt, return_tensors="pt").to(local_model.device)
|
| 109 |
input_ids = input_tokens.input_ids
|
| 110 |
-
|
| 111 |
stop_strings = ["<|im_end|>", "<|end|>", "<|user|>"]
|
| 112 |
-
stopping_criteria = StoppingCriteriaList(
|
| 113 |
-
|
|
|
|
|
|
|
| 114 |
all_tokens = []
|
| 115 |
all_logprobs = []
|
| 116 |
-
|
| 117 |
# Echo (Prompt) Logprobs
|
| 118 |
with torch.no_grad():
|
| 119 |
outputs = local_model(input_ids)
|
|
@@ -122,13 +172,13 @@ def create_demo(api_url, model_id=None):
|
|
| 122 |
token_id = input_ids[0, i].item()
|
| 123 |
token_text = local_tokenizer.decode([token_id])
|
| 124 |
all_tokens.append(token_text)
|
| 125 |
-
|
| 126 |
if i == 0:
|
| 127 |
all_logprobs.append(None)
|
| 128 |
else:
|
| 129 |
-
lp = torch.nn.functional.log_softmax(logits[0, i-1, :], dim=-1)
|
| 130 |
all_logprobs.append(lp[token_id].item())
|
| 131 |
-
|
| 132 |
# Generation
|
| 133 |
with torch.no_grad():
|
| 134 |
gen_out = local_model.generate(
|
|
@@ -137,27 +187,27 @@ def create_demo(api_url, model_id=None):
|
|
| 137 |
temperature=temperature,
|
| 138 |
top_k=top_k,
|
| 139 |
top_p=0.9,
|
| 140 |
-
repetition_penalty=
|
| 141 |
do_sample=temperature > 0,
|
| 142 |
use_cache=False,
|
| 143 |
output_scores=True,
|
| 144 |
return_dict_in_generate=True,
|
| 145 |
stopping_criteria=stopping_criteria,
|
| 146 |
pad_token_id=local_tokenizer.pad_token_id or 32000,
|
| 147 |
-
eos_token_id=[32000, 32007, 32011]
|
| 148 |
)
|
| 149 |
-
|
| 150 |
output_ids = gen_out.sequences
|
| 151 |
scores = gen_out.scores
|
| 152 |
-
generated_ids = output_ids[0, input_ids.shape[1]:]
|
| 153 |
-
|
| 154 |
for i, token_id in enumerate(generated_ids):
|
| 155 |
token_id = token_id.item()
|
| 156 |
token_text = local_tokenizer.decode([token_id])
|
| 157 |
all_tokens.append(token_text)
|
| 158 |
lp = torch.nn.functional.log_softmax(scores[i][0, :], dim=-1)
|
| 159 |
all_logprobs.append(lp[token_id].item())
|
| 160 |
-
|
| 161 |
logprobs = {"tokens": all_tokens, "token_logprobs": all_logprobs}
|
| 162 |
else:
|
| 163 |
# API Mode
|
|
@@ -166,6 +216,7 @@ def create_demo(api_url, model_id=None):
|
|
| 166 |
"max_tokens": max_tokens,
|
| 167 |
"temperature": temperature,
|
| 168 |
"top_k": top_k,
|
|
|
|
| 169 |
"logprobs": 1,
|
| 170 |
"echo": True,
|
| 171 |
}
|
|
@@ -190,14 +241,21 @@ def create_demo(api_url, model_id=None):
|
|
| 190 |
color = "#e2e8f0" # Default
|
| 191 |
if lp is not None:
|
| 192 |
if prob > 0.8:
|
| 193 |
-
color = "#4ade80" # Green
|
| 194 |
elif prob < 0.2:
|
| 195 |
-
color = "#f87171" # Red
|
| 196 |
|
| 197 |
prob_pct = f"{(prob * 100):.1f}%" if lp is not None else "N/A"
|
|
|
|
| 198 |
|
| 199 |
html += f'<span class="token" style="color: {color}">{display_token}'
|
| 200 |
-
html += f'<span class="token-tooltip">
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 201 |
|
| 202 |
html += "</div>"
|
| 203 |
return html
|
|
@@ -205,22 +263,38 @@ def create_demo(api_url, model_id=None):
|
|
| 205 |
traceback.print_exc()
|
| 206 |
return f'<div class="token-container" style="color: #f87171">❌ Error: {str(e)}</div>'
|
| 207 |
|
| 208 |
-
with gr.Blocks(title="Echo-DSRN
|
| 209 |
-
gr.
|
| 210 |
-
|
| 211 |
-
|
| 212 |
-
|
| 213 |
-
|
| 214 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 215 |
|
| 216 |
-
|
| 217 |
-
|
| 218 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 219 |
|
| 220 |
gr.Markdown(
|
| 221 |
"""
|
| 222 |
-
|
| 223 |
-
Hover over tokens to see their prediction confidence.
|
| 224 |
"""
|
| 225 |
)
|
| 226 |
|
|
@@ -238,6 +312,9 @@ def create_demo(api_url, model_id=None):
|
|
| 238 |
)
|
| 239 |
with gr.Row():
|
| 240 |
top_k = gr.Slider(minimum=0, maximum=100, value=40, step=1, label="Top K")
|
|
|
|
|
|
|
|
|
|
| 241 |
|
| 242 |
submit_btn = gr.Button("Generate with Probabilities", variant="primary")
|
| 243 |
|
|
@@ -246,16 +323,18 @@ def create_demo(api_url, model_id=None):
|
|
| 246 |
output_html = gr.HTML(label="Visualized Output")
|
| 247 |
|
| 248 |
submit_btn.click(
|
| 249 |
-
generate_html,
|
|
|
|
|
|
|
| 250 |
)
|
| 251 |
|
| 252 |
gr.Examples(
|
| 253 |
examples=[
|
| 254 |
-
["The capital of France is", 128, 0.1, 40],
|
| 255 |
-
["Responsible AI development requires", 128, 0.7, 40],
|
| 256 |
-
["To be or not to be, that is the", 128, 0.1, 40],
|
| 257 |
],
|
| 258 |
-
inputs=[prompt, max_tokens, temp, top_k],
|
| 259 |
)
|
| 260 |
|
| 261 |
return demo, css
|
|
@@ -267,10 +346,10 @@ if __name__ == "__main__":
|
|
| 267 |
"--api_url", type=str, default="http://localhost:5000", help="URL of the Echo-DSRN API"
|
| 268 |
)
|
| 269 |
parser.add_argument(
|
| 270 |
-
"--model_id",
|
| 271 |
-
type=str,
|
| 272 |
-
default=
|
| 273 |
-
help="Local path or HF ID to the model for Standalone mode"
|
| 274 |
)
|
| 275 |
parser.add_argument("--port", type=int, default=7860, help="Port to run the Gradio app on")
|
| 276 |
args = parser.parse_args()
|
|
|
|
| 26 |
def create_demo(api_url, model_id=None):
|
| 27 |
# Standalone mode check
|
| 28 |
hf_model_id = model_id or os.environ.get("HF_MODEL_ID", "ethicalabs/Echo-DSRN-Small-Instruct")
|
| 29 |
+
|
| 30 |
local_model = None
|
| 31 |
local_tokenizer = None
|
| 32 |
+
model_metadata = None
|
| 33 |
+
|
| 34 |
if hf_model_id:
|
| 35 |
print(f"📦 Standalone Mode: Loading {hf_model_id} via Transformers...")
|
| 36 |
try:
|
| 37 |
local_tokenizer = AutoTokenizer.from_pretrained(hf_model_id, trust_remote_code=True)
|
| 38 |
local_model = AutoModelForCausalLM.from_pretrained(
|
| 39 |
+
hf_model_id,
|
| 40 |
trust_remote_code=True,
|
| 41 |
torch_dtype=torch.bfloat16 if torch.cuda.is_available() else torch.float32,
|
| 42 |
+
device_map="auto",
|
| 43 |
)
|
| 44 |
local_model.eval()
|
| 45 |
print("✅ Model loaded successfully.")
|
| 46 |
+
|
| 47 |
+
# Calculate metadata
|
| 48 |
+
total_params = sum(p.numel() for p in local_model.parameters())
|
| 49 |
+
trainable_params = sum(p.numel() for p in local_model.parameters() if p.requires_grad)
|
| 50 |
+
config = local_model.config
|
| 51 |
+
|
| 52 |
+
model_metadata = {
|
| 53 |
+
"Total Parameters": f"{total_params:,}",
|
| 54 |
+
"Trainable Parameters": f"{trainable_params:,}",
|
| 55 |
+
"Layers": getattr(
|
| 56 |
+
config, "num_hidden_layers", getattr(config, "num_layers", "N/A")
|
| 57 |
+
),
|
| 58 |
+
"Hidden Size": getattr(config, "hidden_size", "N/A"),
|
| 59 |
+
"MLP Ratio": getattr(config, "mlp_ratio", "N/A"),
|
| 60 |
+
"Vocab Size": getattr(config, "vocab_size", "N/A"),
|
| 61 |
+
"Surprise Lambda Init": getattr(config, "surprise_lambda_init", "N/A"),
|
| 62 |
+
"Model Type": getattr(config, "model_type", "echo"),
|
| 63 |
+
}
|
| 64 |
except Exception as e:
|
| 65 |
print(f"❌ Failed to load model: {e}")
|
| 66 |
traceback.print_exc()
|
|
|
|
| 70 |
font-family: 'JetBrains Mono', 'Courier New', monospace;
|
| 71 |
line-height: 1.8;
|
| 72 |
padding: 24px;
|
| 73 |
+
background: #020617;
|
| 74 |
border-radius: 12px;
|
| 75 |
color: #e2e8f0;
|
| 76 |
min-height: 250px;
|
| 77 |
white-space: pre-wrap;
|
| 78 |
overflow-wrap: break-word;
|
| 79 |
border: 1px solid #1e293b;
|
| 80 |
+
position: relative;
|
| 81 |
+
overflow: visible !important;
|
| 82 |
}
|
| 83 |
.token {
|
| 84 |
+
display: inline-block;
|
| 85 |
+
padding: 0 1px;
|
| 86 |
+
border-radius: 3px;
|
| 87 |
position: relative;
|
| 88 |
cursor: help;
|
| 89 |
+
transition: all 0.2s;
|
| 90 |
}
|
| 91 |
.token:hover {
|
| 92 |
+
background: rgba(255, 255, 255, 0.15) !important;
|
| 93 |
+
z-index: 100;
|
| 94 |
}
|
| 95 |
.token-tooltip {
|
| 96 |
+
display: none;
|
| 97 |
position: absolute;
|
| 98 |
+
bottom: 130%;
|
| 99 |
left: 50%;
|
| 100 |
transform: translateX(-50%);
|
| 101 |
background: #1e293b;
|
| 102 |
color: white;
|
| 103 |
+
padding: 12px;
|
| 104 |
+
border-radius: 10px;
|
| 105 |
+
font-size: 0.85rem;
|
| 106 |
+
z-index: 1000;
|
| 107 |
white-space: nowrap;
|
| 108 |
+
box-shadow: 0 20px 25px -5px rgba(0, 0, 0, 0.5);
|
| 109 |
+
border: 1px solid rgba(255, 255, 255, 0.2);
|
| 110 |
+
pointer-events: none;
|
| 111 |
+
text-align: center;
|
| 112 |
+
min-width: 120px;
|
| 113 |
}
|
| 114 |
.token:hover .token-tooltip {
|
| 115 |
+
display: block;
|
| 116 |
}
|
| 117 |
.mode-badge {
|
| 118 |
display: inline-block;
|
|
|
|
| 124 |
}
|
| 125 |
.mode-standalone { background: #059669; color: white; }
|
| 126 |
.mode-api { background: #4f46e5; color: white; }
|
| 127 |
+
|
| 128 |
+
.metadata-card {
|
| 129 |
+
background: #0f172a;
|
| 130 |
+
padding: 15px;
|
| 131 |
+
border-radius: 8px;
|
| 132 |
+
border: 1px solid #1e293b;
|
| 133 |
+
font-family: inherit;
|
| 134 |
+
}
|
| 135 |
+
.metadata-grid {
|
| 136 |
+
display: grid;
|
| 137 |
+
grid-template-columns: repeat(auto-fill, minmax(180px, 1fr));
|
| 138 |
+
gap: 10px;
|
| 139 |
+
margin-top: 10px;
|
| 140 |
+
}
|
| 141 |
+
.metadata-item {
|
| 142 |
+
font-size: 0.85rem;
|
| 143 |
+
color: #94a3b8;
|
| 144 |
+
}
|
| 145 |
+
.metadata-value {
|
| 146 |
+
font-weight: 600;
|
| 147 |
+
color: #f1f5f9;
|
| 148 |
+
display: block;
|
| 149 |
+
}
|
| 150 |
"""
|
| 151 |
|
| 152 |
+
def generate_html(prompt, max_tokens, temperature, top_k, rep_penalty):
|
| 153 |
try:
|
| 154 |
if local_model and local_tokenizer:
|
| 155 |
# Standalone Mode Logic
|
| 156 |
input_tokens = local_tokenizer(prompt, return_tensors="pt").to(local_model.device)
|
| 157 |
input_ids = input_tokens.input_ids
|
| 158 |
+
|
| 159 |
stop_strings = ["<|im_end|>", "<|end|>", "<|user|>"]
|
| 160 |
+
stopping_criteria = StoppingCriteriaList(
|
| 161 |
+
[StringStoppingCriteria(local_tokenizer, stop_strings)]
|
| 162 |
+
)
|
| 163 |
+
|
| 164 |
all_tokens = []
|
| 165 |
all_logprobs = []
|
| 166 |
+
|
| 167 |
# Echo (Prompt) Logprobs
|
| 168 |
with torch.no_grad():
|
| 169 |
outputs = local_model(input_ids)
|
|
|
|
| 172 |
token_id = input_ids[0, i].item()
|
| 173 |
token_text = local_tokenizer.decode([token_id])
|
| 174 |
all_tokens.append(token_text)
|
| 175 |
+
|
| 176 |
if i == 0:
|
| 177 |
all_logprobs.append(None)
|
| 178 |
else:
|
| 179 |
+
lp = torch.nn.functional.log_softmax(logits[0, i - 1, :], dim=-1)
|
| 180 |
all_logprobs.append(lp[token_id].item())
|
| 181 |
+
|
| 182 |
# Generation
|
| 183 |
with torch.no_grad():
|
| 184 |
gen_out = local_model.generate(
|
|
|
|
| 187 |
temperature=temperature,
|
| 188 |
top_k=top_k,
|
| 189 |
top_p=0.9,
|
| 190 |
+
repetition_penalty=rep_penalty,
|
| 191 |
do_sample=temperature > 0,
|
| 192 |
use_cache=False,
|
| 193 |
output_scores=True,
|
| 194 |
return_dict_in_generate=True,
|
| 195 |
stopping_criteria=stopping_criteria,
|
| 196 |
pad_token_id=local_tokenizer.pad_token_id or 32000,
|
| 197 |
+
eos_token_id=[32000, 32007, 32011],
|
| 198 |
)
|
| 199 |
+
|
| 200 |
output_ids = gen_out.sequences
|
| 201 |
scores = gen_out.scores
|
| 202 |
+
generated_ids = output_ids[0, input_ids.shape[1] :]
|
| 203 |
+
|
| 204 |
for i, token_id in enumerate(generated_ids):
|
| 205 |
token_id = token_id.item()
|
| 206 |
token_text = local_tokenizer.decode([token_id])
|
| 207 |
all_tokens.append(token_text)
|
| 208 |
lp = torch.nn.functional.log_softmax(scores[i][0, :], dim=-1)
|
| 209 |
all_logprobs.append(lp[token_id].item())
|
| 210 |
+
|
| 211 |
logprobs = {"tokens": all_tokens, "token_logprobs": all_logprobs}
|
| 212 |
else:
|
| 213 |
# API Mode
|
|
|
|
| 216 |
"max_tokens": max_tokens,
|
| 217 |
"temperature": temperature,
|
| 218 |
"top_k": top_k,
|
| 219 |
+
"repetition_penalty": rep_penalty,
|
| 220 |
"logprobs": 1,
|
| 221 |
"echo": True,
|
| 222 |
}
|
|
|
|
| 241 |
color = "#e2e8f0" # Default
|
| 242 |
if lp is not None:
|
| 243 |
if prob > 0.8:
|
| 244 |
+
color = "#4ade80" # Vibrant Green
|
| 245 |
elif prob < 0.2:
|
| 246 |
+
color = "#f87171" # Soft Red
|
| 247 |
|
| 248 |
prob_pct = f"{(prob * 100):.1f}%" if lp is not None else "N/A"
|
| 249 |
+
logit_str = f"{lp:.4f}" if lp is not None else "N/A"
|
| 250 |
|
| 251 |
html += f'<span class="token" style="color: {color}">{display_token}'
|
| 252 |
+
html += f'<span class="token-tooltip">'
|
| 253 |
+
html += f'<span style="opacity: 0.7; font-size: 0.7rem">CONFIDENCE</span><br>'
|
| 254 |
+
html += f'<span style="font-size: 1.1rem; font-weight: 800; color: {color}">{prob_pct}</span><br>'
|
| 255 |
+
html += f'<hr style="margin: 5px 0; opacity: 0.1; border: none; border-top: 1px solid white;">'
|
| 256 |
+
html += f'<span style="opacity: 0.7; font-size: 0.7rem">LOGIT</span><br>'
|
| 257 |
+
html += f'<span style="font-family: monospace">{logit_str}</span>'
|
| 258 |
+
html += f'</span></span>'
|
| 259 |
|
| 260 |
html += "</div>"
|
| 261 |
return html
|
|
|
|
| 263 |
traceback.print_exc()
|
| 264 |
return f'<div class="token-container" style="color: #f87171">❌ Error: {str(e)}</div>'
|
| 265 |
|
| 266 |
+
with gr.Blocks(title="Echo-DSRN-Small-400m - Next Word Prediction") as demo:
|
| 267 |
+
with gr.Row():
|
| 268 |
+
with gr.Column(scale=4):
|
| 269 |
+
gr.Markdown(
|
| 270 |
+
"""
|
| 271 |
+
# 🌫️ Echo-DSRN-Small-400m - Next Word Prediction
|
| 272 |
+
### Responsible AI Development by ethicalabs.ai
|
| 273 |
+
|
| 274 |
+
**Echo-DSRN** is a high-performance **Dual-State Recurrent Network** inspired by the architecture in [Titans: Learning to Memorize at Test Time (Google Research)](https://arxiv.org/abs/2501.00663). It combines the efficiency of linear recurrence with optimized gating and a novel **Surprise Mechanism** to enable infinite context extrapolation.
|
| 275 |
+
|
| 276 |
+
<p style="margin-top: 15px; border-top: 1px solid rgba(255,255,255,0.1); padding-top: 10px;">
|
| 277 |
+
<span style="font-weight: bold; color: #f87171;">⚠️ Proprietary Technology Disclaimer</span><br>
|
| 278 |
+
<span style="font-style: italic; font-size: 0.9rem; color: #94a3b8;">This research is currently in an internal-only alpha phase. The model source code and weights can be released only when we have the capacity to scale up and expand and fine-tune and safety align it according to our strict responsibility standards.</span>
|
| 279 |
+
</p>
|
| 280 |
+
"""
|
| 281 |
+
)
|
| 282 |
+
with gr.Column(scale=1):
|
| 283 |
+
mode_str = "Standalone Mode" if local_model else f"Connected to API: {api_url}"
|
| 284 |
+
mode_cls = "mode-standalone" if local_model else "mode-api"
|
| 285 |
+
gr.HTML(f'<div class="mode-badge {mode_cls}">{mode_str}</div>')
|
| 286 |
|
| 287 |
+
if model_metadata:
|
| 288 |
+
gr.Markdown("### 📦 Model Architecture Details")
|
| 289 |
+
metadata_html = '<div class="metadata-card"><div class="metadata-grid">'
|
| 290 |
+
for key, value in model_metadata.items():
|
| 291 |
+
metadata_html += f'<div class="metadata-item">{key}<span class="metadata-value">{value}</span></div>'
|
| 292 |
+
metadata_html += "</div></div>"
|
| 293 |
+
gr.HTML(metadata_html)
|
| 294 |
|
| 295 |
gr.Markdown(
|
| 296 |
"""
|
| 297 |
+
Hover over generated tokens to visualize the model's confidence and raw logit values.
|
|
|
|
| 298 |
"""
|
| 299 |
)
|
| 300 |
|
|
|
|
| 312 |
)
|
| 313 |
with gr.Row():
|
| 314 |
top_k = gr.Slider(minimum=0, maximum=100, value=40, step=1, label="Top K")
|
| 315 |
+
rep_penalty = gr.Slider(
|
| 316 |
+
minimum=1.0, maximum=2.0, value=1.2, step=0.05, label="Repetition Penalty"
|
| 317 |
+
)
|
| 318 |
|
| 319 |
submit_btn = gr.Button("Generate with Probabilities", variant="primary")
|
| 320 |
|
|
|
|
| 323 |
output_html = gr.HTML(label="Visualized Output")
|
| 324 |
|
| 325 |
submit_btn.click(
|
| 326 |
+
generate_html,
|
| 327 |
+
inputs=[prompt, max_tokens, temp, top_k, rep_penalty],
|
| 328 |
+
outputs=output_html,
|
| 329 |
)
|
| 330 |
|
| 331 |
gr.Examples(
|
| 332 |
examples=[
|
| 333 |
+
["The capital of France is", 128, 0.1, 40, 1.2],
|
| 334 |
+
["Responsible AI development requires", 128, 0.7, 40, 1.2],
|
| 335 |
+
["To be or not to be, that is the", 128, 0.1, 40, 1.2],
|
| 336 |
],
|
| 337 |
+
inputs=[prompt, max_tokens, temp, top_k, rep_penalty],
|
| 338 |
)
|
| 339 |
|
| 340 |
return demo, css
|
|
|
|
| 346 |
"--api_url", type=str, default="http://localhost:5000", help="URL of the Echo-DSRN API"
|
| 347 |
)
|
| 348 |
parser.add_argument(
|
| 349 |
+
"--model_id",
|
| 350 |
+
type=str,
|
| 351 |
+
default=None,
|
| 352 |
+
help="Local path or HF ID to the model for Standalone mode",
|
| 353 |
)
|
| 354 |
parser.add_argument("--port", type=int, default=7860, help="Port to run the Gradio app on")
|
| 355 |
args = parser.parse_args()
|