Team Ai
Apppublic

pascalx/pathloss-predictor

sourceHugging Facemitupdated 1y agoView on Hugging Face
0likes
app.py298 linesDownload Raw Back to root
1from flask import Flask, render_template, request, jsonify2import numpy as np3import pandas as pd4import tensorflow as tf5from sklearn.preprocessing import StandardScaler, OneHotEncoder6from sklearn.compose import ColumnTransformer7import pickle8import os9 10app = Flask(__name__)11 12# Global variables to store models and preprocessor13models = {}14preprocessor = None15 16def load_models_and_preprocessor():17    """Load the trained models and preprocessor"""18    global models, preprocessor19    20    try:21        # Set TensorFlow to use CPU only to avoid GPU compatibility issues22        tf.config.set_visible_devices([], 'GPU')23        24        # Load models with custom_objects to handle metric compatibility25        custom_objects = {26            'mse': tf.keras.metrics.MeanSquaredError(),27            'mae': tf.keras.metrics.MeanAbsoluteError(),28            'mape': tf.keras.metrics.MeanAbsolutePercentageError()29        }30        31        # Try to load models with different approaches32        model_files = ['models/ann_model.h5', 'models/dnn_model.h5', 'models/cnn_model.h5']33        model_names = ['ann', 'dnn', 'cnn']34        35        for model_file, model_name in zip(model_files, model_names):36            try:37                # First try with custom_objects and compile=False38                models[model_name] = tf.keras.models.load_model(39                    model_file, 40                    custom_objects=custom_objects, 41                    compile=False42                )43                print(f"Successfully loaded {model_name} model")44            except Exception as e:45                print(f"Error loading {model_name} model: {e}")46                # Try alternative loading method47                try:48                    models[model_name] = tf.keras.models.load_model(49                        model_file, 50                        compile=False,51                        options=tf.saved_model.LoadOptions(experimental_io_device='/cpu:0')52                    )53                    print(f"Successfully loaded {model_name} model with alternative method")54                except Exception as e2:55                    print(f"Failed to load {model_name} model with alternative method: {e2}")56                    return False57        58        # Recompile models with current TensorFlow version59        for model_name, model in models.items():60            try:61                model.compile(62                    optimizer='adam',63                    loss='mse',64                    metrics=['mae', 'mape']65                )66                print(f"Successfully compiled {model_name} model")67            except Exception as e:68                print(f"Error compiling {model_name} model: {e}")69        70        # Load preprocessor71        try:72            with open('models/preprocessor.pkl', 'rb') as f:73                global preprocessor74                preprocessor = pickle.load(f)75            print("Preprocessor loaded successfully!")76        except Exception as e:77            print(f"Error loading preprocessor: {e}")78            return False79            80        print("Models and preprocessor loaded successfully!")81        return True82    except Exception as e:83        print(f"Error loading models: {e}")84        return False85 86def make_prediction(model_type, frequency, distance, tx_height, rx_height, environment):87    """Make prediction using the specified model"""88    try:89        # Create input dataframe90        input_data = pd.DataFrame({91            'Frequency_MHz': [frequency],92            'Distance_km': [distance],93            'Tx_Height_m': [tx_height],94            'Rx_Height_m': [rx_height],95            'Environment': [environment]96        })97        print(f"Input data created: {input_data}")98        99        if preprocessor is None:100            print("Error: Preprocessor not loaded")101            return None102            103        # Preprocess the input104        try:105            input_processed = preprocessor.transform(input_data)106            print(f"Input preprocessed successfully. Shape: {input_processed.shape}")107        except Exception as e:108            print(f"Error in preprocessing: {e}")109            return None110        111        # Get the model112        if model_type not in models:113            print(f"Error: Model {model_type} not found in loaded models")114            return None115            116        model = models[model_type]117        118        # For CNN model, we need to reshape the input119        if model_type == 'cnn':120            input_processed = np.expand_dims(input_processed, axis=2)121            print(f"CNN input shape after reshape: {input_processed.shape}")122        123        # Make prediction124        try:125            prediction = model.predict(input_processed)[0][0]126            print(f"Prediction successful: {prediction}")127            return float(prediction)128        except Exception as e:129            print(f"Error during model prediction: {e}")130            return None131    132    except Exception as e:133        print(f"Error in make_prediction: {e}")134        return None135 136@app.route('/')137def index():138    """Main page with the prediction form"""139    return render_template('index.html')140 141@app.route('/predict', methods=['POST'])142def predict():143    """Handle prediction requests"""144    try:145        # Get form data146        frequency = float(request.form['frequency'])147        distance = float(request.form['distance'])148        tx_height = float(request.form['tx_height'])149        rx_height = float(request.form['rx_height'])150        environment = request.form['environment']151        model_type = request.form['model_type']152        153        # Validate inputs154        if frequency <= 0 or distance <= 0 or tx_height <= 0 or rx_height <= 0:155            return render_template('result.html', 156                                 error="All numeric values must be positive")157        158        # Make prediction159        prediction = make_prediction(model_type, frequency, distance, 160                                   tx_height, rx_height, environment)161        162        if prediction is None:163            # Check what might be wrong164            error_msg = "Error making prediction. "165            if not models:166                error_msg += "Models not loaded. "167            if preprocessor is None:168                error_msg += "Preprocessor not loaded. "169            if model_type not in models:170                error_msg += f"Model '{model_type}' not available. "171            error_msg += "Please check the logs or try again."172            173            return render_template('result.html', error=error_msg)174        175        # Prepare result data176        result_data = {177            'prediction': round(prediction, 2),178            'model_type': model_type.upper(),179            'inputs': {180                'frequency': frequency,181                'distance': distance,182                'tx_height': tx_height,183                'rx_height': rx_height,184                'environment': environment185            }186        }187        188        return render_template('result.html', result=result_data)189        190    except ValueError:191        return render_template('result.html', 192                             error="Please enter valid numeric values")193    except Exception as e:194        return render_template('result.html', 195                             error=f"An error occurred: {str(e)}")196 197@app.route('/api/predict', methods=['POST'])198def api_predict():199    """API endpoint for predictions (JSON response)"""200    try:201        data = request.get_json()202        203        prediction = make_prediction(204            data['model_type'],205            data['frequency'],206            data['distance'],207            data['tx_height'],208            data['rx_height'],209            data['environment']210        )211        212        if prediction is None:213            return jsonify({'error': 'Prediction failed'}), 500214        215        return jsonify({216            'prediction': round(prediction, 2),217            'model_type': data['model_type'],218            'status': 'success'219        })220        221    except Exception as e:222        return jsonify({'error': str(e)}), 400223 224@app.route('/health')225def health_check():226    """Health check endpoint to verify models are loaded"""227    try:228        status = {229            'models_loaded': len(models) > 0,230            'preprocessor_loaded': preprocessor is not None,231            'available_models': list(models.keys()),232            'status': 'healthy' if len(models) > 0 and preprocessor is not None else 'unhealthy'233        }234        return jsonify(status)235    except Exception as e:236        return jsonify({'error': str(e), 'status': 'error'}), 500237 238@app.route('/test')239def test_prediction():240    """Test endpoint to verify prediction functionality"""241    try:242        # Test with sample data243        test_data = {244            'frequency': 900.0,245            'distance': 5.0,246            'tx_height': 30.0,247            'rx_height': 1.5,248            'environment': 'Urban',249            'model_type': 'ann'250        }251        252        prediction = make_prediction(253            test_data['model_type'],254            test_data['frequency'],255            test_data['distance'],256            test_data['tx_height'],257            test_data['rx_height'],258            test_data['environment']259        )260        261        if prediction is not None:262            return jsonify({263                'status': 'success',264                'test_prediction': prediction,265                'test_data': test_data,266                'message': 'Prediction system is working correctly!'267            })268        else:269            return jsonify({270                'status': 'error',271                'message': 'Prediction failed - check logs for details'272            }), 500273            274    except Exception as e:275        return jsonify({276            'status': 'error',277            'message': f'Test failed: {str(e)}'278        }), 500279 280# Load models and preprocessor when the module is imported (for Hugging Face Spaces)281print("Starting Pathloss Prediction App...")282print("Loading models and preprocessor...")283 284if load_models_and_preprocessor():285    print("โœ… All models loaded successfully!")286    print("๐Ÿš€ Flask application ready!")287else:288    print("โŒ Failed to load models. Please ensure model files are present.")289    print("Required files:")290    print("  - models/ann_model.h5")291    print("  - models/dnn_model.h5") 292    print("  - models/cnn_model.h5")293    print("  - models/preprocessor.pkl")294 295if __name__ == '__main__':296    # Only run the Flask app directly if this script is executed297    app.run(host='0.0.0.0', port=7860, debug=True)298