Team Ai
Apppublic

luulinh90s/Tabular-LLM-Study-Debugging

sourceHugging Facemitupdated 2y agoView on Hugging Face
0likes
app.py389 linesDownload Raw Back to root
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)