Use bootstrap logger

This commit is contained in:
ciaranbor
2026-01-06 10:51:20 +00:00
parent 8f93a1ff78
commit 26e5b03285
2 changed files with 5 additions and 1 deletions
@@ -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
+1 -1
View File
@@ -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: