@@ -490,10 +490,43 @@ def digest(cls, observation: Any, context: ObservationContext) -> Any:
490490 return observation
491491
492492 @staticmethod
493- def _build_psnr_extension (
493+ def _metric_specs () -> Dict [str , Dict [str , Any ]]:
494+ return {
495+ # Lower value is worse for PSNR/Cosine.
496+ "psnr" : {
497+ "name" : "Per-layer PSNR" ,
498+ "label" : "PSNR" ,
499+ "inverse" : True ,
500+ },
501+ "cosine_sim" : {
502+ "name" : "Per-layer Cosine Similarity" ,
503+ "label" : "Cosine" ,
504+ "inverse" : True ,
505+ },
506+ # Higher value is worse for error metrics.
507+ "mse" : {
508+ "name" : "Per-layer MSE" ,
509+ "label" : "MSE" ,
510+ "inverse" : False ,
511+ },
512+ "abs_err" : {
513+ "name" : "Per-layer AbsErr" ,
514+ "label" : "AbsErr" ,
515+ "inverse" : False ,
516+ },
517+ }
518+
519+ @classmethod
520+ def _build_metric_extension (
521+ cls ,
494522 rows : List [Dict [str , Any ]],
523+ metric_name : str ,
495524 ) -> GraphExtension :
496- ext = GraphExtension (id = "psnr" , name = "Per-layer PSNR" )
525+ spec = cls ._metric_specs ().get (metric_name )
526+ if not spec :
527+ raise ValueError (f"Unsupported per-layer metric extension: { metric_name } " )
528+
529+ ext = GraphExtension (id = metric_name , name = str (spec ["name" ]))
497530 for row in rows :
498531 node_id = str (row ["target_node" ])
499532 info = {
@@ -515,37 +548,36 @@ def _build_psnr_extension(
515548 ext .add_node_data (node_id , info )
516549
517550 ext .set_sync_key ("sparse_match_key" )
518- ext .set_sync_key ("from_node_root" )
519- ext .set_label_formatter (
520- lambda d : [
521- f"PSNR={ float (d .get ('psnr' , 0.0 )):.4f} " ,
522- f"Cos={ float (d .get ('cosine_sim' , 0.0 )):.4f} " ,
523- ]
524- )
525- ext .set_tooltip_formatter (
526- lambda d : [
527- f"match_key={ d .get ('sparse_match_key' , '' )} " ,
528- f"root={ d .get ('from_node_root' , 'n/a' )} " ,
529- f"anchor_node={ d .get ('anchor_node' , 'n/a' )} " ,
551+
552+ def _format_metric_value (value : float , * , tooltip : bool = False ) -> str :
553+ if metric_name in ("mse" , "abs_err" ):
554+ return f"{ value :.6e} " if tooltip else f"{ value :.3e} "
555+ return f"{ value :.6f} " if tooltip else f"{ value :.4f} "
556+
557+ def _label_formatter (d : Dict [str , Any ]) -> List [str ]:
558+ primary = float (d .get (metric_name , 0.0 ))
559+ primary_label = str (spec ["label" ])
560+ return [f"{ primary_label } ={ _format_metric_value (primary )} " ]
561+
562+ ext .set_label_formatter (_label_formatter )
563+
564+ def _tooltip_formatter (d : Dict [str , Any ]) -> List [str ]:
565+ primary = float (d .get (metric_name , 0.0 ))
566+ primary_label = str (spec ["label" ])
567+ return [
530568 f"target_node={ d .get ('target_node' , 'n/a' )} " ,
531- f"anchor_topo={ d .get ('anchor_topo_index' , - 1 )} " ,
532- f"target_topo={ d .get ('target_topo_index' , - 1 )} " ,
533- f"numel={ d .get ('numel_compared' , 0 )} " ,
534- f"shape(anchor)={ d .get ('anchor_shape' , 'n/a' )} " ,
535- f"shape(target)={ d .get ('target_shape' , 'n/a' )} " ,
536- f"PSNR={ float (d .get ('psnr' , 0.0 )):.6f} " ,
537- f"Cosine={ float (d .get ('cosine_sim' , 0.0 )):.6f} " ,
538- f"MSE={ float (d .get ('mse' , 0.0 )):.6e} " ,
539- f"AbsErr={ float (d .get ('abs_err' , 0.0 )):.6e} " ,
569+ f"match_key={ d .get ('sparse_match_key' , '' )} " ,
570+ f"{ primary_label } ={ _format_metric_value (primary , tooltip = True )} " ,
540571 ]
541- )
572+
573+ ext .set_tooltip_formatter (_tooltip_formatter )
542574 ext .set_color_rule (
543575 _MetricNumericColorRule (
544- attribute = "psnr" ,
545- # Low PSNR is severe -> darker red.
576+ attribute = metric_name ,
577+ # Severe values map to darker red.
546578 low_rgb = (254 , 224 , 210 ),
547579 high_rgb = (165 , 15 , 21 ),
548- inverse = True ,
580+ inverse = bool ( spec [ "inverse" ]) ,
549581 )
550582 )
551583 return ext
@@ -569,8 +601,12 @@ def analyze(records: List[RecordDigest], config: Dict[str, Any]) -> AnalysisResu
569601 }
570602 )
571603
572- psnr_ext = PerLayerAccuracyLens ._build_psnr_extension (rows )
573- analysis .add_graph_layer ("psnr" , psnr_ext )
604+ for metric_name in ("cosine_sim" ,):
605+ # TODO other options "psnr" "mse", "abs_err"
606+ metric_ext = PerLayerAccuracyLens ._build_metric_extension (
607+ rows , metric_name
608+ )
609+ analysis .add_graph_layer (metric_name , metric_ext )
574610
575611 result .per_record_data [record .name ] = analysis
576612
@@ -773,8 +809,8 @@ def record(
773809 id = "per_layer_accuracy_graph" ,
774810 title = "Per-layer Accuracy Graph" ,
775811 graph_ref = graph_ref ,
776- default_layers = [f"{ lens_name } /psnr " ],
777- default_color_by = f"{ lens_name } /psnr " ,
812+ default_layers = [f"{ lens_name } /cosine_sim " ],
813+ default_color_by = f"{ lens_name } /cosine_sim " ,
778814 order = 21 ,
779815 ).as_block (),
780816 HtmlBlock (
0 commit comments