
    xi;                     	   d dl Zd dlmZ d dlmZmZ d dlmZm	Z	m
Z
mZmZ d dlmZ d dlZ ej                   ej"                  d        ej$                  e      Z G d d	e      Ze G d
 d             Ze G d d             Ze G d d             Z G d d      Z G d de      Z G d de      Z G d de      Z G d de      Z G d de      Z G d de      Z G d de      Z  G d  d!e      Z! G d" d#e      Z" G d$ d%e      Z# G d& d'e      Z$ G d( d)e      Z% G d* d+e      Z& G d, d-e      Z' G d. d/e      Z( G d0 d1      Z) G d2 d3e      Z* G d4 d5e      Z+ G d6 d7      Z, G d8 d9      Z-ed:k(  r
 e.d;        e,d<=      Z/ej`                  jc                  d>       ej`                  je                  d?      d@z   ejf                  d?       ejh                   ejj                  d?      dAz        dAz   ejl                  ej`                  je                  dB      d@z  ej`                  je                  dC      dDz  g      dEZ7 e.dF       e7jq                         D ]6  \  Z9Z:e:jw                  dGdG      Z< e.dHe9j{                          dI        e.dJ ej|                  e:      dKdL ej~                  e:      dK       e/j                  e<      ZA e.dM        eBeAjq                         dN dOP      D ]=  \  ZCZDeDj                  dQkD  s e.dReC dSeDj                  dKdTeDj                  dUdV       ? e/j                  eA      \  ZHZIZJ e.dWeHdK        e.dXeJj                          eJj                  r e.dYeJj                  d             e.        9  e.dZ       	  ed[      d\   jw                  d]d?      ZM ed^      d\   d_z  j                  eO      ZPeMjw                  d]dGdGd\      ZQ e-d?d`dad<b      ZReRj                  eQePd`c        edd      d\   jw                  d]d?      ZT ede      d\   d_z  j                  eO      ZUeTjw                  d]dGdGd\      ZVeRj                  eVeU      ZXej                  dfeXdK       yy# eZ$ rf  e.dg       ej`                  je                  dhdGdGd\      Z[ej`                  j                  d dadh      Z] e-d?didad<b      ZReRj                  e[e]dic       Y yw xY w)j    N)read)	dataclassfield)DictListTupleOptionalCallable)Enumz)%(asctime)s - %(levelname)s - %(message)s)levelformatc                   ,    e Zd ZdZdZdZdZdZdZdZ	dZ
y	)
LossCategoryz Categories of diagnostic losses.spatialspectralstatistical	componentpatternadversarialregularizationN)__name__
__module____qualname____doc__SPATIALSPECTRALSTATISTICAL	COMPONENTPATTERNADVERSARIALREGULARIZATION     err01.pyr   r      s(    *GHKIGK%Nr#   r   c                   ~    e Zd ZU dZeed<   eed<   dZeed<   dZ	e
ed<   dZee   ed	<   d
 ed      fZeeef   ed<   y)
LossConfigz)Configuration for a single loss function.namecategory      ?weightTenabledNtarget        inf
thresholds)r   r   r   r   str__annotations__r   r*   floatr+   boolr,   r	   r/   r   r"   r#   r$   r&   r&      sN    3
IFEGT"FHUO"'*E%L&9JeUl#9r#   r&   c                   l    e Zd ZU dZeed<   eed<   ej                  ed<   e	ed<   e
ed<   eed<   eed<   y	)

LossResultz&Result from computing a loss function.r'   valuegradientdiagnosticsr(   r*   normalized_valueN)r   r   r   r   r0   r1   r2   npndarrayr   r   r"   r#   r$   r5   r5   &   s0    0
ILjjMr#   r5   c                   r    e Zd ZU dZeeef   ed<   eed<   eeef   ed<   eed<   eeef   ed<   e	e   ed<   y)	MultiLossStatez+State tracking for multi-loss optimization.losses
total_lossweighted_contributionsdominant_lossgradient_normsoptimization_adviceN)
r   r   r   r   r   r0   r5   r1   r2   r   r"   r#   r$   r=   r=   2   sH    5j!! e,,e$$c"r#   r=   c            	           e Zd ZdZdefdZddej                  dej                  dej                  defd	Z	dej                  d
ej                  dej                  fdZ
y)BaseLossFunctionz;Base class for all loss functions derived from diagnostics.configc                      || _         g | _        y N)rF   history)selfrF   s     r$   __init__zBaseLossFunction.__init__D   s    r#   NerrorX	referencereturnc                     t         )zCompute loss and gradient.NotImplementedError)rJ   rL   rM   rN   s       r$   computezBaseLossFunction.computeH       !!r#   r7   c                     t         )z'Compute gradient with respect to input.rQ   )rJ   rM   r7   s      r$   get_gradientzBaseLossFunction.get_gradientL   rT   r#   NN)r   r   r   r   r&   rK   r:   r;   r5   rS   rV   r"   r#   r$   rE   rE   A   se    Ez "RZZ "BJJ ""** "`j ""bjj "BJJ "2:: "r#   rE   c            	       f    e Zd ZdZddej
                  dej
                  dej
                  defdZy)	MSELossz!Baseline mean squared error loss.NrL   rM   rN   rO   c                    |j                         }t        t        j                  |dz              }dt	        |j
                  d      z  |z  }t        | j                  j                  ||dt        t        j                  |            i| j                  j                  | j                  j                  t        |d            S )N          @   rmser)   r'   r6   r7   r8   r(   r*   r9   )ravelr2   r:   meanmaxsizer5   rF   r'   sqrtr(   r*   min)rJ   rL   rM   rN   
