luulinh90s/Tabular-LLM-Study-Debugging
0
1import uuid2from flask import Flask, render_template, request, redirect, url_for, send_from_directory3import json4import random5import os6import string7import logging8from datetime import datetime9from huggingface_hub import login, HfApi, hf_hub_download10 11# Set up logging12logging.basicConfig(level=logging.INFO,13 format='%(asctime)s - %(name)s - %(levelname)s - %(message)s',14 handlers=[15 logging.FileHandler("app.log"),16 logging.StreamHandler()17 ])18logger = logging.getLogger(__name__)19 20# Use the Hugging Face token from environment variables21hf_token = os.environ.get("HF_TOKEN")22if hf_token:23 login(token=hf_token)24else:25 logger.error("HF_TOKEN not found in environment variables")26 27app = Flask(__name__)28app.config['SECRET_KEY'] = 'supersecretkey' # Change this to a random secret key29 30# File-based session storage31SESSION_DIR = '/tmp/sessions'32os.makedirs(SESSION_DIR, exist_ok=True)33 34# Directories for visualizations35VISUALIZATION_DIRS = {36 "No-XAI": "htmls_NO_XAI_mod",37 "Dater": "htmls_DATER_mod2",38 "Chain-of-Table": "htmls_COT_mod",39 "Plan-of-SQLs": "htmls_POS_mod2"40}41 42def get_method_dir(method):43 if method == 'No-XAI':44 return 'NO_XAI'45 elif method == 'Dater':46 return 'DATER'47 elif method == 'Chain-of-Table':48 return 'COT'49 elif method == 'Plan-of-SQLs':50 return 'POS'51 else:52 return None53 54METHODS = ["No-XAI", "Dater", "Chain-of-Table", "Plan-of-SQLs"]55 56def generate_session_id():57 return str(uuid.uuid4())58 59def save_session_data(session_id, data):60 file_path = os.path.join(SESSION_DIR, f'{session_id}.json')61 with open(file_path, 'w') as f:62 json.dump(data, f)63 logger.info(f"Session data saved for session {session_id}")64 65def load_session_data(session_id):66 file_path = os.path.join(SESSION_DIR, f'{session_id}.json')67 if os.path.exists(file_path):68 with open(file_path, 'r') as f:69 return json.load(f)70 return None71 72def save_session_data_to_hf(session_id, data):73 try:74 username = data.get('username', 'unknown')75 seed = data.get('seed', 'unknown')76 start_time = data.get('start_time', datetime.now().isoformat())77 file_name = f'{username}_seed{seed}_{start_time}_{session_id}_session.json'78 file_name = "".join(c for c in file_name if c.isalnum() or c in ['_', '-', '.'])79 80 json_data = json.dumps(data, indent=4)81 temp_file_path = f"/tmp/{file_name}"82 with open(temp_file_path, 'w') as f:83 f.write(json_data)84 85 api = HfApi()86 repo_path = "session_data_debugging"87 88 api.upload_file(89 path_or_fileobj=temp_file_path,90 path_in_repo=f"{repo_path}/{file_name}",91 repo_id="luulinh90s/Tabular-LLM-Study-Data",92 repo_type="space",93 )94 os.remove(temp_file_path)95 logger.info(f"Session data saved for session {session_id} in Hugging Face Data Space")96 except Exception as e:97 logger.exception(f"Error saving session data for session {session_id}: {e}")98 99def load_samples():100 common_samples = []101 categories = ["TP", "TN", "FP", "FN"]102 103 for category in categories:104 files = set(os.listdir(f'htmls_NO_XAI_mod/{category}'))105 for method in ["Dater", "Chain-of-Table", "Plan-of-SQLs"]:106 method_dir = VISUALIZATION_DIRS[method]107 files &= set(os.listdir(f'{method_dir}/{category}'))108 109 for file in files:110 common_samples.append({'category': category, 'file': file})111 112 logger.info(f"Found {len(common_samples)} common samples across all methods")113 return common_samples114 115def select_balanced_samples(samples):116 try:117 # Separate samples into two groups118 tp_fp_samples = [s for s in samples if s['category'] in ['TP', 'TN']]119 tn_fn_samples = [s for s in samples if s['category'] in ['FP', 'FN']]120 121 # Check if we have enough samples in each group122 if len(tp_fp_samples) < 5 or len(tn_fn_samples) < 5:123 logger.warning(f"Not enough samples in each category. TP+FP: {len(tp_fp_samples)}, TN+FN: {len(tn_fn_samples)}")124 return samples if len(samples) <= 10 else random.sample(samples, 10)125 126 # Select 5 samples from each group127 selected_tp_fp = random.sample(tp_fp_samples, 5)128 selected_tn_fn = random.sample(tn_fn_samples, 5)129 130 # Combine and shuffle the selected samples131 selected_samples = selected_tp_fp + selected_tn_fn132 random.shuffle(selected_samples)133 134 logger.info(f"Selected 10 balanced samples: 5 from TP+FP, 5 from TN+FN")135 return selected_samples136 except Exception as e:137 logger.exception("Error selecting balanced samples")138 return []139 140@app.route('/')141def introduction():142 return render_template('introduction.html')143 144@app.route('/attribution')145def attribution():146 return render_template('attribution.html')147 148@app.route('/index', methods=['GET', 'POST'])149def index():150 if request.method == 'POST':151 username = request.form.get('username')152 seed = request.form.get('seed')153 method = request.form.get('method')154 if not username or not seed or not method:155 return render_template('index.html', error="Please fill in all fields and select a method.")156 if method not in ['Chain-of-Table', 'Plan-of-SQLs', 'Dater']:157 return render_template('index.html', error="Invalid method selected.")158 try:159 seed = int(seed)160 random.seed(seed)161 all_samples = load_samples()162 selected_samples = select_balanced_samples(all_samples)163 if len(selected_samples) == 0:164 return render_template('index.html', error="No common samples were found")165 start_time = datetime.now().isoformat()166 session_id = generate_session_id()167 session_data = {168 'username': username,169 'seed': str(seed),170 'method': method,171 'selected_samples': selected_samples,172 'current_index': 0,173 'responses': [],174 'start_time': start_time,175 'session_id': session_id176 }177 save_session_data(session_id, session_data)178 logger.info(f"Session data stored for user {username}, method {method}, session_id {session_id}")179 180 # Redirect to explanation for all methods181 return redirect(url_for('explanation', session_id=session_id))182 except Exception as e:183 logger.exception(f"Error in index route: {e}")184 return render_template('index.html', error="An error occurred. Please try again.")185 return render_template('index.html', show_no_xai=False)186 187@app.route('/explanation/<session_id>')188def explanation(session_id):189 session_data = load_session_data(session_id)190 if not session_data:191 logger.error(f"No session data found for session ID: {session_id}")192 return redirect(url_for('index'))193 194 method = session_data.get('method')195 if not method:196 logger.error(f"No method found in session data for session ID: {session_id}")197 return redirect(url_for('index'))198 199 if method == 'Chain-of-Table':200 return render_template('cot_intro.html', session_id=session_id)201 elif method == 'Plan-of-SQLs':202 return render_template('pos_intro.html', session_id=session_id)203 elif method == 'Dater':204 return render_template('dater_intro.html', session_id=session_id)205 else:206 logger.error(f"Invalid method '{method}' for session ID: {session_id}")207 return redirect(url_for('index'))208 209@app.route('/experiment/<session_id>', methods=['GET', 'POST'])210def experiment(session_id):211 try:212 session_data = load_session_data(session_id)213 if not session_data:214 return redirect(url_for('index'))215 216 selected_samples = session_data['selected_samples']217 method = session_data['method']218 current_index = session_data['current_index']219 220 if current_index >= len(selected_samples):221 return redirect(url_for('completed', session_id=session_id))222 223 sample = selected_samples[current_index]224 visualization_dir = VISUALIZATION_DIRS[method]225 visualization_path = f"{visualization_dir}/{sample['category']}/{sample['file']}"226 227 statement = """228Please note that in select row function, starting index is 0 for Chain-of-Table and 1 for Dater and Index * represents the selection for all rows.229 """230 231 return render_template('experiment.html',232 sample_id=current_index,233 statement=statement,234 visualization=url_for('send_visualization', filename=visualization_path),235 session_id=session_id,236 method=method)237 except Exception as e:238 logger.exception(f"An error occurred in the experiment route: {e}")239 return "An error occurred", 500240 241@app.route('/subjective/<session_id>', methods=['GET', 'POST'])242def subjective(session_id):243 if request.method == 'POST':244 understanding = request.form.get('understanding')245 246 session_data = load_session_data(session_id)247 if not session_data:248 logger.error(f"No session data found for session: {session_id}")249 return redirect(url_for('index'))250 251 session_data['subjective_feedback'] = understanding252 save_session_data(session_id, session_data)253 254 return redirect(url_for('completed', session_id=session_id))255 256 return render_template('subjective.html', session_id=session_id)257 258@app.route('/feedback', methods=['POST'])259def feedback():260 try:261 session_id = request.form['session_id']262 prediction = request.form['prediction']263 264 session_data = load_session_data(session_id)265 if not session_data:266 logger.error(f"No session data found for session: {session_id}")267 return redirect(url_for('index'))268 269 session_data['responses'].append({270 'sample_id': session_data['current_index'],271 'user_prediction': prediction272 })273 274 session_data['current_index'] += 1275 save_session_data(session_id, session_data)276 logger.info(f"Prediction saved for session {session_id}, sample {session_data['current_index'] - 1}")277 278 if session_data['current_index'] >= len(session_data['selected_samples']):279 return redirect(url_for('subjective', session_id=session_id))280 281 return redirect(url_for('experiment', session_id=session_id))282 except Exception as e:283 logger.exception(f"Error in feedback route: {e}")284 return "An error occurred", 500285 286@app.route('/completed/<session_id>')287def completed(session_id):288 try:289 session_data = load_session_data(session_id)290 if not session_data:291 logger.error(f"No session data found for session: {session_id}")292 return redirect(url_for('index'))293 294 session_data['end_time'] = datetime.now().isoformat()295 responses = session_data['responses']296 method = session_data['method']297 298 if method == "Chain-of-Table":299 json_file = 'Tabular_LLMs_human_study_vis_6_COT.json'300 elif method == "Plan-of-SQLs":301 json_file = 'Tabular_LLMs_human_study_vis_6_POS.json'302 elif method == "Dater":303 json_file = 'Tabular_LLMs_human_study_vis_6_DATER.json'304 elif method == "No-XAI":305 json_file = 'Tabular_LLMs_human_study_vis_6_NO_XAI.json'306 else:307 return "Invalid method", 400308 309 with open(json_file, 'r') as f:310 ground_truth = json.load(f)311 312 correct_predictions = 0313 true_predictions = 0314 false_predictions = 0315 316 for response in responses:317 sample_id = response['sample_id']318 user_prediction = response['user_prediction']319 visualization_file = session_data['selected_samples'][sample_id]['file']320 index = visualization_file.split('-')[1].split('.')[0]321 322 ground_truth_key = f"{get_method_dir(method)}_test-{index}.html"323 logger.info(f"ground_truth_key: {ground_truth_key}")324 325 if ground_truth_key in ground_truth:326 model_prediction = ground_truth[ground_truth_key]['prediction'].upper()327 ground_truth_label = ground_truth[ground_truth_key]['answer'].upper()328 329 correctness = "TRUE" if model_prediction.upper() == ground_truth_label.upper() else "FALSE"330 331 if user_prediction.upper() == correctness:332 correct_predictions += 1333 334 if user_prediction.upper() == "TRUE":335 true_predictions += 1336 elif user_prediction.upper() == "FALSE":337 false_predictions += 1338 else:339 logger.warning(f"Missing key in ground truth: {ground_truth_key}")340 341 accuracy = (correct_predictions / len(responses)) * 100 if responses else 0342 accuracy = round(accuracy, 2)343 344 true_percentage = (true_predictions / len(responses)) * 100 if len(responses) else 0345 false_percentage = (false_predictions / len(responses)) * 100 if len(responses) else 0346 347 true_percentage = round(true_percentage, 2)348 false_percentage = round(false_percentage, 2)349 350 session_data['accuracy'] = accuracy351 session_data['true_percentage'] = true_percentage352 session_data['false_percentage'] = false_percentage353 354 # Save all the data to Hugging Face at the end355 save_session_data_to_hf(session_id, session_data)356 357 # Remove the local session data file358 os.remove(os.path.join(SESSION_DIR, f'{session_id}.json'))359 360 return render_template('completed.html',361 accuracy=accuracy,362 true_percentage=true_percentage,363 false_percentage=false_percentage)364 except Exception as e:365 logger.exception(f"An error occurred in the completed route: {e}")366 return "An error occurred", 500367 368@app.route('/visualizations/<path:filename>')369def send_visualization(filename):370 logger.info(f"Attempting to serve file: {filename}")371 base_dir = os.getcwd()372 file_path = os.path.normpath(os.path.join(base_dir, filename))373 if not file_path.startswith(base_dir):374 return "Access denied", 403375 376 if not os.path.exists(file_path):377 return "File not found", 404378 379 directory = os.path.dirname(file_path)380 file_name = os.path.basename(file_path)381 logger.info(f"Serving file from directory: {directory}, filename: {file_name}")382 return send_from_directory(directory, file_name)383 384@app.route('/visualizations/<path:filename>')385def send_examples(filename):386 return send_from_directory('', filename)387 388if __name__ == "__main__":389 app.run(host="0.0.0.0", port=7860, debug=True)