FosterSystemsDatabase commited on
Commit
b1b854b
·
verified ·
1 Parent(s): 2bbc234

Hopefully fix again!

Browse files
Files changed (1) hide show
  1. app.py +13 -16
app.py CHANGED
@@ -22,7 +22,7 @@ model, tokenizer = FastLanguageModel.from_pretrained(
22
  dtype=None,
23
  load_in_4bit=True
24
  )
25
- adapter_path = "FosterSystemsDatabase/model" # Path to LoRA adapter on Hugging Face
26
  model = PeftModel.from_pretrained(model, adapter_path)
27
 
28
  # Load CSV data
@@ -30,17 +30,14 @@ file_path = 'Clean Missouri Data.csv'
30
  df = pd.read_csv(file_path, encoding='MacRoman')
31
 
32
  def search_relevant_policies(query, df, top_n=10, max_chars=40000):
33
- # Convert policies into a TF-IDF matrix
34
  tfidf = TfidfVectorizer(stop_words='english')
35
  tfidf_matrix = tfidf.fit_transform(df['Content'])
36
  query_vector = tfidf.transform([query])
37
  cosine_sim = cosine_similarity(query_vector, tfidf_matrix).flatten()
38
-
39
- # Get the top N relevant policies
40
  top_indices = cosine_sim.argsort()[-top_n:][::-1]
41
  relevant_policies = df.iloc[top_indices].copy()
42
 
43
- # Ensure total text is capped at max_chars
44
  char_count = 0
45
  valid_indices = []
46
  for idx, row in relevant_policies.iterrows():
@@ -49,7 +46,7 @@ def search_relevant_policies(query, df, top_n=10, max_chars=40000):
49
  break
50
  char_count += content_length
51
  valid_indices.append(idx)
52
-
53
  truncated_policies = relevant_policies.loc[valid_indices]
54
  return truncated_policies
55
 
@@ -65,12 +62,12 @@ def process_query(query, tokenizer):
65
  relevant_policies = search_relevant_policies(query, df)
66
  formatted_policies = [row['Content'] for _, row in relevant_policies.iterrows()]
67
  relevant_policy_text = "\n\n".join(formatted_policies)
68
-
69
  messages_with_relevant_policies = [
70
  {"role": "system", "content": relevant_policy_text},
71
  {"role": "user", "content": query},
72
  ]
73
-
74
  tokenizer = get_chat_template(tokenizer, chat_template="llama-3.1")
75
  inputs = tokenizer.apply_chat_template(
76
  messages_with_relevant_policies,
@@ -78,20 +75,19 @@ def process_query(query, tokenizer):
78
  add_generation_prompt=True,
79
  return_tensors="pt"
80
  ).to("cuda")
81
-
82
  FastLanguageModel.for_inference(model)
83
  outputs = model.generate(input_ids=inputs, max_new_tokens=512, use_cache=True, temperature=0.7, min_p=0.1)
84
  generated_response = tokenizer.batch_decode(outputs, skip_special_tokens=True)[0]
85
  response = get_content_after_query(generated_response, query)
86
-
87
- # Rank the top policy using SBERT
88
  model_sbert = SentenceTransformer('all-MiniLM-L6-v2')
89
  response_embedding = model_sbert.encode(generated_response, convert_to_tensor=True)
90
  policy_embeddings = model_sbert.encode(relevant_policies['Content'].tolist(), convert_to_tensor=True)
91
  cosine_similarities = util.cos_sim(response_embedding, policy_embeddings).flatten()
92
  most_relevant_index = cosine_similarities.argmax().item()
93
  most_relevant_link = relevant_policies.iloc[most_relevant_index]['Link to Content']
