oookiku/route-explainer
1
1# templates2import numpy as np3import streamlit as st4from typing import Dict, List5from models.prompts.identify_question import Template4IdentifyQuestion6from models.prompts.generate_explanation import Template4GenerateExplanation7from langchain.callbacks.base import BaseCallbackHandler8from langchain.schema import AIMessage9import utils.util_app as util_app10 11class StreamingChatCallbackHandler(BaseCallbackHandler):12 def __init__(self):13 pass14 15 def on_llm_start(self, *args, **kwargs):16 self.container = st.empty()17 self.text = ""18 19 def on_llm_new_token(self, token: str, *args, **kwargs):20 self.text += token21 self.container.markdown(22 body=self.text,23 unsafe_allow_html=False,24 )25 26 def on_llm_end(self, response: str, *args, **kwargs):27 self.container.markdown(28 body=response.generations[0][0].text,29 unsafe_allow_html=False,30 )31 32class RouteExplainer():33 template_identify_question = Template4IdentifyQuestion()34 template_generate_explanation = Template4GenerateExplanation()35 36 def __init__(self,37 llm,38 cf_generator, 39 classifier) -> None:40 assert cf_generator.problem == classifier.problem, "Problem type of cf_generator and predictor should coincide!"41 self.coord_dim = 242 self.problem = cf_generator.problem43 self.cf_generator = cf_generator44 self.classifier = classifier45 self.actual_route = None46 self.cf_route = None47 # templates48 self.question_extractor = self.template_identify_question.sandwiches(llm)49 self.explanation_generator = self.template_generate_explanation.sandwiches(llm)50 51 #----------------52 # whole pipeline53 #----------------54 def generate_explanation(self, 55 tour_list,56 whynot_question: str,57 actual_routes: list,58 actual_labels: list,59 node_feats: dict,60 dist_matrix: np.array) -> str:61 #--------------------------------62 # define why & why-not questions63 #--------------------------------64 route_info_text = self.get_route_info_text(tour_list, actual_routes)65 inputs = self.question_extractor.invoke({66 "whynot_question": whynot_question,67 "route_info": route_info_text68 })69 util_app.stream_words(inputs["summary"] + " " + inputs["intent"])70 st.session_state.chat_history.append(AIMessage(content=inputs["summary"] + inputs["intent"]))71 if not inputs["success"]:72 return ""73 74 #----------------------75 # validate the CF edge76 #----------------------77 is_cf_edge_feasible, reason = self.validate_cf_edge(node_feats,78 dist_matrix,79 actual_routes[0],80 inputs["cf_step"],81 inputs["cf_visit"]-1)82 # exception83 if not is_cf_edge_feasible:84 util_app.stream_words(reason)85 return reason86 87 #---------------------88 # generate a cf route89 #---------------------90 cf_routes = self.cf_generator(actual_routes,91 vehicle_id=0,92 cf_step=inputs["cf_step"],93 cf_next_node_id=inputs["cf_visit"]-1,94 node_feats=node_feats,95 dist_matrix=dist_matrix)96 st.session_state.generated_cf_route = True97 st.session_state.close_chat = True98 st.session_state.cf_step = inputs["cf_step"]99 100 #--------------------------------------101 # classify the intentions of each edge102 #--------------------------------------103 cf_labels = self.classifier(self.classifier.get_inputs(cf_routes,104 0,105 node_feats,106 dist_matrix))107 st.session_state.cf_routes = cf_routes108 st.session_state.cf_labels = cf_labels109 110 #-------------------------------------111 # generate a constrastive explanation112 #-------------------------------------113 comparison_results = self.get_comparison_results(question_summary=inputs["summary"],114 tour_list=tour_list,115 actual_routes=actual_routes,116 actual_labels=actual_labels,117 cf_routes=cf_routes,118 cf_labels=cf_labels,119 cf_step=inputs["cf_step"])120 121 explanation = self.explanation_generator.invoke({122 "comparison_results": comparison_results,123 "intent": inputs["intent"]124 }, config={"callbacks": [StreamingChatCallbackHandler()]})125 126 return explanation127 128 #-------------------------129 # for exctracting inputs130 #-------------------------131 def get_route_info_text(self, tour_list, routes) -> str:132 route_info = ""133 # nodes134 route_info += "Nodes(node id, name): "135 for i, destination in enumerate(tour_list):136 if i != len(tour_list) - 1:137 route_info += f"({i+1}, {destination['name']}), "138 else:139 route_info += f"({i+1}, {destination['name']})\n"140 141 # routes142 route_info += "Route: "143 for i, node_id in enumerate(routes[0]):144 if i == 0:145 route_info += f"{tour_list[node_id]['name']} "146 else:147 route_info += f"> (step {i}) > {tour_list[node_id]['name']})"148 if i == len(routes[0]) - 1:149 route_info += "\n"150 else:151 route_info += " "152 return route_info153 154 #--------------------------155 # for validating a CF edge156 #--------------------------157 def validate_cf_edge(self,158 node_feats: Dict[str, np.array],159 dist_matrix: np.array,160 route: List[int],161 cf_step: int,162 cf_visit: int) -> bool:163 # calc current time164 curr_time = node_feats["time_window"][route[0]][0] # start point's open time165 for step in range(1, cf_step):166 curr_node_id = route[step-1]167 next_node_id = route[step]168 curr_time += node_feats["service_time"][curr_node_id] + dist_matrix[curr_node_id][next_node_id]169 curr_time = max(curr_time, node_feats["time_window"][next_node_id][0]) # waiting170 171 # validate the cf edge172 curr_node_id = route[cf_step-1]173 next_node_id = cf_visit174 next_node_close_time = node_feats["time_window"][next_node_id][1] 175 arrival_time = curr_time + node_feats["service_time"][curr_node_id] + dist_matrix[curr_node_id][next_node_id]176 if next_node_close_time < arrival_time:177 exceed_time = (arrival_time - next_node_close_time)178 return False, f"Oops, your CF edge is infeasible because it does not meet the destination's close time by {util_app.add_time_unit(exceed_time)}."179 else:180 return True, "The CF edge is feasible!"181 182 #-------------------------------183 # for generating an explanation184 #-------------------------------185 def get_comparison_results(self,186 tour_list,187 question_summary,188 actual_routes: List[List[int]],189 actual_labels: List[List[int]],190 cf_routes: List[List[int]],191 cf_labels: List[List[int]],192 cf_step: int) -> str:193 comparison_results = "Question:\n" + question_summary + "\n"194 comparison_results += "Actual route:\n" + \195 self.get_route_info(tour_list, actual_routes[0], actual_labels[0], cf_step-1, "actual") + \196 self.get_representative_values(actual_routes[0], actual_labels[0], cf_step-1, "actual")197 comparison_results += "CF route:\n" + \198 self.get_route_info(tour_list, cf_routes[0], cf_labels[0], cf_step-1, "CF") + \199 self.get_representative_values(cf_routes[0], cf_labels[0], cf_step-1, "CF")200 comparison_results += "Difference between two routes:\n" + self.get_diff(cf_step-1, actual_routes[0], cf_routes[0])201 comparison_results += "Planed desination information:\n" + self.get_node_info()202 return comparison_results203 204 def get_route_info(self,205 tour_list,206 route: List[int],207 label: List[int], 208 ex_step: int, 209 type: str) -> str:210 def get_labelname(label_number):211 return "route_len" if label_number == 0 else "time_window"212 route_info = "- route: "213 for i, node_id in enumerate(route):214 if i == ex_step and i != len(route) - 1:215 if type == "actual":216 edge_label = {get_labelname(label[i])}217 else:218 edge_label = "user_preference"219 route_info += f"{tour_list[node_id]['name']} > ({type} edge: {edge_label}) > "220 elif i != len(route) - 1:221 route_info += f"{tour_list[node_id]['name']} > ({get_labelname(label[i])}) > "222 else:223 route_info += f"{tour_list[node_id]['name']}\n"224 return route_info225 226 def get_representative_values(self, route, labels, ex_step, type) -> str:227 time_window_ratio = self.get_intention_ratio(1, labels, ex_step) * 100228 route_len_ratio = self.get_intention_ratio(0, labels, ex_step) * 100229 return f"- short-term effect (immediate travel time): {self.get_immediate_state(route, ex_step)//60} minutes\n- long-term effect (total travel time): {self.get_route_length(route)//60} minutes\n- missed nodes: {self.get_infeasible_node_name(route)}\n- edge-intention ratio after the {type} edge: time_window {time_window_ratio: .1f}%, route_len {route_len_ratio: .1f}%"230 231 def get_immediate_state(self, route, ex_step) -> str:232 return st.session_state.dist_matrix[route[ex_step]][route[ex_step+1]]233 234 def get_route_length(self, route) -> float:235 route_length = 0.0236 for i in range(len(route)-1):237 route_length += st.session_state.dist_matrix[route[i]][route[i+1]]238 return route_length239 240 def get_infeasible_nodes(self, route) -> int:241 return len(route) - (len(st.session_state.dist_matrix) - 1)242 243 def get_infeasible_node_name(self, route) -> str:244 if len(route) == len(st.session_state.dist_matrix) - 1:245 return "none"246 else:247 num_nodes = np.arange(len(st.session_state.dist_matrix))248 for node_id in route:249 num_nodes = num_nodes[num_nodes != node_id]250 return ",".join([st.session_state.tour_list[node_id]["name"] for node_id in num_nodes])251 252 def get_intention_ratio(self, 253 intention: int, 254 labels: List[int], 255 ex_step: int) -> float:256 np_labels = np.array(labels)257 return np.sum(np_labels[ex_step:] == intention) / len(labels[ex_step:])258 259 def get_diff(self, ex_step, actual_route, cf_route) -> str:260 def get_str(effect: float):261 long_effect_str = "The actual route increases it by" if effect > 0 else "The actual route reduces it by"262 long_effect_str += util_app.add_time_unit(abs(effect))263 return long_effect_str264 265 def get_str2(num_nodes: int, num_missed_nodes):266 if num_nodes < 0:267 num_nodes_str = f"The actual route visits {abs(num_nodes)} more nodes" 268 elif num_nodes == 0:269 if num_missed_nodes == 0:270 num_nodes_str = f"Both routes missed no node,"271 else:272 num_nodes_str = f"Both routes missed the same number of nodes ({abs(num_missed_nodes)} node(s))"273 else:274 num_nodes_str = f"The actual route visits {abs(num_nodes)} less nodes" 275 return num_nodes_str276 277 # short/long-term effects278 short_effect = self.get_immediate_state(actual_route, ex_step) - self.get_immediate_state(cf_route, ex_step)279 long_effect = self.get_route_length(actual_route) - self.get_route_length(cf_route)280 short_effect_str = get_str(short_effect)281 long_effect_str = get_str(long_effect)282 283 # missed nodes284 missed_nodes = self.get_infeasible_nodes(actual_route) - self.get_infeasible_nodes(cf_route)285 missed_nodes_str = get_str2(missed_nodes, self.get_infeasible_nodes(actual_route))286 287 return f"- short-term effect: {short_effect_str}\n - long-term effect: {long_effect_str}\n- missed nodes: {missed_nodes_str}\n"288 289 def get_node_info(self) -> str:290 node_info = ""291 for i in range(len(st.session_state.df_tour)):292 node_info += f"- {st.session_state.df_tour['destination'][i]}: {st.session_state.df_tour['remarks'][i]}\n"293 return node_info