Spaces:
Sleeping
Sleeping
Update app.py
Browse files
app.py
CHANGED
|
@@ -6,17 +6,13 @@ from transformers import AutoTokenizer, AutoModelForCausalLM
|
|
| 6 |
MODEL_ID = "PraneetNS/codesentinel-full"
|
| 7 |
TOKENIZER_ID = "PraneetNS/codesentinel-adapter"
|
| 8 |
|
|
|
|
|
|
|
| 9 |
tokenizer = AutoTokenizer.from_pretrained(
|
| 10 |
TOKENIZER_ID,
|
| 11 |
trust_remote_code=True,
|
| 12 |
)
|
| 13 |
|
| 14 |
-
model = AutoModelForCausalLM.from_pretrained(
|
| 15 |
-
MODEL_ID,
|
| 16 |
-
dtype=torch.float16,
|
| 17 |
-
device_map="auto",
|
| 18 |
-
trust_remote_code=True,
|
| 19 |
-
)
|
| 20 |
SYSTEM_PROMPT = """You are CodeSentinel.
|
| 21 |
|
| 22 |
You are an expert AI software engineer.
|
|
@@ -33,12 +29,10 @@ Capabilities:
|
|
| 33 |
Never generate malware or unsafe code.
|
| 34 |
"""
|
| 35 |
|
| 36 |
-
# Global model variable
|
| 37 |
model = None
|
| 38 |
|
| 39 |
|
| 40 |
def load_model():
|
| 41 |
-
"""Load the model only once after GPU allocation."""
|
| 42 |
global model
|
| 43 |
|
| 44 |
if model is None:
|
|
@@ -67,9 +61,30 @@ def generate(message, history):
|
|
| 67 |
}
|
| 68 |
]
|
| 69 |
|
|
|
|
| 70 |
if history:
|
| 71 |
-
for
|
| 72 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 73 |
|
| 74 |
messages.append(
|
| 75 |
{
|
|
@@ -82,13 +97,18 @@ def generate(message, history):
|
|
| 82 |
messages,
|
| 83 |
tokenize=True,
|
| 84 |
add_generation_prompt=True,
|
| 85 |
-
return_tensors="pt",
|
| 86 |
return_dict=True,
|
| 87 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 88 |
|
| 89 |
with torch.no_grad():
|
| 90 |
outputs = model.generate(
|
| 91 |
-
inputs,
|
| 92 |
max_new_tokens=512,
|
| 93 |
temperature=0.2,
|
| 94 |
top_p=0.95,
|
|
@@ -97,8 +117,10 @@ def generate(message, history):
|
|
| 97 |
pad_token_id=tokenizer.eos_token_id,
|
| 98 |
)
|
| 99 |
|
|
|
|
|
|
|
| 100 |
response = tokenizer.decode(
|
| 101 |
-
|
| 102 |
skip_special_tokens=True,
|
| 103 |
)
|
| 104 |
|
|
@@ -114,7 +136,7 @@ demo = gr.ChatInterface(
|
|
| 114 |
"Explain this C++ function.",
|
| 115 |
"Optimize this SQL query.",
|
| 116 |
"Write a secure FastAPI login API.",
|
| 117 |
-
"Review this Java code for security vulnerabilities."
|
| 118 |
],
|
| 119 |
)
|
| 120 |
|
|
|
|
| 6 |
MODEL_ID = "PraneetNS/codesentinel-full"
|
| 7 |
TOKENIZER_ID = "PraneetNS/codesentinel-adapter"
|
| 8 |
|
| 9 |
+
print("Loading tokenizer...")
|
| 10 |
+
|
| 11 |
tokenizer = AutoTokenizer.from_pretrained(
|
| 12 |
TOKENIZER_ID,
|
| 13 |
trust_remote_code=True,
|
| 14 |
)
|
| 15 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 16 |
SYSTEM_PROMPT = """You are CodeSentinel.
|
| 17 |
|
| 18 |
You are an expert AI software engineer.
|
|
|
|
| 29 |
Never generate malware or unsafe code.
|
| 30 |
"""
|
| 31 |
|
|
|
|
| 32 |
model = None
|
| 33 |
|
| 34 |
|
| 35 |
def load_model():
|
|
|
|
| 36 |
global model
|
| 37 |
|
| 38 |
if model is None:
|
|
|
|
| 61 |
}
|
| 62 |
]
|
| 63 |
|
| 64 |
+
# Convert Gradio history
|
| 65 |
if history:
|
| 66 |
+
for item in history:
|
| 67 |
+
if isinstance(item, (list, tuple)) and len(item) == 2:
|
| 68 |
+
user_msg, assistant_msg = item
|
| 69 |
+
|
| 70 |
+
if user_msg:
|
| 71 |
+
messages.append(
|
| 72 |
+
{
|
| 73 |
+
"role": "user",
|
| 74 |
+
"content": user_msg,
|
| 75 |
+
}
|
| 76 |
+
)
|
| 77 |
+
|
| 78 |
+
if assistant_msg:
|
| 79 |
+
messages.append(
|
| 80 |
+
{
|
| 81 |
+
"role": "assistant",
|
| 82 |
+
"content": assistant_msg,
|
| 83 |
+
}
|
| 84 |
+
)
|
| 85 |
+
|
| 86 |
+
elif isinstance(item, dict):
|
| 87 |
+
messages.append(item)
|
| 88 |
|
| 89 |
messages.append(
|
| 90 |
{
|
|
|
|
| 97 |
messages,
|
| 98 |
tokenize=True,
|
| 99 |
add_generation_prompt=True,
|
|
|
|
| 100 |
return_dict=True,
|
| 101 |
+
return_tensors="pt",
|
| 102 |
+
)
|
| 103 |
+
|
| 104 |
+
inputs = {
|
| 105 |
+
k: v.to(model.device)
|
| 106 |
+
for k, v in inputs.items()
|
| 107 |
+
}
|
| 108 |
|
| 109 |
with torch.no_grad():
|
| 110 |
outputs = model.generate(
|
| 111 |
+
**inputs,
|
| 112 |
max_new_tokens=512,
|
| 113 |
temperature=0.2,
|
| 114 |
top_p=0.95,
|
|
|
|
| 117 |
pad_token_id=tokenizer.eos_token_id,
|
| 118 |
)
|
| 119 |
|
| 120 |
+
generated_tokens = outputs[0][inputs["input_ids"].shape[-1]:]
|
| 121 |
+
|
| 122 |
response = tokenizer.decode(
|
| 123 |
+
generated_tokens,
|
| 124 |
skip_special_tokens=True,
|
| 125 |
)
|
| 126 |
|
|
|
|
| 136 |
"Explain this C++ function.",
|
| 137 |
"Optimize this SQL query.",
|
| 138 |
"Write a secure FastAPI login API.",
|
| 139 |
+
"Review this Java code for security vulnerabilities.",
|
| 140 |
],
|
| 141 |
)
|
| 142 |
|