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: