luulinh90s/Tabular-LLM-Study-Preference-Ranking
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_download10from statistics import mean11 12# Set up logging13logging.basicConfig(level=logging.INFO,14 format='%(asctime)s - %(name)s - %(levelname)s - %(message)s',15 handlers=[16 logging.FileHandler("app.log"),17 logging.StreamHandler()18 ])19logger = logging.getLogger(__name__)20 21# Use the Hugging Face token from environment variables22hf_token = os.environ.get("HF_TOKEN")23if hf_token:24 login(token=hf_token)25else:26 logger.error("HF_TOKEN not found in environment variables")27 28app = Flask(__name__)29app.config['SECRET_KEY'] = 'supersecretkey'30 31# File-based session storage32SESSION_DIR = '/tmp/sessions'33os.makedirs(SESSION_DIR, exist_ok=True)34 35# Update visualization directories for the 4 methods36VISUALIZATION_DIRS = {37 "Text2SQL": "htmls_Text2SQL",38 "Dater": "htmls_DATER_mod2",39 "Chain-of-Table": "htmls_COT_mod",40 "Plan-of-SQLs": "htmls_POS_mod2"41}42 43 44# Update method directory mapping45def get_method_dir(method):46 method_mapping = {47 'Text2SQL': 'Text2SQL',48 'Dater': 'DATER',49 'Chain-of-Table': 'COT',50 'Plan-of-SQLs': 'POS'51 }52 return method_mapping.get(method)53 54 55# Update methods list to only include the 4 methods we want to rank56METHODS = ["Text2SQL", "Dater", "Chain-of-Table", "Plan-of-SQLs"]57 58 59def generate_session_id():60 return str(uuid.uuid4())61 62 63def save_session_data(session_id, data):64 file_path = os.path.join(SESSION_DIR, f'{session_id}.json')65 with open(file_path, 'w') as f:66 json.dump(data, f)67 logger.info(f"Session data saved for session {session_id}")68 69 70def load_session_data(session_id):71 file_path = os.path.join(SESSION_DIR, f'{session_id}.json')72 if os.path.exists(file_path):73 with open(file_path, 'r') as f:74 return json.load(f)75 return None76 77 78def save_session_data_to_hf(session_id, data):79 try:80 username = data.get('username', 'unknown')81 seed = data.get('seed', 'unknown')82 start_time = data.get('start_time', datetime.now().isoformat())83 file_name = f'{username}_seed{seed}_{start_time}_{session_id}_session.json'84 file_name = "".join(c for c in file_name if c.isalnum() or c in ['_', '-', '.'])85 86 json_data = json.dumps(data, indent=4)87 temp_file_path = f"/tmp/{file_name}"88 with open(temp_file_path, 'w') as f:89 f.write(json_data)90 91 api = HfApi()92 repo_path = "session_data_preference_ranking"93 94 api.upload_file(95 path_or_fileobj=temp_file_path,96 path_in_repo=f"{repo_path}/{file_name}",97 repo_id="luulinh90s/Tabular-LLM-Study-Data",98 repo_type="space",99 )100 os.remove(temp_file_path)101 logger.info(f"Session data saved for session {session_id} in Hugging Face Data Space")102 except Exception as e:103 logger.exception(f"Error saving session data for session {session_id}: {e}")104 105 106def load_samples_for_all_methods(metadata_files):107 samples_by_method = {}108 common_samples = []109 110 # First, load all samples for each method111 for method in METHODS:112 method_samples = []113 categories = ["TP", "TN", "FP", "FN"]114 115 for category in categories:116 method_dir = VISUALIZATION_DIRS[method]117 try:118 files = set(os.listdir(f'{method_dir}/{category}'))119 120 for file in files:121 index = file.split('-')[1].split('.')[0]122 metadata_key = f"{get_method_dir(method)}_test-{index}.html"123 124 # Get metadata for this sample125 sample_metadata = metadata_files[method].get(metadata_key, {})126 127 method_samples.append({128 'category': category,129 'file': file,130 'metadata': sample_metadata131 })132 except Exception as e:133 logger.error(f"Error loading samples for method {method}, category {category}: {e}")134 135 samples_by_method[method] = method_samples136 137 # Find common samples across all methods138 file_sets = []139 for method, samples in samples_by_method.items():140 file_set = {s['file'] for s in samples}141 file_sets.append(file_set)142 143 common_files = set.intersection(*file_sets)144 145 # Create groups of samples that exist across all methods146 for file_name in common_files:147 sample_group = {}148 for method in METHODS:149 sample = next((s for s in samples_by_method[method] if s['file'] == file_name), None)150 if sample:151 sample_group[method] = sample152 if len(sample_group) == len(METHODS):153 common_samples.append(sample_group)154 155 return common_samples156 157 158def select_balanced_samples(samples):159 try:160 # Get the category from any method (they should all be the same)161 sample_categories = [(s, next(iter(s.values()))['category']) for s in samples]162 163 # Separate samples into two groups164 tp_fp_samples = [s for s, cat in sample_categories if cat in ['TP', 'FP']]165 tn_fn_samples = [s for s, cat in sample_categories if cat in ['TN', 'FN']]166 167 # Select balanced samples168 if len(tp_fp_samples) >= 5 and len(tn_fn_samples) >= 5:169 selected_tp_fp = random.sample(tp_fp_samples, 5)170 selected_tn_fn = random.sample(tn_fn_samples, 5)171 selected_samples = selected_tp_fp + selected_tn_fn172 random.shuffle(selected_samples)173 else:174 logger.warning(175 f"Not enough samples for balanced selection. TP+FP: {len(tp_fp_samples)}, TN+FN: {len(tn_fn_samples)}")176 selected_samples = random.sample(samples, min(10, len(samples)))177 178 return selected_samples179 except Exception as e:180 logger.exception("Error selecting balanced samples")181 return []182 183 184@app.route('/')185def root():186 return redirect(url_for('consent'))187 188 189@app.route('/consent', methods=['GET', 'POST'])190def consent():191 if request.method == 'POST':192 return redirect(url_for('introduction'))193 return render_template('consent.html')194 195 196@app.route('/introduction')197def introduction():198 return render_template('introduction.html')199 200 201@app.route('/attribution')202def attribution():203 return render_template('attribution.html')204 205 206@app.route('/index', methods=['GET', 'POST'])207def index():208 if request.method == 'POST':209 username = request.form.get('username')210 seed = request.form.get('seed')211 212 if not username or not seed:213 return render_template('index.html', error="Please fill in all fields.")214 215 try:216 seed = int(seed)217 random.seed(seed)218 219 # Load metadata for all methods220 metadata_files = {}221 for method in METHODS:222 json_file = f'Tabular_LLMs_human_study_vis_6_{get_method_dir(method)}.json'223 with open(json_file, 'r') as f:224 metadata_files[method] = json.load(f)225 226 # Load and select samples227 all_samples = load_samples_for_all_methods(metadata_files)228 selected_samples = select_balanced_samples(all_samples)229 230 if len(selected_samples) == 0:231 return render_template('index.html', error="No common samples were found")232 233 # Create session234 session_id = generate_session_id()235 session_data = {236 'username': username,237 'seed': str(seed),238 'selected_samples': selected_samples,239 'current_index': 0,240 'responses': [],241 'start_time': datetime.now().isoformat(),242 'session_id': session_id243 }244 save_session_data(session_id, session_data)245 246 return redirect(url_for('experiment', session_id=session_id))247 248 except Exception as e:249 logger.exception(f"Error in index route: {e}")250 return render_template('index.html', error="An error occurred. Please try again.")251 252 return render_template('index.html')253 254 255@app.route('/experiment/<session_id>', methods=['GET', 'POST'])256def experiment(session_id):257 try:258 session_data = load_session_data(session_id)259 if not session_data:260 return redirect(url_for('index'))261 262 selected_samples = session_data['selected_samples']263 current_index = session_data['current_index']264 265 if current_index >= len(selected_samples):266 return redirect(url_for('completed', session_id=session_id))267 268 if request.method == 'POST':269 # Validate and save rankings270 rankings = {method: int(request.form.get(method)) for method in METHODS}271 272 if not all(1 <= rank <= 4 for rank in rankings.values()):273 return "Invalid rankings. Please use numbers 1-4.", 400274 if len(set(rankings.values())) != 4:275 return "Each method must have a unique rank.", 400276 277 session_data['responses'].append({278 'sample_id': current_index,279 'rankings': rankings280 })281 session_data['current_index'] += 1282 save_session_data(session_id, session_data)283 return redirect(url_for('experiment', session_id=session_id))284 285 # Get current sample group and prepare visualizations286 sample_group = selected_samples[current_index]287 visualizations = {288 method: url_for('send_visualization',289 filename=f"{VISUALIZATION_DIRS[method]}/{sample['category']}/{sample['file']}")290 for method, sample in sample_group.items()291 }292 293 # Get metadata from any method (they should all have the same statement)294 sample_metadata = next(iter(sample_group.values()))['metadata']295 statement = sample_metadata.get('statement', '')296 297 return render_template('experiment.html',298 sample_id=current_index,299 statement=statement,300 visualizations=visualizations,301 methods=METHODS,302 session_id=session_id)303 304 except Exception as e:305 logger.exception(f"An error occurred in the experiment route: {e}")306 return "An error occurred", 500307 308 309@app.route('/completed/<session_id>')310def completed(session_id):311 try:312 session_data = load_session_data(session_id)313 if not session_data:314 return redirect(url_for('index'))315 316 session_data['end_time'] = datetime.now().isoformat()317 responses = session_data['responses']318 319 # Calculate average ranking for each method320 average_rankings = {321 method: mean(r['rankings'][method] for r in responses)322 for method in METHODS323 }324 325 # Sort methods by average ranking (ascending)326 sorted_methods = sorted(327 average_rankings.items(),328 key=lambda x: x[1]329 )330 331 session_data['average_rankings'] = average_rankings332 save_session_data_to_hf(session_id, session_data)333 334 # Clean up local session file335 try:336 os.remove(os.path.join(SESSION_DIR, f'{session_id}.json'))337 except Exception as e:338 logger.warning(f"Error removing session file: {e}")339 340 return render_template(341 'completed.html',342 average_rankings=average_rankings,343 sorted_methods=sorted_methods344 )345 346 except Exception as e:347 logger.exception(f"An error occurred in the completed route: {e}")348 return "An error occurred", 500349 350 351@app.route('/visualizations/<path:filename>')352def send_visualization(filename):353 base_dir = os.getcwd()354 file_path = os.path.normpath(os.path.join(base_dir, filename))355 if not file_path.startswith(base_dir):356 return "Access denied", 403357 if not os.path.exists(file_path):358 return "File not found", 404359 directory = os.path.dirname(file_path)360 file_name = os.path.basename(file_path)361 return send_from_directory(directory, file_name)362 363 364if __name__ == "__main__":365 app.run(host="0.0.0.0", port=7860, debug=True)