Skip to content

Commit 2309367

Browse files
author
nightcityblade
committed
feat: optionally preserve sequences in WandB logs
1 parent 82bb998 commit 2309367

3 files changed

Lines changed: 29 additions & 7 deletions

File tree

ignite/handlers/base_logger.py

Lines changed: 17 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -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),

ignite/handlers/wandb_logger.py

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -174,6 +174,8 @@ class OutputHandler(BaseOutputHandler):
174174
uses function output as global_step. To setup global step from another engine, please use
175175
:meth:`~ignite.handlers.wandb_logger.global_step_from_engine`.
176176
sync: Deprecated, has no function. Argument is kept here for compatibility with existing code.
177+
flatten_sequences: Whether to flatten sequences into separate metrics. Set to ``False`` to pass sequences of
178+
scalar values directly to Weights & Biases.
177179
178180
Examples:
179181
.. code-block:: python
@@ -282,8 +284,10 @@ def __init__(
282284
global_step_transform: Callable[[Engine, str | Events], int] | None = None,
283285
sync: bool | None = None,
284286
state_attributes: list[str] | None = None,
287+
flatten_sequences: bool = True,
285288
):
286289
super().__init__(tag, metric_names, output_transform, global_step_transform, state_attributes)
290+
self.flatten_sequences = flatten_sequences
287291
if sync is not None:
288292
warn("The sync argument for the WandBLoggers is no longer used, and may be removed in the future")
289293

@@ -297,7 +301,9 @@ def __call__(self, engine: Engine, logger: WandBLogger, event_name: str | Events
297301
f"global_step must be int, got {type(global_step)}. Please check the output of global_step_transform."
298302
)
299303

300-
metrics = self._setup_output_metrics_state_attrs(engine, log_text=True, key_tuple=False)
304+
metrics = self._setup_output_metrics_state_attrs(
305+
engine, log_text=True, key_tuple=False, flatten_sequences=self.flatten_sequences
306+
)
301307
logger.log(metrics, step=global_step)
302308

303309

tests/ignite/handlers/test_wandb_logger.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -131,6 +131,11 @@ def test_output_handler_metric_names():
131131
wrapper(mock_engine, mock_logger, Events.ITERATION_STARTED)
132132
mock_logger.log.assert_called_once_with({f"tag/a/{i}": v for i, v in enumerate(data)}, step=7)
133133

134+
wrapper = OutputHandler("tag", metric_names=["a"], flatten_sequences=False)
135+
mock_logger.log = MagicMock()
136+
wrapper(mock_engine, mock_logger, Events.ITERATION_STARTED)
137+
mock_logger.log.assert_called_once_with({"tag/a": data}, step=7)
138+
134139
wrapper = OutputHandler("tag", metric_names="all")
135140
mock_engine = MagicMock()
136141
mock_engine.state = State(

0 commit comments

Comments
 (0)