@@ -116,7 +116,11 @@ def global_step_transform(engine: Engine, event_name: str | Events) -> int:
116116 self .state_attributes = state_attributes
117117
118118 def _setup_output_metrics_state_attrs (
119- self , engine : Engine , log_text : bool | None = False , key_tuple : bool | None = True
119+ self ,
120+ engine : Engine ,
121+ log_text : bool | None = False ,
122+ key_tuple : bool | None = True ,
123+ flatten_sequences : bool = True ,
120124 ) -> dict [Any , Any ]:
121125 """Helper method to setup metrics and state attributes to log"""
122126 metrics_state_attrs = OrderedDict ()
@@ -144,7 +148,7 @@ def _setup_output_metrics_state_attrs(
144148 if self .state_attributes is not None :
145149 metrics_state_attrs .update ({name : getattr (engine .state , name , None ) for name in self .state_attributes })
146150
147- metrics_state_attrs_dict : dict [Any , str | float | numbers .Number ] = OrderedDict ()
151+ metrics_state_attrs_dict : dict [Any , str | float | numbers .Number | list [ numbers . Number ] ] = OrderedDict ()
148152
149153 def key_tuple_fn (parent_key : str | tuple [str , ...] | None , * args : str ) -> tuple [str , ...]:
150154 if parent_key is None :
@@ -161,19 +165,25 @@ def key_str_fn(parent_key: str, *args: str) -> str:
161165
162166 def handle_value_fn (
163167 value : str | int | float | numbers .Number | torch .Tensor ,
164- ) -> None | str | float | numbers .Number :
168+ ) -> None | str | float | numbers .Number | list [ numbers . Number ] :
165169 if isinstance (value , numbers .Number ):
166170 return value
167171 elif isinstance (value , torch .Tensor ) and value .ndimension () == 0 :
168172 return value .item ()
173+ elif not flatten_sequences and isinstance (value , Sequence ):
174+ scalar_values = [item .item () if isinstance (item , torch .Tensor ) else item for item in value ]
175+ if all (isinstance (item , numbers .Number ) for item in scalar_values ):
176+ return scalar_values
169177 else :
170178 if isinstance (value , str ) and log_text :
171179 return value
172180 else :
173181 warnings .warn (f"Logger output_handler can not log metrics value type { type (value )} " )
174182 return None
175183
176- metrics_state_attrs_dict = _flatten_dict (metrics_state_attrs , key_fn , handle_value_fn , parent_key = self .tag )
184+ metrics_state_attrs_dict = _flatten_dict (
185+ metrics_state_attrs , key_fn , handle_value_fn , parent_key = self .tag , flatten_sequences = flatten_sequences
186+ )
177187 return metrics_state_attrs_dict
178188
179189
@@ -182,13 +192,14 @@ def _flatten_dict(
182192 key_fn : Callable ,
183193 value_fn : Callable ,
184194 parent_key : str | tuple [str , ...] | None = None ,
195+ flatten_sequences : bool = True ,
185196) -> dict :
186197 items = {}
187198 for key , value in in_dict .items ():
188199 new_key = key_fn (parent_key , key )
189200 if isinstance (value , Mapping ):
190- items .update (_flatten_dict (value , key_fn , value_fn , new_key ))
191- elif any (
201+ items .update (_flatten_dict (value , key_fn , value_fn , new_key , flatten_sequences ))
202+ elif flatten_sequences and any (
192203 [
193204 isinstance (value , tuple ) and hasattr (value , "_fields" ), # namedtuple
194205 not isinstance (value , str ) and isinstance (value , Sequence ),
0 commit comments