94
-
95
  return {
96
  "response": response,
97
  "most_relevant_link": most_relevant_link
@@ -115,9 +111,9 @@ def greet(query):
115
  def choose_preference(name, output1, output2, preference, query, broken):
116
  if not name:
117
  return "Please enter your name before submitting."
118
-
119
  broken_flag = "Yes" if broken else "No"
120
-
121
  if preference == "Output 1":
122
  new_row = [query, output1, output2, name, broken_flag]
123
  spreadsheet.append_row(new_row)
@@ -133,18 +129,19 @@ def choose_preference(name, output1, output2, preference, query, broken):
133
  with gr.Blocks() as demo:
134
  name_input = gr.Textbox(label="Enter your name")
135
  query_input = gr.Textbox(label="Enter your query")
136
- # "Generate Outputs" button is now placed right after the query input
137
  generate_button = gr.Button("Generate Outputs")
138
  output_1 = gr.Textbox(label="Output 1", interactive=False)
139
  output_2 = gr.Textbox(label="Output 2", interactive=False)
140
  preference = gr.Radio(["Output 1", "Output 2"], label="Choose your preferred output")
141
  broken_flag = gr.Checkbox(label="Mark as Broken Response")
 
142
  submit_button = gr.Button("Submit Preference")
143
-
144
  generate_button.click(greet, inputs=query_input, outputs=[output_1, output_2])
145
  submit_button.click(
146
  choose_preference,
147
  inputs=[name_input, output_1, output_2, preference, query_input, broken_flag],
 
148
  ).then(
149
  fn=lambda: ("", "", "", "", "", False, ""),
150
  inputs=[],
 
22
  dtype=None,
23
  load_in_4bit=True
24
  )
25
+ adapter_path = "FosterSystemsDatabase/model"
26
  model = PeftModel.from_pretrained(model, adapter_path)
27
 
28
  # Load CSV data
 
30
  df = pd.read_csv(file_path, encoding='MacRoman')
31
 
32
  def search_relevant_policies(query, df, top_n=10, max_chars=40000):
 
33
  tfidf = TfidfVectorizer(stop_words='english')
34
  tfidf_matrix = tfidf.fit_transform(df['Content'])
35
  query_vector = tfidf.transform([query])
36
  cosine_sim = cosine_similarity(query_vector, tfidf_matrix).flatten()
37
+
 
38
  top_indices = cosine_sim.argsort()[-top_n:][::-1]
39
  relevant_policies = df.iloc[top_indices].copy()
40
 
 
41
  char_count = 0
42
  valid_indices = []
43
  for idx, row in relevant_policies.iterrows():
 
46
  break
47
  char_count += content_length
48
  valid_indices.append(idx)
49
+
50
  truncated_policies = relevant_policies.loc[valid_indices]
51
  return truncated_policies
52
 
 
62
  relevant_policies = search_relevant_policies(query, df)
63
  formatted_policies = [row['Content'] for _, row in relevant_policies.iterrows()]
64
  relevant_policy_text = "\n\n".join(formatted_policies)
65
+
66
  messages_with_relevant_policies = [
67
  {"role": "system", "content": relevant_policy_text},
68
  {"role": "user", "content": query},
69
  ]
70
+
71
  tokenizer = get_chat_template(tokenizer, chat_template="llama-3.1")
72
  inputs = tokenizer.apply_chat_template(
73
  messages_with_relevant_policies,
 
75
  add_generation_prompt=True,
76
  return_tensors="pt"
77
  ).to("cuda")
78
+
79
  FastLanguageModel.for_inference(model)
80
  outputs = model.generate(input_ids=inputs, max_new_tokens=512, use_cache=True, temperature=0.7, min_p=0.1)
81
  generated_response = tokenizer.batch_decode(outputs, skip_special_tokens=True)[0]
82
  response = get_content_after_query(generated_response, query)
83
+
 
84
  model_sbert = SentenceTransformer('all-MiniLM-L6-v2')
85
  response_embedding = model_sbert.encode(generated_response, convert_to_tensor=True)
86
  policy_embeddings = model_sbert.encode(relevant_policies['Content'].tolist(), convert_to_tensor=True)
87
  cosine_similarities = util.cos_sim(response_embedding, policy_embeddings).flatten()
88
  most_relevant_index = cosine_similarities.argmax().item()
89
  most_relevant_link = relevant_policies.iloc[most_relevant_index]['Link to Content']
90
+
91
  return {
92
  "response": response,
93
  "most_relevant_link": most_relevant_link
 
111
  def choose_preference(name, output1, output2, preference, query, broken):
112
  if not name:
113
  return "Please enter your name before submitting."
114
+
115
  broken_flag = "Yes" if broken else "No"
116
+
117
  if preference == "Output 1":
118
  new_row = [query, output1, output2, name, broken_flag]
119
  spreadsheet.append_row(new_row)
 
129
  with gr.Blocks() as demo:
130
  name_input = gr.Textbox(label="Enter your name")
131
  query_input = gr.Textbox(label="Enter your query")
 
132
  generate_button = gr.Button("Generate Outputs")
133
  output_1 = gr.Textbox(label="Output 1", interactive=False)
134
  output_2 = gr.Textbox(label="Output 2", interactive=False)
135
  preference = gr.Radio(["Output 1", "Output 2"], label="Choose your preferred output")
136
  broken_flag = gr.Checkbox(label="Mark as Broken Response")
137
+ preference_result = gr.Textbox(label="Preference Result", interactive=False)
138
  submit_button = gr.Button("Submit Preference")
139
+
140
  generate_button.click(greet, inputs=query_input, outputs=[output_1, output_2])
141
  submit_button.click(
142
  choose_preference,
143
  inputs=[name_input, output_1, output_2, preference, query_input, broken_flag],
144
+ outputs=preference_result
145
  ).then(
146
  fn=lambda: ("", "", "", "", "", False, ""),
147
  inputs=[],