PraneetNS commited on
Commit
c1b072e
·
verified ·
1 Parent(s): 3780be1

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +37 -15
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 msg in history:
72
- messages.append(msg)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
- ).to(model.device)
 
 
 
 
 
 
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
- outputs[0][inputs.shape[-1]:],
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