flat_error
loss_valuer7   s          r$   rS   zMSELoss.computeT   s    [[]
277:?34
#jooq11U:!!rwwz':!;<[[));;%% S1
 	
r#   rW   r   r   r   r   r:   r;   r5   rS   r"   r#   r$   rY   rY   Q   s3    +
RZZ 
BJJ 
"** 
`j 
r#   rY   c            	           e Zd ZdZd
dej
                  dej
                  dej
                  defdZdej
                  dej
                  dej
                  fd	Zy)SpatialConcentrationLossa  
    Loss based on spatial concentration of errors.
    
    Encourages errors to be spread uniformly (like good convolution outputs).
    Penalizes clustered errors which indicate model focusing too much
    on specific regions.
    
    Loss = concentration_score (higher = worse)
    NrL   rM   rN   rO   c                 
   |j                         d d j                  dd      }g }d}t        dd|z
  |      D ]L  }t        dd|z
  |      D ]7  }||||z   |||z   f   }	|j                  t	        j
                  |	             9 N |rt	        j                  |      nd}
t	        j
                  |      dz   }|
|z  }t	        j                  |d      }t	        j                  |d      }t	        j                  |dz  |dz  z         }|d	kD  r|nd
}t        | j                  j                  ||j                         d t        |       ||
|dt        j                  | j                  j                  t!        |d      dz        S )N        r   :0yE>axisr]   r[   r)   r-   )concentration	local_var
global_varr\   r_   )r`   reshaperangeappendr:   varra   r7   rd   r5   rF   r'   lenr   r   r*   re   )rJ   rL   rM   rN   error_2dlocal_windowswindow_sizeyxwindowrs   rt   rr   grad_ygrad_xgradient_magnituderg   s                    r$   rS   z SpatialConcentrationLoss.computes   s   ;;=#&..r26 q"{*K8 	5A1b;.< 5!!AkM/1Q{]?"BC$$RVVF^45	5
 /<BGGM*	VVH%,
 "J. XA.XA.WWVQY%:; '4c&9]s
!!'--/U<*7i_ij!));;%% 4s:
 	
r#   output_gradc                     |S rH   r"   )rJ   rM   r   s      r$   rV   z%SpatialConcentrationLoss.get_gradient   s    r#   rW   )	r   r   r   r   r:   r;   r5   rS   rV   r"   r#   r$   rj   rj   h   s\    !
RZZ !
BJJ !
"** !
`j !
Fbjj rzz bjj r#   rj   c            	       f    e Zd ZdZddej
                  dej
                  dej
                  defdZy)	SpatialEntropyLossz
    Loss based on spatial entropy of error distribution.
    
    Encourages uniform error distribution (high entropy).
    Penalizes focused/predictable error patterns.
    
    Target: Maximum entropy = uniform distribution
    NrL   rM   rN   rO   c           	      z   |j                         d d j                  dd      }t        j                  |      j                         }|t        j                  |      dz   z  }t        j                  |t        j
                  |dz         z         }t        j
                  t        |            }||dz   z  }d|z
  }	t        j                  |      d|z
  z  }
t        | j                  j                  |	|
j                         d t        |       |||dt        j                  | j                  j                  |	      S )Nrl   rm   ro   r)   r]   )entropymax_entropy
normalizedr_   )r`   ru   r:   abssumlogry   r5   rF   r'   r   r   r*   )rJ   rL   rM   rN   rz   	error_magr   r   normalized_entropyrg   r7   s              r$   rS   zSpatialEntropyLoss.compute   s   ;;=#&..r26 FF8$**,		!2T!9:	 66)bffY-=&>>??ffS^, %d(:; --
 66(#q+='=>!!^^%ks5z2$+KWij!));;%%'
 	
r#   rW   rh   r"   r#   r$   r   r      s5    
RZZ 
BJJ 
"** 
`j 
r#   r   c            	            e Zd ZdZddededef fdZddej                  dej                  dej                  d	e
fd
Z xZS )HotspotPenaltyLossz
    Loss that specifically penalizes error hotspots.
    
    Identifies top-k highest error regions and applies
    additional penalty to reduce them.
    
    Inspired by attention mechanism - focus correction
    on worst regions.
    rF   
n_hotspotspenalty_scalec                 @    t         |   |       || _        || _        y rH   )superrK   r   r   )rJ   rF   r   r   	__class__s       r$   rK   zHotspotPenaltyLoss.__init__   s     $*r#   rL   rM   rN   rO   c                 B   |j                         d d j                  dd      }|j                  }t        j                  |      j                         }t        j
                  |dd| j                  t        |      z  z
  z        }||k\  }t        j                  |      }	||   }
|	dkD  rt        j                  |
      nd}|	dkD  rt        j                  |      nd}t        j                  |
| j                  z        |	dz   z  }t        j                  |      }|
| j                  z  ||<   t        | j                  j                  ||d | j                  |j                         |	|||dt"        j$                  | j                  j&                  t)        |dz  d      	      S )
Nrl   rm   d   r]   r   )r   mean_hotspot_errormax_hotspot_error	threshold      $@r)   r_   )r`   ru   rc   r:   r   
percentiler   ry   r   ra   rb   r   
zeros_liker5   rF   r'   shaper   r   r*   re   )rJ   rL   rM   rN   rz   
error_sizerf   r   hotspot_maskn_hotspots_actualhotspot_errorr   r   rg   hotspot_gradients                  r$   rS   zHotspotPenaltyLoss.compute   s|   ;;=#&..r26ZZ
 VVH%++-
MM*cQ3z?9Z5Z.[\	!Y. FF<0"<07H17LRWW]3RS2Ca2GBFF:.Q VVMT-?-??@DUXYDYZ
 ==4)69K9K)K&!!%kz2::5;;G/&8%6&	 "));;%% d!2C8
 	
r#   )rn   r\   rW   )r   r   r   r   r&   intr2   rK   r:   r;   r5   rS   __classcell__r   s   @r$   r   r      sP    +z +s +u +
$
RZZ $
BJJ $
"** $
`j $
r#   r   c            	            e Zd ZdZddededef fdZddej                  dej                  dej                  d	e
fd
Z xZS )SpectralBandLossa  
    Loss based on error energy in different frequency bands.
    
    Allows separate control over:
    - Low frequency (overall structure)
    - Mid frequency (shapes and patterns)
    - High frequency (edges and details)
    
    Each band can have different target weights.
    rF   bandtarget_ratioc                 @    t         |   |       || _        || _        y rH   )r   rK   r   r   )rJ   rF   r   r   r   s       r$   rK   zSpectralBandLoss.__init__  s     	(r#   rL   rM   rN   rO   c                 
   |j                         d d j                  dd      }t        j                  j	                  |      }t        j                  j                  |      }t        j                  |      }d\  }}	|dz  |	dz  }}
t        j                  d |d |	f   \  }}t        j                  ||
z
  dz  ||z
  dz  z         }t        j                  |
dz  |dz  z         }| j                  dk(  r	||dz  k  }n)| j                  dk(  r||dz  k\  ||dz  k  z  }n||dz  k\  }t        j                  ||   dz        }t        j                  |dz        d	z   }||z  }t        || j                  z
        }t        j                  |      }| j                  d
k(  r,||   t        j                  || j                  z
        z  ||<   nf| j                  dk(  r,||   t        j                  | j                  |z
        z  ||<   n+||   t        j                  || j                  z
        z  ||<   t        j                  j                  |      }t        j                  t        j                  j!                  |            }t#        | j$                  j&                  ||j                         d t)        |       | j                  || j                  ||dt*        j,                  | j$                  j.                  t1        |d            S )Nrl   rm   rm   rm   r[   lowg      ?midg      ?ro   high)r   energy_ratior   band_energytotal_energyr)   r_   )r`   ru   r:   fftfft2fftshiftr   ogridrd   r   r   r   r   sign	ifftshiftrealifft2r5   rF   r'   ry   r   r   r*   re   )rJ   rL   rM   rN   rz   r   	fft_shift	magnitudeHWcenter_ycenter_xr}   r~   distancemax_distmaskr   r   ratiorg   	grad_maskgrad_fftgrad_spatials                           r$   rS   zSpectralBandLoss.compute  s   ;;=#&..r26 ffkk(#FFOOC(	FF9%	 1!VQ!V(xxBQB177AL1,Hq/@@A778Q;145 99ho-DYY%4/Hx$4NODx$.D ffYt_12vvi1n-4l* !2!223
 MM),	99'o@Q@Q8Q0RRIdOYY%'o8I8IE8Q0RRIdO'o@Q@Q8Q0RRIdO 66##I.wwrvv||H56!!!'')+3u:6		 % $ 1 1* , "**;;%% S1
 	
r#   )r   r-   rW   )r   r   r   r   r&   r0   r2   rK   r:   r;   r5   rS   r   r   s   @r$   r   r      sP    	)z ) )U )
=
RZZ =
BJJ =
"** =
`j =
r#   r   c            	            e Zd ZdZd
dedef fdZddej                  dej                  dej                  de	fd	Z
 xZS )SpectralSkewnessLossz
    Loss based on spectral skewness.
    
    Penalizes asymmetric frequency distributions which indicate
    systematic bias in certain directions/frequencies.
    rF   target_skewnessc                 2    t         |   |       || _        y rH   )r   rK   r   )rJ   rF   r   r   s      r$   rK   zSpectralSkewnessLoss.__init__X       .r#   rL   rM   rN   rO   c                    |j                         d d j                  dd      }t        j                  j	                  |      }t        j                  j                  |      }t        j                  |      }t        j                  |      }t        j                  |      dz   }	t        j                  |d      }
t        j                  |
t        j                  |
      z
  |	z  dz        }t        j                  |d      }t        j                  |t        j                  |      z
  |	z  dz        }t        |      t        |      z   dz  }t        || j                  z
        }|t        j                  |      z  }t        | j                  j                  ||j                         d t        |       |||| j                  d	t        j                   | j                  j"                  t%        |d
      d
z        S )Nrl   rm   ro   r   rp      r]   r[   )h_skewv_skewcombined_skewr,   r\   r_   )r`   ru   r:   r   r   r   r   ra   stdr   r   r5   rF   r'   ry   r   r   r*   re   )rJ   rL   rM   rN   rz   r   r   r   mean_magstd_magh_meanr   v_meanr   skewnessrg   grads                    r$   rS   zSpectralSkewnessLoss.compute\  s   ;;=#&..r26 ffkk(#FFOOC(	FF9%	 779%&&#d* +6BGGFO3w>1DE +6BGGFO3w>1DE K#f+-2 D$8$889
 "''&/)!!ZZ\+3u:.  !)..	 "**;;%% 3/#5
 	
