Files
remote-sensing/test_cloud_upload.py
T

166 lines
5.3 KiB
Python
Executable File

#!/usr/bin/env python3
"""
Test Cloud Removal Model Upload Feature
"""
import requests
import json
from pathlib import Path
# API base URL
BASE_URL = "http://localhost:8000"
def test_upload_cloud_model(file_path):
"""Test uploading a cloud removal model"""
print(f"\n{'='*60}")
print("TEST 1: Upload Cloud Removal Model")
print(f"{'='*60}")
if not Path(file_path).exists():
print(f"❌ File not found: {file_path}")
print(" Create a dummy .pth file for testing:")
print(f" touch {file_path}")
return None
with open(file_path, 'rb') as f:
files = {'file': (Path(file_path).name, f, 'application/octet-stream')}
print(f"📤 Uploading: {file_path}")
response = requests.post(f"{BASE_URL}/api/cloud-removal/upload", files=files)
if response.status_code == 200:
result = response.json()
print(f"✅ Upload successful!")
print(f" Filename: {result['filename']}")
print(f" Size: {result['size_mb']} MB")
print(f" Path: {result['path']}")
return result['filename']
else:
print(f"❌ Upload failed: {response.status_code}")
print(f" {response.json().get('detail', 'Unknown error')}")
return None
def test_list_cloud_models():
"""Test listing cloud removal models"""
print(f"\n{'='*60}")
print("TEST 2: List Cloud Removal Models")
print(f"{'='*60}")
response = requests.get(f"{BASE_URL}/api/cloud-removal/models")
if response.status_code == 200:
data = response.json()
print(f"✅ Found {data['count']} models:")
for i, model in enumerate(data['models'], 1):
print(f"\n {i}. {model['filename']}")
print(f" Size: {model['size_mb']} MB")
print(f" Created: {model['created']}")
if 'epoch' in model:
print(f" Epoch: {model['epoch']}, Val Loss: {model['val_loss']:.4f}")
return data['models']
else:
print(f"❌ Failed to list models: {response.status_code}")
return []
def test_prediction_with_cloud_model(model_filename, cloud_model_filename):
"""Test prediction using uploaded cloud removal model"""
print(f"\n{'='*60}")
print("TEST 3: Prediction with Custom Cloud Removal Model")
print(f"{'='*60}")
config = {
"model_filename": model_filename,
"min_lon": 105.80,
"min_lat": 10.00,
"max_lon": 105.82,
"max_lat": 10.02,
"start_date": "2024-01-15",
"end_date": "2024-01-17",
"max_scenes": 2,
"cloud_cover": 30,
"resolution": 20,
"use_gpu": False,
"export_ndvi": True,
"export_classification": True,
"cloud_removal_method": "deep",
"cloud_removal_model": cloud_model_filename
}
print("📊 Prediction Config:")
print(json.dumps(config, indent=2))
print(f"\n🚀 Starting prediction with cloud removal model: {cloud_model_filename}")
response = requests.post(
f"{BASE_URL}/api/predict/with-ndvi",
json=config,
headers={'Content-Type': 'application/json'}
)
if response.status_code == 200:
result = response.json()
print(f"✅ Prediction started!")
print(f" Message: {result.get('message')}")
return result
else:
print(f"❌ Prediction failed: {response.status_code}")
print(f" {response.json().get('detail', 'Unknown error')}")
return None
def test_delete_cloud_model(filename):
"""Test deleting a cloud removal model"""
print(f"\n{'='*60}")
print("TEST 4: Delete Cloud Removal Model")
print(f"{'='*60}")
print(f"🗑️ Deleting: {filename}")
response = requests.delete(f"{BASE_URL}/api/cloud-removal/models/{filename}")
if response.status_code == 200:
result = response.json()
print(f"✅ {result['message']}")
return True
else:
print(f"❌ Delete failed: {response.status_code}")
return False
def main():
print("="*60)
print("CLOUD REMOVAL MODEL UPLOAD - FEATURE TEST")
print("="*60)
# Test file path (create a dummy file for testing)
test_file = "test_cloud_removal_model.pth"
# Create dummy file if it doesn't exist
if not Path(test_file).exists():
print(f"\n📝 Creating dummy test file: {test_file}")
Path(test_file).write_bytes(b"dummy_pytorch_model_data")
# Run tests
uploaded_filename = test_upload_cloud_model(test_file)
if uploaded_filename:
models = test_list_cloud_models()
# Test prediction (requires a real land classification model)
print(f"\n{'='*60}")
print("NOTE: Prediction test requires a trained land classification model")
print(" Skipping prediction test in this demo")
print(f"{'='*60}")
# Cleanup - delete test model
if input("\nDelete test model? (y/n): ").lower() == 'y':
test_delete_cloud_model(uploaded_filename)
# Cleanup dummy file
if Path(test_file).exists():
Path(test_file).unlink()
print(f"\n🗑️ Cleaned up dummy file: {test_file}")
print(f"\n{'='*60}")
print("TESTS COMPLETED")
print(f"{'='*60}")
if __name__ == "__main__":
main()