hoàn thành tính cận trên và cận dưới của tất cả các thuật toán
This commit is contained in:
Regular → Executable
+245
-2
@@ -88,6 +88,30 @@ prediction_status = {
|
||||
batch_queue = []
|
||||
batch_results = []
|
||||
|
||||
# Cloud removal training status (with baseline PSNR tracking)
|
||||
cloud_training_status = {
|
||||
"is_training": False,
|
||||
"progress": "",
|
||||
"error": None,
|
||||
"training_id": None,
|
||||
"current_epoch": 0,
|
||||
"total_epochs": 0,
|
||||
"train_loss": 0.0,
|
||||
"val_loss": 0.0,
|
||||
"val_psnr": 0.0,
|
||||
"baseline_psnr": 0.0,
|
||||
"improvement": 0.0,
|
||||
"start_time": None,
|
||||
"end_time": None,
|
||||
"history": {
|
||||
"epochs": [],
|
||||
"train_losses": [],
|
||||
"val_losses": [],
|
||||
"val_psnrs": [],
|
||||
"improvements": []
|
||||
}
|
||||
}
|
||||
|
||||
# Label mapping from training data (from 01.train_ODC.ipynb)
|
||||
DEFAULT_LABEL_MAPPING = {
|
||||
"Lua tom": "0",
|
||||
@@ -561,17 +585,23 @@ async def list_cloud_removal_models():
|
||||
epoch = checkpoint.get('epoch', 0) if isinstance(checkpoint, dict) else 0
|
||||
train_loss = checkpoint.get('train_loss', 0) if isinstance(checkpoint, dict) else 0
|
||||
val_loss = checkpoint.get('val_loss', 0) if isinstance(checkpoint, dict) else 0
|
||||
val_psnr = checkpoint.get('val_psnr', None) if isinstance(checkpoint, dict) else None
|
||||
baseline_psnr = checkpoint.get('baseline_psnr', None) if isinstance(checkpoint, dict) else None
|
||||
use_s1 = checkpoint.get('use_s1', True) if isinstance(checkpoint, dict) else True
|
||||
in_channels = checkpoint.get('in_channels', 6) if isinstance(checkpoint, dict) else 6
|
||||
out_channels = checkpoint.get('out_channels', 4) if isinstance(checkpoint, dict) else 4
|
||||
hyperparams = checkpoint.get('hyperparameters', {}) if isinstance(checkpoint, dict) else {}
|
||||
except:
|
||||
# If checkpoint format is different or corrupted, use defaults
|
||||
epoch = 0
|
||||
train_loss = 0
|
||||
val_loss = 0
|
||||
val_psnr = None
|
||||
baseline_psnr = None
|
||||
use_s1 = True
|
||||
in_channels = 3
|
||||
out_channels = 3
|
||||
hyperparams = {}
|
||||
|
||||
models.append({
|
||||
"filename": model_file.name,
|
||||
@@ -580,9 +610,12 @@ async def list_cloud_removal_models():
|
||||
"epoch": epoch,
|
||||
"train_loss": train_loss,
|
||||
"val_loss": val_loss,
|
||||
"val_psnr": val_psnr,
|
||||
"baseline_psnr": baseline_psnr,
|
||||
"use_s1": use_s1,
|
||||
"in_channels": in_channels,
|
||||
"out_channels": out_channels,
|
||||
"hyperparameters": hyperparams,
|
||||
"description": "",
|
||||
"created": model_file.stat().st_mtime,
|
||||
"size_mb": model_file.stat().st_size / (1024 * 1024),
|
||||
@@ -611,6 +644,26 @@ async def list_cloud_removal_models():
|
||||
return {"models": models, "count": len(models)}
|
||||
|
||||
|
||||
@app.get("/api/cloud-removal/training/status")
|
||||
async def get_cloud_training_status():
|
||||
"""Lấy trạng thái training cloud removal với baseline PSNR"""
|
||||
global cloud_training_status
|
||||
return cloud_training_status
|
||||
|
||||
|
||||
@app.post("/api/cloud-removal/training/stop")
|
||||
async def stop_cloud_training():
|
||||
"""Dừng cloud training đang chạy"""
|
||||
global cloud_training_status
|
||||
|
||||
if not cloud_training_status["is_training"]:
|
||||
raise HTTPException(status_code=400, detail="No training is running")
|
||||
|
||||
cloud_training_status["progress"] = "Stopping..."
|
||||
# Training loop should check this flag
|
||||
return {"message": "Stopping cloud removal training..."}
|
||||
|
||||
|
||||
@app.post("/api/cloud-removal/train")
|
||||
async def train_cloud_removal(config: CloudRemovalTrainingConfig, background_tasks: BackgroundTasks):
|
||||
"""Bắt đầu train cloud removal model"""
|
||||
@@ -627,28 +680,161 @@ async def train_cloud_removal(config: CloudRemovalTrainingConfig, background_tas
|
||||
training_id = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
|
||||
async def run_cloud_training():
|
||||
global cloud_training_status
|
||||
|
||||
try:
|
||||
# Reset status
|
||||
cloud_training_status = {
|
||||
"is_training": True,
|
||||
"progress": "Initializing...",
|
||||
"error": None,
|
||||
"training_id": training_id,
|
||||
"current_epoch": 0,
|
||||
"total_epochs": config.num_epochs,
|
||||
"train_loss": 0.0,
|
||||
"val_loss": 0.0,
|
||||
"val_psnr": 0.0,
|
||||
"baseline_psnr": 0.0,
|
||||
"improvement": 0.0,
|
||||
"start_time": datetime.now().isoformat(),
|
||||
"end_time": None,
|
||||
"history": {
|
||||
"epochs": [],
|
||||
"train_losses": [],
|
||||
"val_losses": [],
|
||||
"val_psnrs": [],
|
||||
"improvements": []
|
||||
}
|
||||
}
|
||||
|
||||
from train_cloud_removal import train_cloud_removal_model
|
||||
|
||||
print(f"[CLOUD REMOVAL TRAINING] Starting training {training_id}")
|
||||
|
||||
model, train_losses, val_losses = train_cloud_removal_model(
|
||||
# Status callback function
|
||||
def update_status(epoch, train_loss, val_loss, val_psnr, baseline_psnr):
|
||||
cloud_training_status["current_epoch"] = epoch
|
||||
cloud_training_status["train_loss"] = train_loss
|
||||
cloud_training_status["val_loss"] = val_loss
|
||||
cloud_training_status["val_psnr"] = val_psnr
|
||||
cloud_training_status["baseline_psnr"] = baseline_psnr
|
||||
cloud_training_status["improvement"] = val_psnr - baseline_psnr
|
||||
cloud_training_status["progress"] = f"Epoch {epoch}/{config.num_epochs}"
|
||||
|
||||
# Add to history
|
||||
cloud_training_status["history"]["epochs"].append(epoch)
|
||||
cloud_training_status["history"]["train_losses"].append(train_loss)
|
||||
cloud_training_status["history"]["val_losses"].append(val_loss)
|
||||
cloud_training_status["history"]["val_psnrs"].append(val_psnr)
|
||||
cloud_training_status["history"]["improvements"].append(val_psnr - baseline_psnr)
|
||||
|
||||
print(f"[STATUS UPDATE] Epoch {epoch}: PSNR={val_psnr:.2f}dB, Baseline={baseline_psnr:.2f}dB, Improvement={val_psnr-baseline_psnr:+.2f}dB")
|
||||
|
||||
model, train_losses, val_losses, val_psnrs, baseline_psnr = train_cloud_removal_model(
|
||||
data_dir=config.data_dir,
|
||||
use_s1=config.use_s1,
|
||||
batch_size=config.batch_size,
|
||||
num_epochs=config.num_epochs,
|
||||
learning_rate=config.learning_rate,
|
||||
device="cuda" if config.use_gpu else "cpu",
|
||||
save_dir="cloud_removal_model"
|
||||
save_dir="cloud_removal_model",
|
||||
status_callback=update_status
|
||||
)
|
||||
|
||||
print(f"[CLOUD REMOVAL TRAINING] Completed {training_id}")
|
||||
|
||||
cloud_training_status["is_training"] = False
|
||||
cloud_training_status["progress"] = "Completed!"
|
||||
cloud_training_status["end_time"] = datetime.now().isoformat()
|
||||
cloud_training_status["result"] = {
|
||||
"success": True,
|
||||
"training_id": training_id,
|
||||
"final_train_loss": train_losses[-1],
|
||||
"final_val_loss": val_losses[-1],
|
||||
"final_val_psnr": val_psnrs[-1],
|
||||
"baseline_psnr": baseline_psnr,
|
||||
"improvement": val_psnrs[-1] - baseline_psnr,
|
||||
"epochs": len(train_losses)
|
||||
}
|
||||
|
||||
# Calculate best checkpoint (cao nhất - cận trên)
|
||||
best_epoch_idx = val_psnrs.index(max(val_psnrs)) if val_psnrs else 0
|
||||
best_checkpoint = {
|
||||
'modelPSNR': val_psnrs[best_epoch_idx] if val_psnrs else 0,
|
||||
'epoch': best_epoch_idx + 1,
|
||||
'trainLoss': train_losses[best_epoch_idx] if train_losses else 0,
|
||||
'valLoss': val_losses[best_epoch_idx] if val_losses else 0,
|
||||
'baselinePSNR': baseline_psnr,
|
||||
'improvement': (val_psnrs[best_epoch_idx] - baseline_psnr) if val_psnrs else 0
|
||||
}
|
||||
|
||||
# Calculate worst checkpoint (thấp nhất - cận dưới)
|
||||
worst_epoch_idx = val_psnrs.index(min(val_psnrs)) if val_psnrs else 0
|
||||
worst_checkpoint = {
|
||||
'modelPSNR': val_psnrs[worst_epoch_idx] if val_psnrs else 0,
|
||||
'epoch': worst_epoch_idx + 1,
|
||||
'trainLoss': train_losses[worst_epoch_idx] if train_losses else 0,
|
||||
'valLoss': val_losses[worst_epoch_idx] if val_losses else 0,
|
||||
'baselinePSNR': baseline_psnr,
|
||||
'improvement': (val_psnrs[worst_epoch_idx] - baseline_psnr) if val_psnrs else 0
|
||||
}
|
||||
|
||||
# Generate training report
|
||||
try:
|
||||
training_result = {
|
||||
'training_id': training_id,
|
||||
'model_type': 'cloud_removal_unet',
|
||||
'train_accuracy': 0, # N/A for cloud removal
|
||||
'test_accuracy': 0, # N/A for cloud removal
|
||||
'val_psnr': val_psnrs[-1],
|
||||
'baseline_psnr': baseline_psnr,
|
||||
'train_loss': train_losses[-1],
|
||||
'val_loss': val_losses[-1],
|
||||
'best_checkpoint': best_checkpoint,
|
||||
'worst_checkpoint': worst_checkpoint,
|
||||
'training_samples': len(train_dataset) if 'train_dataset' in locals() else 0,
|
||||
'testing_samples': len(val_dataset) if 'val_dataset' in locals() else 0,
|
||||
'test_size': 0.2,
|
||||
'classes': [], # N/A for cloud removal
|
||||
'classification_report': {},
|
||||
'confusion_matrix': [],
|
||||
'model_path': str(Path('cloud_removal_model') / 'cloud_removal_unet_best.pth'),
|
||||
'bbox': [],
|
||||
'time_range': '',
|
||||
'resolution': 10,
|
||||
'data_source': f'SEN12MS-CR Dataset ({config.data_dir})',
|
||||
'collections': ['Sentinel-2 L2A', 'Sentinel-1 RTC'] if config.use_s1 else ['Sentinel-2 L2A'],
|
||||
'features': ['B02', 'B03', 'B04', 'B08', 'B11'] + (['VV', 'VH'] if config.use_s1 else []),
|
||||
'feature_mode': 'image',
|
||||
'n_features': 6 if config.use_s1 else 4,
|
||||
'hyperparameters': {
|
||||
'data_dir': config.data_dir,
|
||||
'use_s1': config.use_s1,
|
||||
'batch_size': config.batch_size,
|
||||
'num_epochs': config.num_epochs,
|
||||
'learning_rate': config.learning_rate,
|
||||
'use_gpu': config.use_gpu,
|
||||
'model_type': 'unet',
|
||||
'architecture': 'U-Net',
|
||||
'optimizer': 'Adam',
|
||||
'criterion': 'L1Loss',
|
||||
'scheduler': 'ReduceLROnPlateau'
|
||||
}
|
||||
}
|
||||
|
||||
report_path, _ = generate_training_report(training_result, config=None)
|
||||
print(f"[REPORT] Generated training report: {report_path}")
|
||||
except Exception as report_err:
|
||||
print(f"[REPORT ERROR] Failed to generate report: {report_err}")
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"training_id": training_id,
|
||||
"final_train_loss": train_losses[-1],
|
||||
"final_val_loss": val_losses[-1],
|
||||
"final_val_psnr": val_psnrs[-1],
|
||||
"baseline_psnr": baseline_psnr,
|
||||
"improvement": val_psnrs[-1] - baseline_psnr,
|
||||
"epochs": len(train_losses)
|
||||
}
|
||||
|
||||
@@ -656,6 +842,12 @@ async def train_cloud_removal(config: CloudRemovalTrainingConfig, background_tas
|
||||
print(f"[CLOUD REMOVAL TRAINING ERROR] {e}")
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
|
||||
cloud_training_status["is_training"] = False
|
||||
cloud_training_status["error"] = str(e)
|
||||
cloud_training_status["progress"] = f"Error: {str(e)}"
|
||||
cloud_training_status["end_time"] = datetime.now().isoformat()
|
||||
|
||||
return {
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
@@ -1429,6 +1621,57 @@ async def stop_training():
|
||||
return {"message": "Đang dừng training..."}
|
||||
|
||||
|
||||
@app.post("/api/training/report/regenerate")
|
||||
async def regenerate_training_report(request: dict):
|
||||
"""
|
||||
Regenerate training report with best_checkpoint, worst_checkpoint, and random_baseline data from frontend
|
||||
|
||||
Body:
|
||||
- training_result: dict (original training result)
|
||||
- best_checkpoint: dict (bestCheckpoint data from frontend)
|
||||
- worst_checkpoint: dict (worstCheckpoint data from frontend, optional)
|
||||
- random_baseline: float (random baseline accuracy for land classification, optional)
|
||||
"""
|
||||
try:
|
||||
training_result = request.get('training_result', {})
|
||||
best_checkpoint = request.get('best_checkpoint', None)
|
||||
worst_checkpoint = request.get('worst_checkpoint', None)
|
||||
random_baseline = request.get('random_baseline', None)
|
||||
|
||||
if not training_result:
|
||||
return {"success": False, "error": "Missing training_result"}
|
||||
|
||||
# Add best_checkpoint to training_result
|
||||
if best_checkpoint:
|
||||
# Convert camelCase to snake_case if needed
|
||||
if 'trainAcc' in best_checkpoint:
|
||||
# Frontend uses camelCase, keep it as is
|
||||
training_result['best_checkpoint'] = best_checkpoint
|
||||
else:
|
||||
training_result['best_checkpoint'] = best_checkpoint
|
||||
|
||||
# Add worst_checkpoint to training_result
|
||||
if worst_checkpoint:
|
||||
training_result['worst_checkpoint'] = worst_checkpoint
|
||||
|
||||
# Add random_baseline to training_result
|
||||
if random_baseline is not None:
|
||||
training_result['random_baseline'] = random_baseline
|
||||
|
||||
# Regenerate report
|
||||
report_path, _ = generate_training_report(training_result)
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"report_path": report_path,
|
||||
"report_filename": Path(report_path).name
|
||||
}
|
||||
except Exception as e:
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
return {"success": False, "error": str(e)}
|
||||
|
||||
|
||||
@app.post("/api/cache/clear")
|
||||
async def clear_cache():
|
||||
"""Xóa cache dataset"""
|
||||
|
||||
Reference in New Issue
Block a user