Team Ai
Apppublic

luulinh90s/Tabular-LLM-Study-Preference-Ranking

sourceHugging Facemitupdated 2y agoView on Hugging Face
0likes
app.py365 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_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)