Team Ai
Apppublic

oookiku/route-explainer

sourceHugging Faceotherupdated 3y agoView on Hugging Face
1likes
route_explainer.py293 linesDownload Raw Back to models
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