pascalx/pathloss-predictor
0
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 