r#   r-   rW   r   r   r   r   r&   r2   rK   r:   r;   r5   rS   r   r   s   @r$   r   r   P  sI    /z /E /+
RZZ +
BJJ +
"** +
`j +
r#   r   c            	       f    e Zd ZdZddej
                  dej
                  dej
                  defdZy)	MeanErrorLossu   
    Loss based on mean error (bias correction).
    
    Penalizes non-zero mean error which indicates
    systematic over/under-estimation.
    
    Target: mean ≈ 0
    NrL   rM   rN   rO   c                    t        j                  |      }t        |      }t        j                  |      t        j                  |      z  }t        | j                  j                  ||j                         d t        |       |t        j                  |      dt        j                  | j                  j                  t        t        |      dz  d            S )N)ra   r   
   r)   r_   )r:   ra   r   	ones_liker   r5   rF   r'   r`   ry   r   r   r   r*   re   )rJ   rL   rM   rN   
mean_errorrg   r7   s          r$   rS   zMeanErrorLoss.compute  s    WWU^
 _
 <<&)<<!!^^%ks5z2!+BFF5MB!--;;%% Z2!5s;
 	
r#   rW   rh   r"   r#   r$   r   r     s5    
RZZ 
BJJ 
"** 
`j 
r#   r   c            	            e Zd ZdZd
dedef fdZddej                  dej                  dej                  de	fd	Z
 xZS )SkewnessLossz
    Loss based on error distribution skewness.
    
    Encourages symmetric error distribution.
    Penalizes right-skewed (many small + few large positive errors)
    or left-skewed (many small + few large negative errors) distributions.
    rF   target_skewc                 2    t         |   |       || _        y rH   )r   rK   r   )rJ   rF   r   r   s      r$   rK   zSkewnessLoss.__init__       &r#   rL   rM   rN   rO   c                 R   |j                         }t        j                  |      }t        j                  |      dz   }t        j                  ||z
  |z  dz        }t	        || j
                  z
        }||z
  |z  }	|	dz  t        j                  || j
                  z
        z  }
t        | j                  j                  ||
j                         d t        |       || j
                  dt        j                  | j                  j                  t        t	        |      d      dz        S )Nro   r   r[   )r   r,         @r_   )r`   r:   ra   r   r   r   r   r5   rF   r'   ry   r   r   r*   re   )rJ   rL   rM   rN   
