From 26e5b03285e63cb3cf79fc31e13ef98aa910e853 Mon Sep 17 00:00:00 2001 From: ciaranbor Date: Fri, 5 Dec 2025 14:17:09 +0000 Subject: [PATCH] Use bootstrap logger --- src/exo/worker/engines/mflux/shard_mflux.py | 4 ++++ src/exo/worker/engines/mflux/utils_mflux.py | 2 +- 2 files changed, 5 insertions(+), 1 deletion(-) diff --git a/src/exo/worker/engines/mflux/shard_mflux.py b/src/exo/worker/engines/mflux/shard_mflux.py index 2ffc2afe..66b523bb 100644 --- a/src/exo/worker/engines/mflux/shard_mflux.py +++ b/src/exo/worker/engines/mflux/shard_mflux.py @@ -4,6 +4,7 @@ from mflux.models.flux.variants.txt2img.flux import Flux1 from exo.shared.types.worker.shards import PipelineShardMetadata from exo.worker.engines.mflux.pipefusion.pipefusion import apply_pipefusion_transformer from exo.worker.engines.mlx.utils_mlx import mx_barrier +from exo.worker.runner.bootstrap import logger def shard_flux_transformer( @@ -26,6 +27,7 @@ def shard_flux_transformer( The model with sharded transformer """ model = apply_pipefusion_transformer(model, group, shard_metadata) + logger.info("applied pipefusion transformations") mx.eval(model.parameters()) @@ -33,6 +35,8 @@ def shard_flux_transformer( mx.eval(model) # Synchronize processes before generation to avoid timeout + logger.info("before barrier") mx_barrier(group) + logger.info("after barrier") return model diff --git a/src/exo/worker/engines/mflux/utils_mflux.py b/src/exo/worker/engines/mflux/utils_mflux.py index 1c678b7e..afebc0b0 100644 --- a/src/exo/worker/engines/mflux/utils_mflux.py +++ b/src/exo/worker/engines/mflux/utils_mflux.py @@ -1,4 +1,3 @@ -from loguru import logger from mflux.config.model_config import ModelConfig from mflux.models.flux.variants.txt2img.flux import Flux1 @@ -8,6 +7,7 @@ from exo.worker.download.download_utils import build_model_path from exo.worker.engines.mflux.distributed_flux import DistributedFlux1 from exo.worker.engines.mflux.shard_mflux import shard_flux_transformer from exo.worker.engines.mlx.utils_mlx import mlx_distributed_init +from exo.worker.runner.bootstrap import logger def initialize_mflux(bound_instance: BoundInstance) -> DistributedFlux1: