Update W&B logging crash in MLX-LM-LORA (#316)
* update * Ensure all mx array or tensor-like values are safely converted to Python-native types * format
This commit is contained in:
@@ -32,12 +32,19 @@ class WandBCallback(TrainingCallback):
|
||||
self.wrapped_callback = wrapped_callback
|
||||
wandb.init(project=project_name, dir=log_dir, config=config)
|
||||
|
||||
def _convert_to_serializable(self, data: dict) -> dict:
|
||||
return {k: v.tolist() if hasattr(v, "tolist") else v for k, v in data.items()}
|
||||
|
||||
def on_train_loss_report(self, train_info: dict):
|
||||
wandb.log(train_info, step=train_info.get("iteration"))
|
||||
wandb.log(
|
||||
self._convert_to_serializable(train_info), step=train_info.get("iteration")
|
||||
)
|
||||
if self.wrapped_callback:
|
||||
self.wrapped_callback.on_train_loss_report(train_info)
|
||||
|
||||
def on_val_loss_report(self, val_info: dict):
|
||||
wandb.log(val_info, step=val_info.get("iteration"))
|
||||
wandb.log(
|
||||
self._convert_to_serializable(val_info), step=val_info.get("iteration")
|
||||
)
|
||||
if self.wrapped_callback:
|
||||
self.wrapped_callback.on_val_loss_report(val_info)
|
||||
|
||||
Reference in New Issue
Block a user