error_flatra   r   r   rg   standardizedr7   s              r$   rS   zSkewnessLoss.compute  s    [[]
wwz"ffZ 4'77Z$.#5!;< D$4$445
 #T)S01$rwwx$:J:J/J'KK!!^^%ks5z2%-9I9IJ!--;;%% X4s:
 	
r#   r   rW   r   r   s   @r$   r   r     sI    'z ' '
RZZ 
BJJ 
"** 
`j 
r#   r   c            	            e Zd ZdZd
dedef fdZddej                  dej                  dej                  de	fd	Z
 xZS )KurtosisLossu   
    Loss based on error distribution kurtosis.
    
    Encourages normal-like kurtosis (≈3).
    Penalizes heavy tails (>3) or light tails (<3).
    
    Heavy tails indicate outliers - important for robust learning.
    rF   target_kurtosisc                 2    t         |   |       || _        y rH   )r   rK   r   )rJ   rF   r   r   s      r$   rK   zKurtosisLoss.__init__  r   r#   rL   rM   rN   rO   c                 X   |j                         }t        j                  |      }t        j                  |      dz   }t        j                  ||z
  |z  dz        }t	        || j
                  z
        }||z
  |z  }	|	dz  t        j                  || j
                  z
        z  }
t        | j                  j                  ||
j                         d t        |       || j
                  dt        j                  | j                  j                  t        t	        |dz
        dz  d            S )Nro      r   )kurtosisr,   rn   r)   r_   )r`   r:   ra   r   r   r   r   r5   rF   r'   ry   r   r   r*   re   )rJ   rL   rM   rN   r   ra   r   r   rg   r   r7   s              r$   rS   zKurtosisLoss.compute  s   [[]
wwz"ffZ 4'77Z$.#5!;< D$8$889
 #T)S01$rwwx$:N:N/N'OO!!^^%ks5z2%-9M9MN!--;;%% X\!2Q!6<
 	
r#   )r   rW   r   r   s   @r$   r   r     sI    /z /E /
RZZ 
BJJ 
"** 
`j 
r#   r   c            	            e Zd ZdZd
dedef fdZddej                  dej                  dej                  de	fd	Z
 xZS )BimodalityLossz
    Loss to detect and penalize bimodal error distributions.
    
    Bimodal errors indicate the model is making two different
    types of mistakes - suggests a mixture of failure modes.
    
    Target: unimodal distribution
    rF   n_binsc                 2    t         |   |       || _        y rH   )r   rK   r   )rJ   rF   r   r   s      r$   rK   zBimodalityLoss.__init__  s     r#   rL   rM   rN   rO   c                 (   |j                         }t        j                  || j                  d      \  }}g }t	        dt        |      dz
        D ]6  }||   ||dz
     kD  s||   ||dz      kD  s!|j                  |||   f       8 t        |      dk\  rw|j                  d d       |d   |d   }
}	t        |	d   |
d         }t        |	d   |
d         }t        j                  |||dz          }|	d   |
d   z   dz  }d||d	z   z  z
  }nd
}t        d|dz
        }t        j                  |      }t        j                  ||      dz
  }t        j                  |d| j                  dz
        }t	        | j                        D ]  }||k(  }t        |      dk\  r|cxk  rk  r	n nd||<   *t        |      dk\  r||d   d   |d   d   fv rd||<   Pdt        j                  ||   t        j                  |      z
        z  ||<    t        | j                   j"                  ||j                         d t        |       |t        |      |d d D cg c]  }|d   	 c}dt$        j&                  | j                   j(                  |      S c c}w )NT)binsdensityr]   r[   c                     | d   S Nr]   r"   r~   s    r$   <lambda>z(BimodalityLoss.compute.<locals>.<lambda>  s
    QqT r#   keyreverser   r)   ro   r-   333333?g      皙?r   )bimodality_scoren_peakspeak_heightsr_   )r`   r:   	histogramr   rv   ry   rw   sortre   rb   r   digitizeclipr   ra   r5   rF   r'   r   r   r*   )rJ   rL   rM   rN   r   hist	bin_edgespeaksipeak1peak2valley_start
valley_endvalley_heightr  r  rg   r7   bin_idxbr   ps                         r$   rS   zBimodalityLoss.compute	  s   [[]
 ,,zTRi q#d)a-( 	+AAwac"tAwac':aa\*	+
 u:?JJ>4J8 8U1X5E uQxq2LU1XuQx0JFF4Z\#BCM "!HuQx/14L"m|d7J&KL" ,s23
 ==,++j)4q8'''1dkkAo6t{{# 	WAa<D5zQ<1#B
#B!$UqQ58A;a*D%D!%!$rwwz$/?"''*BU/U'V!V	W !!^^%ks5z2$4u:/4Ray 9!1 9
 "--;;%%-
 	
 !:s   J)   rW   )r   r   r   r   r&   r   rK   r:   r;   r5   rS   r   r   s   @r$   r   r     sI    z 3 :
RZZ :
BJJ :
"** :
`j :
r#   r   c            	            e Zd ZdZd
dedef fdZddej                  dej                  dej                  de	fd	Z
 xZS )ChannelErrorLossz
    Loss based on per-channel error magnitudes.
    
    Penalizes channels with high error, encouraging
    balanced error distribution across all channels.
    rF   target_equalc                 2    t         |   |       || _        y rH   )r   rK   r  )rJ   rF   r  r   s      r$   rK   zChannelErrorLoss.__init__R  s     (r#   rL   rM   rN   rO   c                    |j                         }|j                  dk(  r=t        j                  t        j                  t        j
                  |            g      }nCt        |j                  dz  d      }|d |dz   j                  |d      }|j	                  d      }t        j
                  |      }| j                  r9t        j                  |      }	t        j                  ||	z
  dz        }
||	z
  dz  }n*t        j                  |      }
t        j                  |      }t        | j                  j                  |
t        j                  |t        j                  |            j                  |j                        |t        j                  |      t        j                  |      t        j                   |      dt"        j$                  | j                  j&                  t        |
dz  d            S )	Nrl   r]   rp   r[   )mean_errorsmax_channel_errormin_channel_errorchannel_variancer   r)   r_   )r`   rc   r:   arrayra   r   rb   ru   r  r   r5   rF   r'   	full_liker   re   rx   r   r   r*   )rJ   rL   rM   rN   rf   channel_errors
n_channelsusabler  r,   rg   r7   s               r$   rS   zChannelErrorLoss.computeV  s   [[]
 ??c!XXrwwrvvj/A'B&CDN Z__3Q7J 1c!12:::sKF#[[a[0Nff^,WW[)F+"61!<=J $f,1H -Jww~.H!!\\*bggh.?@HHU*%'VVK%8%'VVK%8$&FF;$7	 "++;;%% b#6
 	
r#   TrW   )r   r   r   r   r&   r3   rK   r:   r;   r5   rS   r   r   s   @r$   r  r  J  sI    )z ) )'
RZZ '
BJJ '
"** '
`j '
r#   r  c            	            e Zd ZdZd
dedef fdZddej                  dej                  dej                  de	fd	Z
 xZS )CorrelationLossz
    Loss based on error correlation structure.
    
    Penalizes highly correlated errors across dimensions.
    Encourages errors to be independent (each dimension
    has independent noise).
    
    Target: identity covariance matrix
    rF   target_correlationc                 2    t         |   |       || _        y rH   )r   rK   r(  )rJ   rF   r(  r   s      r$   rK   zCorrelationLoss.__init__  s     "4r#   rL   rM   rN   rO   c                 "   |j                         }|j                  dkD  r|j                  dd      n|j                  dd      }|j                  d   dk  r_t	        | j
                  j                  dt        j                  |      ddddt        j                  | j
                  j                  d      S t        j                  |d	z         }|j                  d   }g }t        |      D ],  }	t        |	dz   |      D ]  }
|j                  ||	|
f           . t        |      dkD  rQt        j                   t        j"                  |            }t        j$                  t        j"                  |            }nd}d}t#        || j&                  z
        }||z  }t	        | j
                  j                  ||j                         d t        |       ||t        |      dt        j                  | j
                  j                  t)        |d
            S )Nrl   r]   r   r[   r-   )mean_correlationmax_correlationn_pairsr_   ro   r)   )r`   rc   ru   r   r5   rF   r'   r:   r   r   r   r*   corrcoefrv   rw   ry   ra   r   rb   r(  re   )rJ   rL   rM   rN   rf   rz   corr_matrixnoff_diagonalr  jmean_off_diagonalmax_off_diagonalrg   r7   s                  r$   rS   zCorrelationLoss.compute  s   [[]
2<//C2G:%%b#.ZM_M_`aceMf>>!q [[%%u-(+'* 
 &//{{))!$  kk(T/2 a q 	7A1Q3] 7##K1$567	7 |q  "|(< =!vvbff\&:; #" *T-D-DDE
 ,,!!^^%ks5z2$5#3|,
 "++;;%% !2C8
 	
r#   r   rW   r   r   s   @r$   r'  r'    sI    5z 5u 56
RZZ 6
BJJ 6
"** 6
`j 6
r#   r'  c            	            e Zd ZdZddededef fdZddej                  dej                  dej                  d	e	fd
Z
 xZS )EdgeErrorLossz
    Loss based on error at edge regions.
    
    Separately tracks error in edge regions vs smooth regions.
    Allows differential treatment of edge vs smooth errors.
    rF   edge_weightsmooth_weightc                 @    t         |   |       || _        || _        y rH   )r   rK   r8  r9  )rJ   rF   r8  r9  r   s       r$   rK   zEdgeErrorLoss.__init__  s      &*r#   rL   rM   rN   rO   c                 D   |j                         d d j                  dd      }|$|j                         d d j                  dd      }n8|4t        |d      r#|j                         d d j                  dd      n|}|}n|}t        j                  t        j
                  |d            }t        j                  t        j
                  |d            }||z   }	t        j                  |	d      }
|	|
kD  }| }t        j                  |      r+t        j                  t        j                  |      |         nd}t        j                  |      r+t        j                  t        j                  |      |         nd}| j                  |z  | j                  |z  z   }t        j                  |      }| j                  t        j                  ||         z  ||<   | j                  t        j                  ||         z  ||<   t        | j                  j                  ||j                         d t!        |       ||||dz   z  t        j"                  |      t        j"                  |      d	t$        j&                  | j                  j(                  t+        |d
z  d            S )Nrl   rm   r`   r   rp   r]   K   ro   )
edge_errorsmooth_error
edge_ratioedge_pixelssmooth_pixelsr   r)   r_   )r`   ru   hasattrr:   r   r7   r   anyra   r8  r9  r   r   r5   rF   r'   ry   r   r   r   r*   re   )rJ   rL   rM   rN   rz   ref_2dX_2dr   r   edge_magnitudeedge_threshold	edge_masksmooth_maskr=  r>  rg   r7   s                    r$   rS   zEdgeErrorLoss.compute  s9   ;;=#&..r26  __&t,44R<F]6=a6I1779Tc?**2r2qDFF F34F34& ~r:"^3	 j >@VVI=NRWWRVVH-i89TU
ACATrwwrvvh/<=Z[ %%
2T5G5G,5VV
 ==*"..)9L1MM $ 2 2RWWXk=R5S S!!^^%ks5z2( ,(L4,?@!vvi0!#!4 "));;%% d!2C8
 	
r#   )r)   r)   rW   r   r   s   @r$   r7  r7    sQ    +z + +TY +
0
RZZ 0
BJJ 0
"** 0
`j 0
r#   r7  c            	            e Zd ZdZddedef fdZddej                  dej                  dej                  de	fd	Z
dej                  d
ej                  dej                  fdZ xZS )StructuralSimilarityLossz
    Loss based on structural similarity (SSIM-like).
    
    Measures structural preservation between prediction
    and reference, not just pixel-level differences.
    rF   r|   c                 2    t         |   |       || _        y rH   )r   rK   r|   )rJ   rF   r|   r   s      r$   rK   z!StructuralSimilarityLoss.__init__  r   r#   rL   rM   rN   rO   c                 6   |j                         d d j                  dd      }|)|j                         d d j                  dd      }||z
  }n|}| }| j                  dz  }t        j                  ||d      }t        j                  ||d      }	t        j
                  | j                  | j                  f      | j                  dz  z  }
| j                  ||
      || || f   }| j                  |	|
      || || f   }|dz  }|dz  }||z  }| j                  |dz  |
      || || f   |z
  }| j                  |	dz  |
      || || f   |z
  }| j                  ||	z  |
      || || f   |z
  }d}d}d|z  |z   d|z  |z   z  }||z   |z   ||z   |z   z  }||dz   z  }t        j                  |      }d	|z
  }|d
|z
  z  }t        | j                  j                  ||j                         d t        |       |t        j                  ||z
  dz        t        j                  ||z
  dz        t        j                  |t        j                  ||z        dz   z  dz        dt        j                  | j                  j                  |      S )Nrl   rm   r[   reflect)modeg-C6?gH}M?ro   r)   r]   )ssimluminance_componentcontrast_componentstructure_componentr_   )r`   ru   r|   r:   padones_convolve2dra   r5   rF   r'   ry   rd   r   r   r*   )rJ   rL   rM   rN   rz   original	predictedrT  original_paddedpredicted_paddedkernelmu_originalmu_predictedmu_original_sqmu_predicted_sqmu_original_predictedsigma_original_sqsigma_predicted_sqsigma_original_predictedC1C2	numeratordenominatorrP  	mean_ssimrg   r7   s                              r$   rS   z StructuralSimilarityLoss.compute  sA   ;;=#&..r26   (#.66r2>H 8+IH!	I !#&&3Y?66)SyA $**D,<,<=>$BRBRVWBWX&&?SD#sd(@RS''(8&A#sd(CQTPTHBTU$)&!+ +l : ,,_-A6J3PSt8UXZ]Y]U]K]^aoo!--.>!.CVLSRUQUXWZ\_[_W_M_`crr#'#3#3OFV4VX^#_`cehdh`hjmornrjr`r#s  wL  $L   ..3<T8TWY8YZ	%7"<ARUgAgjlAlmK$./GGDM	 9_
 q9}-!!^^%ks5z2!')wwl0Jq/P'Q&(gg/@CU/UZ[.[&\')ww0HBGGTehzTzL{  C  MC  1D  IJ  0J  (K	 "));;%%'
 	
r#   r[  c                    t        j                  t        |j                  d   t	        |      z
  dz         D cg c]n  }t        |j                  d   t	        |      z
  dz         D cg c];  }t        j
                  |||t	        |      z   ||t	        |      z   f   |z        = c}p c}}      S c c}w c c}}w )zSimple 2D convolution.r   r]   )r:   r   rv   r   ry   r   )rJ   rM   r[  r}   r~   s        r$   rV  z$StructuralSimilarityLoss._convolve2dW  s    xx
 1771:F3a78

  qwwqzCK7!;< q1S[=!Ac&kM/9:VCD
  	
s   .B:
(A B5(B:
5B:
)rn   rW   )r   r   r   r   r&   r   rK   r:   r;   r5   rS   rV  r   r   s   @r$   rK  rK    sp    'z ' '=
RZZ =
BJJ =
"** =
`j =
~RZZ  

 r#   rK  c            	            e Zd ZdZddeddf fdZddej                  dej                  dej                  d	efd
Z	 xZ
S )AdversarialLossa  
    Adversarial loss using discriminator to judge error quality.
    
    Generator tries to minimize error in ways that fool discriminator.
    Discriminator learns to distinguish "good" errors from "bad" errors.
    
    This creates a game-theoretic formulation where the model
    learns not just to minimize numerical error, but to produce
    errors that look "natural" to the discriminator.
    rF   discriminatorErrorDiscriminatorc                 2    t         |   |       || _        y rH   )r   rK   rl  )rJ   rF   rl  r   s      r$   rK   zAdversarialLoss.__init__r  s     *r#   rL   rM   rN   rO   c                    | j                   t        j                  |dz        }d|z  }nC| j                   j                  |      \  }}|r|}nd|z
  }| j                   j	                  |      }t        | j                  j                  ||j                         d t        |       dt        j                  | j                  j                  t        |d            S )Nr[   r)   )discriminator_scoreis_realr_   )rl  r:   ra   judgerV   r5   rF   r'   r`   ry   r   r    r*   re   )rJ   rL   rM   rN   rg   r7   quality_scorerq  s           r$   rS   zAdversarialLoss.computev  s    %!,J5yH &*%7%7%=%=e%D"M7 *
 =0
 ))66u=H!!^^%ks5z20='R!--;;%% S1
 	
r#   rH   rW   )r   r   r   r   r&   rK   r:   r;   r5   rS   r   r   s   @r$   rk  rk  f  sJ    	+z +:N +
RZZ 
BJJ 
"** 
`j 
r#   rk  c                       e Zd ZdZddefdZdej                  defdZ	dej                  de
eef   fdZddej                  d	ej                  d
efdZdej                  dej                  fdZy)rm  a  
    Discriminator that judges error quality.
    
    Learned to distinguish:
    - Natural/easy errors (low frequency, uniform, symmetric)
    - Artificial/hard errors (high frequency, clustered, asymmetric)
    
    Used in adversarial training to guide error generation.
    hidden_sizec                 ,   t         j                  j                  d|      dz  | _        t        j                  d|f      | _        t         j                  j                  |d      dz  | _        t        j                  d      | _        d| _        d| _	        y )Nrl   {Gz?r]   )r]   r]         ?)
r:   randomrandnW1zerosb1W2b2
real_score
fake_score)rJ   ru  s     r$   rK   zErrorDiscriminator.__init__  sk    ))//#{3d:((A{+,))//+q1D8((6"r#   rL   rO   c                 N   |j                         dd }t        j                  || j                        | j                  z   }t        j
                  d|      }t        j                  || j                        | j                  z   }t        t        j                  |d   dd            S )z&Judge error quality, return score 0-1.Nrl   r   )r   r   r]   )
r`   r:   dotr{  r}  maximumr~  r  r2   r	  )rJ   rL   r   z1a1scores         r$   forwardzErrorDiscriminator.forward  sz    [[]4C(
VVJ(4772ZZ2r477#dgg-RWWU4[!Q/00r#   c                 >    | j                  |      }|dkD  r|dfS |dfS )z%Judge error, return (score, is_real).rx  TF)r  )rJ   rL   r  s      r$   rr  zErrorDiscriminator.judge  s-    U# 3;$;%<r#   real_errorsfake_errorslrc                 $   t        j                  |D cg c]  }| j                  |       c}      }t        j                  |D cg c]  }| j                  |       c}      }t        j                  |dz
  dz        }t        j                  |dz        }||z   }	| xj                  |t        j                  |d      t        j                  |d      z
  z  dz  z  c_        | xj
                  |t        j                  |      t        j                  |      z
  z  dz  z  c_        t        t        j                  |            | _        t        t        j                  |            | _        |	| j                  | j                  fS c c}w c c}w )z7Train discriminator to distinguish real vs fake errors.r]   r[   r   rp   rw  )	r:   r   r  ra   r{  r}  r2   r  r  )
rJ   r  r  r  ereal_scoresfake_scores	real_loss	fake_lossr?   s
             r$   
train_stepzErrorDiscriminator.train_step  s,    hhEAQEFhhEAQEF GG[1_23	GGK1,-	*
 	215RS8TTUX\\\2-0DDELL   45 454??DOO;;#  FEs   FFc                 ^    |j                         dd }| j                  |      }d|z
  |z  }|S )z%Get gradient for generator to reduce.Nrl   rx  )r`   r  )rJ   rL   r   r  
adjustments        r$   rV   zErrorDiscriminator.get_gradient  s:    [[]4C(
 U# EkZ/
r#   N)    )rw  )r   r   r   r   r   rK   r:   r;   r2   r  r   r3   rr  r  rV   r"   r#   r$   rm  rm    s    C 1RZZ 1E 1 2::  %t*<  <bjj <rzz <u <,"**  r#   rm  c            	            e Zd ZdZd
dedef fdZddej                  dej                  dej                  de	fd	Z
 xZS )GradientPenaltyLossz
    Gradient penalty to ensure smooth error surfaces.
    
    Penalizes large gradients (unstable predictions).
    Encourages smooth, well-behaved error landscapes.
    rF   penaltyc                 2    t         |   |       || _        y rH   )r   rK   r  )rJ   rF   r  r   s      r$   rK   zGradientPenaltyLoss.__init__  s     r#   rL   rM   rN   rO   c                    |j                         d d j                  dd      }t        j                  |d      }t        j                  |d      }t        j                  |dz  |dz  z         }t        j
                  |dz        }t        j                  |d      t        j                  |d      z   }	t        | j                  j                  ||	j                         d t        |       t        j
                  |      t        j                  |      t        j                  |      dt        j                  | j                  j                  t        |dz  d	      
      S )Nrl   rm   r   rp   r]   r[   )mean_gradientmax_gradientgradient_stdr   r)   r_   )r`   ru   r:   r7   rd   ra   r5   rF   r'   ry   rb   r   r   r!   r*   re   )
rJ   rL   rM   rN   rz   r   r   grad_magrg   	laplacians
             r$   rS   zGradientPenaltyLoss.compute  s   ;;=#&..r26 XA.XA. 77619vqy01 WWX]+
 KKQ/"++f12MM	!!__&{E
3!#!2 "x 0 "x 0
 "00;;%% d!2C8
 	
r#   r)   rW   r   r   s   @r$   r  r    sI    z E 
RZZ 
BJJ 
"** 
`j 
r#   r  c            	            e Zd ZdZd
dedef fdZddej                  dej                  dej                  de	fd	Z
 xZS )LipschitzPenaltyLossz
    Lipschitz penalty to ensure smooth function behavior.
    
    Penalizes large local changes (K-Lipschitz constraint).
    Encourages stable, predictable error surfaces.
    rF   Kc                 2    t         |   |       || _        y rH   )r   rK   r  )rJ   rF   r  r   s      r$   rK   zLipschitzPenaltyLoss.__init__  s     r#   rL   rM   rN   rO   c                 >   |j                         d d j                  dd      }d\  }}|d d dd f   |d d d df   z
  }|dd d d f   |d dd d f   z
  }t        j                  t        j                  |            }	t        j                  t        j                  |            }
t        |	|
      }t        d|| j
                  z
        }t        j                  |      }t        |      D ]^  }t        |dz
        D ]K  }t	        |||f         | j
                  kD  s!|||fxx   |||f   z  cc<   |||dz   fxx   |||f   z  cc<   M ` t        |dz
        D ][  }t        |      D ]K  }t	        |||f         | j
                  kD  s!|||fxx   |||f   z  cc<   ||dz   |fxx   |||f   z  cc<   M ] t        | j                  j                  ||j                         d t        |       || j
                  t        j                  t        j                  |      | j
                  kD        t        j                  t        j                  |      | j
                  kD        z   dt        j                  | j                  j                  t!        |dz  d	      
      S )Nrl   rm   r   r]   r+  r   )max_diffr  
violationsg      @r)   r_   )r`   ru   r:   rb   r   r  r   rv   r5   rF   r'   ry   r   r   r!   r*   re   )rJ   rL   rM   rN   rz   r   r   h_diffv_diff
max_h_diff
max_v_diffr  rg   r7   r}   r~   s                   r$   rS   zLipschitzPenaltyLoss.compute  sv   ;;=#&..r26 1 !QR%8AssF#33!"a%8CRCF#33 VVBFF6N+
VVBFF6N+
z:. Htvv-.
 ==*q 	7A1q5\ 7vad|$tvv-QTNfQTl2NQAX&&A,6&7	7 q1u 	7A1X 7vad|$tvv-QTNfQTl2NQUAX&&A,6&7	7 !!^^%ks5z2$VV ffRVVF^dff%<=rvvf~X\X^X^G^@__
 "00;;%% c!137
 	
r#   r  rW   r   r   s   @r$   r  r    sI    z e .
RZZ .
BJJ .
"** .
`j .
r#   r  c                      e Zd ZdZ	 	 ddeedf   deeee	f      fdZ
d Z	 	 ddej                  d	ej                  d
ej                  deeef   fdZ	 ddeeef   dedee	ej                  ef   fdZdeeef   dedee   fdZ	 	 	 ddej                  d	ej                  d
ej                  dedeej                  ef   f
dZy)MultiLossOptimizera(  
    Optimizer that combines multiple diagnostic losses.
    
    Key features:
    - Dynamic loss weighting based on which losses are large
    - Gradient normalization to prevent dominance
    - Automatic tuning of loss weights based on training progress
    - Adversarial training support
    Nr   .default_weightsc                     || _         i | _        t               | _        |xs5 i ddddddddddd	dd
ddddddddddddddddddddd| _        | j                          g | _        g | _        y )Nmser)   spatial_concentrationrx  spatial_entropyr  hotspot_penaltyspectral_highspectral_lowr   r   r   
bimodality皙?channel_errorcorrelationr=  structural_similarityr   gradient_penalty皙?lipschitz_penalty)r   loss_functionsrm  rl  r  _initialize_loss_functionsloss_historygradient_history)rJ   r   r  s      r$   rK   zMultiLossOptimizer.__init__Y  s   
 
 /1  /  
 3
33
#S3
 s3
 s	3

 S3
 C3
 #3
 3
 3
 #3
 S3
 33
 #3
 $S3
 33
  !3
"  #3
( 	'')  "r#   c                    t        t        dt        j                  d            | j                  d<   t        t        dt        j                  d            | j                  d<   t        t        dt        j                  d            | j                  d<   t        t        dt        j                  d      d	d
      | j                  d<   t        t        dt        j                  d      dd      | j                  d<   t        t        dt        j                  d      dd      | j                  d<   t        t        dt        j                  d      dd      | j                  d<   t        t        dt        j                  d      d      | j                  d<   t        t        dt        j                  d            | j                  d<   t        t        dt        j                  d      d      | j                  d<   t        t        dt        j                  d      d      | j                  d<   t        t        dt        j                  d      d       | j                  d<   t!        t        d!t        j"                  d      d"#      | j                  d!<   t%        t        d$t        j"                  d      d%      | j                  d$<   t'        t        d&t        j(                  d      d'd(      | j                  d&<   t+        t        d)t        j(                  d      d	*      | j                  d)<   t-        t        d+t        j.                  d      | j0                  ,      | j                  d+<   t3        t        d-t        j4                  d      d.      | j                  d-<   t7        t        d/t        j4                  d      d0      | j                  d/<   y1)2z1Initialize all loss functions with their configs.r  r)   )r*   r  rx  r  r  r  rn   r\   )r   r   r  r   r  )r   r   spectral_midr   r  r   spectral_skewnessr-   )r   r   r   )r   r   r   )r   r  r  r  )r   r  T)r  r  )r(  r=  g      ?)r8  r9  r  )r|   r   )rl  r  )r  r  )r  N)rY   r&   r   r   r  rj   r   r   r   r   r   r   r   r   r   r   r  r   r'  r7  r   rK  rk  r    rl  r  r!   r  )rJ   s    r$   r  z-MultiLossOptimizer._initialize_loss_functions}  sb    &-ul66sC&
E"
 8P.0D0DSQ8
34 2D(,*>*>sK2
-. 2D(,*>*>sK2
-. 0@(=(=cJc0
O, /?~|'<'<SIS/
N+ /?~|'<'<SIS/
N+ 4H*L,A,A#N4
/0 -:|\%=%=cJ-
L) +7z<#;#;CH+
J' +7z<#;#;CH+
J' -;|\%=%=cJ-
L) 0@(>(>sK0
O, .=}l&<&<SI".
M* -:|\%9%9#F3-
L) 8P.0D0DSQ8
34 .=}l&>&>sK,,.
M* 3F)<+F+FsS3
./ 4H*L,G,GPST4
/0r#   rL   rM   rN   rO   c                 N   i }| j                   j                         D ]  \  }}	 |j                  |||      }|||<     |S # t        $ r[}t        j                  d| d|        t        |dt        j                  |      i t        j                  dd      ||<   Y d}~d}~ww xY w)zCompute all enabled losses.zLoss z	 failed: r-   r_   N)r  itemsrS   	Exceptionloggerwarningr5   r:   r   r   r!   )	rJ   rL   rM   rN   resultsr'   loss_fnresultr  s	            r$   compute_all_lossesz%MultiLossOptimizer.compute_all_losses  s     !00668 	MD' 9= &	"   tfIaS9: *]]51 ")88%(!s   A  	B$	ABB$r>   dynamic_weightingc                    d}t        j                  t        |j                               d   j                        }i }i }|r|j                         D cg c]  }|j
                  dkD  s|j                    }}|rot        |      }	t        |      }
|j                         D ]F  \  }}|j
                  dkD  s|j                  |
z
  |	|
z
  dz   z  }|j
                  d|z   z  }||_        H |j                         D ]  \  }}|j
                  dkD  s|j                  |j
                  z  }||z  }|||<   t         j                  j                  |j                        dz   }|j                  |z  }|||j
                  z  z  }|||<    |rt        ||j                        nd}t        |||||| j                  ||            }|||fS c c}w )a/  
        Aggregate all losses into single loss and gradient.
        
        Args:
            losses: Dict of loss results
            dynamic_weighting: If True, adjust weights based on loss magnitudes
            
        Returns:
            Tuple of (total_loss, combined_gradient, state)
        r-   r   ro   r]   )r   r  )r>   r?   r@   rA   rB   rC   )r:   r   listvaluesr7   r*   r9   rb   re   r  r6   linalgnormgetr=   _generate_advice)rJ   r>   r  r?   combined_gradientr@   rB   lloss_valuesmax_lossmin_lossr'   lossr   adjusted_weightweighted_loss	grad_normnormalized_gradientrA   states                       r$   aggregate_lossesz#MultiLossOptimizer.aggregate_losses  s    
MM$v}}*?*B*K*KL!# 7=}}W!!((UV,1--WKW{+{+"(,,. 6JD${{Q&*&;&;h&F8V^K^aeKe%f
*.++Z*H&56 !,,. 	1JD${{Q $

T[[ 8m+
/<&t, IINN4==9D@	&*mmi&?#!%84;;%FF!'0t$	1 Xn28N8R8RSsx !#9') $ 5 5fm L
 ,e33O Xs   G)GrA   c                    g }|j                         D cg c]%  \  }}|j                  t        j                  k(  s$|' }}}|j                         D cg c]%  \  }}|j                  t        j                  k(  s$|' }}}|j                         D cg c]%  \  }}|j                  t        j
                  k(  s$|' }}}|rCt        j                  |D cg c]  }|j                   c}      }	|	dkD  r|j                  d       |r5t        d |D        d      }
|
r |
j                  dkD  r|j                  d       |rjt        d |D        d      }|r |j                  dkD  r|j                  d	       t        d
 |D        d      }|r |j                  dkD  r|j                  d       |dk(  r|j                  d       |S |dk(  r|j                  d       |S |dk(  r|j                  d       |S c c}}w c c}}w c c}}w c c}w )z9Generate optimization advice based on current loss state.rx  z>Focus on reducing spatial concentration - errors are clusteredc              3   >   K   | ]  }d |j                   v s|  yw)r   Nr'   .0r  s     r$   	<genexpr>z6MultiLossOptimizer._generate_advice.<locals>.<genexpr>D  s     MAFaff<LaM   Ng333333?z>High-frequency errors dominant - smooth the prediction surfacec              3   >   K   | ]  }d |j                   v s|  yw)r  Nr  r  s     r$   r  z6MultiLossOptimizer._generate_advice.<locals>.<genexpr>J  s     PQ9OqPr  r  z9Bimodal error detected - model has multiple failure modesc              3   \   K   | ]$  }d |j                   v sd|j                   vs!| & yw)r   r   Nr  r  s     r$   r  z6MultiLossOptimizer._generate_advice.<locals>.<genexpr>N  s*     i1zQVV7KPZbcbhbhPhQis   ,,,z4Error distribution is skewed - apply bias correctionr  z/Focus training on identified high-error regionsr=  zImprove edge rendering accuracyr   zApply global offset correction)r  r(   r   r   r   r   r:   ra   r9   rw   next)rJ   r>   rA   advicer1  r  spatial_lossesspectral_lossesstat_lossesavg_spatial	high_freqr  r   s                r$   r  z#MultiLossOptimizer._generate_advice3  s    )/]1!**H\H\:\!]])/_A1::I^I^;^1__%+\\^^TQqzz\E]E]7]q^^ ''~"N!1#5#5"NOKS ^_ MMtTIY77#=^_ P+PRVWJj99C?YZiikopHH55;TU --MMKL  l*MM;<  l*MM:;E ^_^ #Os(   %G3G3%G9<G9%G?<G?Hc                 \    | j                  |||      }| j                  ||      \  }}}||fS )z
        Get combined gradient from all losses.
        
        This is the main entry point for gradient-based optimization.
        )r  r  )	rJ   rL   rM   rN   r  r>   _r  r  s	            r$   get_combined_gradientz(MultiLossOptimizer.get_combined_gradient\  s>     ((9=&*&;&;FDU&V#e %''r#   )rm   rm   r]   NrW   r%  )NNT)r   r   r   r   r   r   r	   r   r0   r2   rK   r  r:   r;   r5   r  r3   r=   r  r   r  r  r"   r#   r$   r  r  N  sQ    "-6:"#S#X"# "$sEz"23"#HY
|  $	zz :: ::	
 
c:o	> #'=4S*_%=4  =4 
ubjj.0	1	=4~'tCO'< 'S 'UYZ]U^ 'X  $"&(zz( ::( ::	(
  ( 
rzz>)	*(r#   r  c                   "   e Zd ZdZ	 	 	 	 ddedededeedf   fdZdej                  d	ej                  fd
Z	dej                  d	ej                  fdZ
dej                  d	ej                  fdZdej                  dej                  dej                  d	efdZdej                  dej                  fdZddej                  dej                  defdZdej                  d	ej                  fdZdej                  dej                  d	efdZy)MultiLossMLPClassifierz
    MLP Classifier using multi-loss diagnostic learning.
    
    Instead of single MSE loss, uses rich diagnostic losses
    to learn from multiple error aspects simultaneously.
    
input_sizeru  output_sizer   .c                    || _         || _        || _        || _        t        j
                  j                  ||      dz  | _        t	        j                  d|f      | _	        t        j
                  j                  ||      dz  | _
        t	        j                  d|f      | _        t	        j                  g d      | _        t        |      | _        g | _        y )Nrw  r]   )rw  rw  rw  rw  r   )r  ru  r  r   r:   ry  rz  r{  r|  r}  r~  r  r   r  r  
multi_lossreference_bank)rJ   r  ru  r  r   s        r$   rK   zMultiLossMLPClassifier.__init__z  s     %&&
 ))//*k:TA((A{+,))//+{;dB((A{+,((34 -59 !r#   rM   rO   c                 b    |j                   dkD  r|j                  |j                  d   d      S |S )z(Flatten image-shaped inputs for the MLP.r[   r   r+  )ndimru   r   rJ   rM   s     r$   _prepare_inputz%MultiLossMLPClassifier._prepare_input  s+    66A:99QWWQZ,,r#   c                    | j                  |      }|dz  dz
  }t        j                  || j                        | j                  z   | _        t        j                  | j
                        t        j                  t        j                  | j
                        dz         z  | _	        t        j                  | j                  | j                        | j                  z   | _        | j                  | j                        | _        | j                  S )N     _@r)   ro   )r  r:   r  r{  r}  r  sinrd   r   r  r~  r  z2_softmaxoutput)rJ   rM   X_norms      r$   r  zMultiLossMLPClassifier.forward  s    "US&&)DGG3&&/BGGBFF477Od,B$CC&&$''*TWW4mmDGG,{{r#   r~   c                     t        j                  |t        j                  |dd      z
        }|t        j                  |dd      z  S )Nr]   Trq   keepdims)r:   exprb   r   )rJ   r~   exp_xs      r$   r  zMultiLossMLPClassifier._softmax  s:    q266!!d;;<rvve!d;;;r#   y_truey_predc                    | j                  |      }|j                  d   }|dz  dz
  }||z
  }t        j                  | j                  j
                  |      |z  }t        j                  |dd      |z  }t        j                  || j                  j
                        }	t        j                  |d      }
| j                  j                  |
|d      \  }}t        j                  | j                        d	z   }t        j                  |      }|t        j                  | j                        z  t        j                  | j                        d
|z  z  t        j                   | j                        z  z   }|	|z  }|d| j"                   }t        j$                  |      }|dd|z  z   z  }t        j                  |j
                  |      |z  }t        j                  |dd      |z  }||||||dS )z
        Backward pass with multi-loss diagnostic gradients.
        
        Instead of just computing dz2 = y_pred - y_true,
        we analyze the error in detail and create multiple gradient signals.
        r   r  r)   Tr  rp   N)rM   rN   ro   r[   r  )dW1db1dW2db2
loss_statedz2)r  r   r:   r  r  Tr   r~  ra   r  r  r   r  rd   cosr   r  ru  tanh)rJ   rM   r  r  mr	  r  dW2_standarddb2_standardda1error_for_analysismulti_gradientr  abs_z1sqrt_abs_z1spder_derivdz1_standardmulti_gradient_reshapedmulti_modulation	dz1_multidW1_standarddb1_standards                         r$   backward_with_multi_lossz/MultiLossMLPClassifier.backward_with_multi_loss  s    "LLOUS vovvdggii-1vvcD9A= ffS$''))$  WWV!4 &*__%J%J &K &
"
 4'ggfo!BFF477O3rwwtww7G1{?7[_a_e_efjfmfm_n6nn [( #11B$2B2B"C77#:; !C#0@*@$@A	vvfhh	2Q6vvia$?!C  $
 	
r#   c                    | j                  |      }| j                  |      }| j                  |||      }| xj                  | j                  d   |d   z  z  c_        | xj
                  | j                  d   |d   z  z  c_        | xj                  | j                  d   |d   z  z  c_        | xj                  | j                  d   |d   z  z  c_        |d	   S )
Nr   r  r]   r  r[   r  r   r  r  )r  r  r*  r{  r  r}  r~  r  )rJ   rM   r  r  gradss        r$   updatezMultiLossMLPClassifier.update  s    "a--a@4771:e,,4771:e,,4771:e,,4771:e,,\""r#   X_trainy_trainepochsc                 2   t         j                  d       t        |      D ]s  }t        j                  j                  dt        |      d      }||   }||   }| j                  |t        j                  | j                        |         }|dz  dk(  ss| j                  |dd |dd       }	t         j                  d| d       t         j                  d	|	d
       t         j                  d|j                  d
       t         j                  d|j                          t        |j                  j                         d d      dd }
t         j                  d|
        |j                   sRt         j                  d|j                           v y)z'Train with detailed loss visualization.z/Training with multi-loss diagnostic learning...r   r   r  N  z
=== Epoch  ===z
Accuracy: .4fzTotal loss: Dominant loss: c                     | d   S r   r"   r   s    r$   r   zAMultiLossMLPClassifier.train_with_visualization.<locals>.<lambda>  s
    !A$ r#   Tr   r   zTop contributors: Advice: )r  inforv   r:   ry  randintry   r-  eyer  r  r?   rA   sortedr@   r  rC   )rJ   r.  r/  r0  r  idxX_batchy_batchr  accuracysorted_lossess              r$   train_with_visualizationz/MultiLossMLPClassifier.train_with_visualization  sl   EFv 	MA))##As7|S9CclGclG WbffT5E5E.Fw.OPJ 2v{::getngetnEl1#T23j#78l:+@+@*EFGoj.F.F-GHI !'55;;=& ! 1	!
 0@A 11KK(:+I+I*J KL7	Mr#   c                 N    t        j                  | j                  |      d      S )Nr]   rp   )r:   argmaxr  r  s     r$   predictzMultiLossMLPClassifier.predict  s    yyaq11r#   c                 P    t        j                  | j                  |      |k(        S rH   )r:   ra   rD  )rJ   rM   r  s      r$   r  zMultiLossMLPClassifier.score  s    wwt||A&011r#   N)rl   r   r   r  )r   )r   r   r   r   r   r   rK   r:   r;   r  r  r  r   r*  r-  rA  rD  r2   r  r"   r#   r$   r  r  r  sV    !,!! ! 	!
 S#X!2

 rzz  

 <"** < <A
::A
 

A
 

	A

 
A
F
#

 
#BJJ 
#M

 MRZZ MY\ MB2 2

 22rzz 22:: 2% 2r#   r  __main__z,=== Multi-Loss Diagnostic Learning Demo ===
r  r  *   rl   r  rx  i  i  r\   )uniform_small	clusteredr  skewedz:--- Computing Multi-Loss for Different Error Patterns ---
rm   z=== Error Pattern: r3  zMean: r4  z, Std: z
Loss breakdown:c                      | d   j                   S r   )r6   r   s    r$   r   r   >  s    qtzz r#   Tr   rw  z  z: z (normalized: z.2f)z
Total weighted loss: r5  r7  z%
--- Training MLP with Multi-Loss ---z../X_train.wavr]   r+  z../y_train.wav	   r   r   )r  ru  r  r   )r0  z../X_test.wavz../y_test.wavz
Final Test Accuracy: z-Data files not found. Running synthetic demo.r2  2   )^numpyr:   scipy.io.wavfiler   dataclassesr   r   typingr   r   r   r	   r
   enumr   loggingbasicConfigINFO	getLoggerr   r  r   r&   r5   r=   rE   rY   rj   r   r   r   r   r   r   r   r   r  r'  r7  rK  rk  rm  r  r  r  r  printr  ry  seedrz  r|  r  arangeconcatenatetest_errorsr  
error_namerL   ru   rz   upperra   r   r  r>   r;  r'   r  r6   r9   r  totalr7   r  rA   rC   r.  astyper   r/  X_train_img
classifierrA  X_testy_test
X_test_imgr  r?  r8  FileNotFoundErrorX_syntheticr9  y_syntheticr"   r#   r$   <module>ri     sj    ! ( 8 8     ',,/Z [			8	$&4 & : : :    # # #" " 
 
.// /d'
) '
T4
) 4
vN
' N
b7
+ 7
|
$ 
<$
# $
N%
# %
PH
% H
^3
' 3
lE
& E
X=
$ =
@S/ St*
& *
ZH H^(
* (
V:
+ :
B]( ](H	h2 h2^ z	
9: $+6J IINN2 -3RXXc]RVVIBIIcNS01C7 ".."))//#"6"<biiooc>RUX>X!YZ	K 

GH(..0 
E==R(#J$4$4$6#7t<=wrwwu~c*'&"&&-1DEF ..x8 	!# 5ISWX 	_JD$zzD 4&4::c"2.AVAVWZ@[[\]^	_
 ",!<!<V!Dx'c{34 3 3456$$HU66q9:;<16 

23%Q'(+33B<()!,q088=oob"b!4+	

 	++K+M o&q)11"c:'*Q.66s;^^BB2
##J7-hs^<=I L  Q=> iioodBA6ii''2t4+	

 	++KR+PQs   &CR A(S21S2