def visualize_predictions(model, dataset, n_samples=5):
"""Visualize model predictions vs ground truth"""
images, masks = next(iter(dataset))
predictions = model.predict(images[:n_samples], verbose=0)
fig, axes = plt.subplots(n_samples, 4, figsize=(20, n_samples*5))
for i in range(n_samples):
# Original SAR VV
axes[i, 0].imshow(images[i, :, :, 0], cmap='gray', vmin=0, vmax=1)
axes[i, 0].set_title(f'SAR VV (Normalized)')
# Ground Truth
axes[i, 1].imshow(masks[i, :, :, 0], cmap='Blues', vmin=0, vmax=1)
axes[i, 1].set_title('Ground Truth Mask')
# Prediction
iou_val = iou_score(masks[i:i+1], predictions[i:i+1]).numpy()
axes[i, 2].imshow(predictions[i, :, :, 0], cmap='Blues', vmin=0, vmax=1)
axes[i, 2].set_title(f'Prediction (IoU: {iou_val:.3f})')
# Error Overlay: TP=Green, FP=Red, FN=Yellow
overlay = np.zeros((256, 256, 3))
gt = masks[i, :, :, 0] > 0.5
pred = predictions[i, :, :, 0] > 0.5
overlay[gt & pred] = [0, 1, 0]
overlay[~gt & pred] = [1, 0, 0]
overlay[gt & ~pred] = [1, 1, 0]
axes[i, 3].imshow(overlay)
axes[i, 3].set_title('Errors: Green=TP, Red=FP, Yellow=FN')