Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 12 additions & 0 deletions flytekit/core/node.py
Original file line number Diff line number Diff line change
Expand Up @@ -232,7 +232,19 @@ def with_overrides(
logger.warning("This override is beta. We may want to revisit this in the future.")
if not isinstance(task_config, type(self.run_entity._task_config)):
raise ValueError("can't change the type of the task config")

# Add debugging for task config overrides
old_config = self.run_entity._task_config
logger.info(f"Applying task_config override on node {self.id}")
logger.info(f" Old config: {old_config}")
logger.info(f" New config: {task_config}")

# For Elastic PyTorch tasks, specifically log nnodes changes
if hasattr(old_config, 'nnodes') and hasattr(task_config, 'nnodes'):
logger.info(f" Elastic PyTorch nnodes override: {old_config.nnodes} -> {task_config.nnodes}")

self.run_entity._task_config = task_config
logger.info(f"Task config override applied successfully")

if container_image is not None:
assert_not_promise(container_image, "container_image")
Expand Down
39 changes: 38 additions & 1 deletion flytekit/tools/translator.py
Original file line number Diff line number Diff line change
Expand Up @@ -465,7 +465,44 @@ def get_serializable_node(
override_pod_spec = _serialize_pod_spec(
entity._pod_template, entity.flyte_entity._get_container(settings), settings
)
task_spec = get_serializable(entity_mapping, settings, entity.flyte_entity, options=options)

# Fix 2: Override-aware serialization
# Create a temporary copy of the task with overrides applied for serialization
# This ensures that get_custom() sees the correct configuration without affecting the original task
from flytekit import logger
import copy

task_entity = entity.flyte_entity

# Check if we need to create a temporary task copy with overrides applied
# This happens when the node's task config differs from the original task config
needs_override_copy = False
if hasattr(entity, 'run_entity') and hasattr(entity.run_entity, '_task_config'):
node_task_config = entity.run_entity._task_config
original_task_config = task_entity._task_config

# Compare configs to see if they differ (indicating an override was applied)
if node_task_config != original_task_config:
needs_override_copy = True
logger.info(f"Fix 2: Node {entity.id} has task config override, creating temporary copy for serialization")
logger.info(f" Original task config: {original_task_config}")
logger.info(f" Node task config: {node_task_config}")

# For Elastic PyTorch tasks, specifically log nnodes changes
if hasattr(original_task_config, 'nnodes') and hasattr(node_task_config, 'nnodes'):
logger.info(f" Elastic PyTorch nnodes override: {original_task_config.nnodes} -> {node_task_config.nnodes}")

if needs_override_copy:
# Create a shallow copy of the task entity and apply the node's task config
# This ensures get_custom() sees the overridden config during serialization
temp_task_entity = copy.copy(task_entity)
temp_task_entity._task_config = node_task_config

logger.info(f"Fix 2: Using temporary task copy with overridden config for serialization")
task_spec = get_serializable(entity_mapping, settings, temp_task_entity, options=options)
else:
# No override needed, use the original task entity
task_spec = get_serializable(entity_mapping, settings, task_entity, options=options)
node_model = workflow_model.Node(
id=_dnsify(entity.id),
metadata=entity.metadata,
Expand Down
13 changes: 11 additions & 2 deletions plugins/flytekit-kf-pytorch/flytekitplugins/kfpytorch/task.py
Original file line number Diff line number Diff line change
Expand Up @@ -507,21 +507,30 @@ def execute(self, **kwargs) -> Any:
return self._execute(**kwargs)

def get_custom(self, settings: SerializationSettings) -> Optional[Dict[str, Any]]:
if self._task_config.nnodes == 1:
# Get the current task config, which should include any overrides applied via with_overrides()
current_nnodes = self._task_config.nnodes

# Add debugging to help diagnose override issues
logger.info(f"PytorchElasticFunctionTask.get_custom() called with nnodes={current_nnodes} (type: {type(current_nnodes)})")

if current_nnodes == 1:
"""
Torch elastic distributed training is executed in a normal k8s pod so that this
works without the kubeflow train operator.
"""
logger.info("Using single-node execution (normal Kubernetes pod)")
return super().get_custom(settings)
else:
logger.info(f"Using multi-node execution (PyTorchJob) with nnodes={current_nnodes}")
from flyteidl.plugins.kubeflow.pytorch_pb2 import ElasticConfig

try:
from torch.distributed import run
except ImportError:
raise ImportError(TORCH_IMPORT_ERROR_MESSAGE)

min_nodes, max_nodes = run.parse_min_max_nnodes(str(self._task_config.nnodes))
min_nodes, max_nodes = run.parse_min_max_nnodes(str(current_nnodes))
logger.info(f"Parsed nnodes '{current_nnodes}' to min_nodes={min_nodes}, max_nodes={max_nodes}")

elastic_config = ElasticConfig(
rdzv_backend=self.rdzv_backend,
Expand Down
116 changes: 116 additions & 0 deletions plugins/flytekit-kf-pytorch/tests/test_elastic_task.py
Original file line number Diff line number Diff line change
Expand Up @@ -310,3 +310,119 @@ def test_task():
test_task()

assert e.value.timestamp is not None


def test_nnodes_override() -> None:
"""Test that nnodes overrides work correctly in get_custom()."""
from flytekit import workflow
from flytekit.configuration import SerializationSettings

# Create a task with default nnodes=2 (multi-node)
@task(task_config=Elastic(nnodes=2, nproc_per_node=1))
def multi_node_task():
return "done"

@workflow
def test_workflow():
# Override to single-node
return multi_node_task().with_overrides(
task_config=Elastic(nnodes=1, nproc_per_node=1)
)

# Get the workflow node
wf = test_workflow
node = list(wf.nodes)[0]

# Verify the override was applied to the task config
assert node.run_entity._task_config.nnodes == 1

# Test serialization - this is where the bug would manifest
settings = SerializationSettings(image_config=None)

# Original task should return custom config (PyTorchJob for multi-node)
original_custom = multi_node_task.get_custom(settings)
assert original_custom is not None, "Multi-node task should have custom config"

# Overridden task should return None (normal pod for single-node)
overridden_custom = node.run_entity.get_custom(settings)
assert overridden_custom is None, "Single-node override should have no custom config"

# Verify they're different (override took effect)
assert original_custom != overridden_custom, "Override should change the custom config"


def test_nnodes_override_reverse() -> None:
"""Test overriding from single-node to multi-node."""
from flytekit import workflow
from flytekit.configuration import SerializationSettings

# Create a task with default nnodes=1 (single-node)
@task(task_config=Elastic(nnodes=1, nproc_per_node=1))
def single_node_task():
return "done"

@workflow
def test_workflow():
# Override to multi-node
return single_node_task().with_overrides(
task_config=Elastic(nnodes=2, nproc_per_node=1)
)

# Get the workflow node
wf = test_workflow
node = list(wf.nodes)[0]

# Verify the override was applied to the task config
assert node.run_entity._task_config.nnodes == 2

# Test serialization
settings = SerializationSettings(image_config=None)

# Original task should return None (normal pod for single-node)
original_custom = single_node_task.get_custom(settings)
assert original_custom is None, "Single-node task should have no custom config"

# Overridden task should return custom config (PyTorchJob for multi-node)
overridden_custom = node.run_entity.get_custom(settings)
assert overridden_custom is not None, "Multi-node override should have custom config"

# Verify they're different (override took effect)
assert original_custom != overridden_custom, "Override should change the custom config"


def test_fix2_serialization_isolation() -> None:
"""Test Fix 2: Ensure serialization with overrides doesn't affect original task."""
from flytekit import workflow
from flytekit.configuration import SerializationSettings
from flytekit.tools.translator import get_serializable_workflow
from collections import OrderedDict

# Create a task with default nnodes=2 (multi-node)
@task(task_config=Elastic(nnodes=2, nproc_per_node=1))
def multi_node_task():
return "done"

@workflow
def test_workflow():
# Override to single-node
return multi_node_task().with_overrides(
task_config=Elastic(nnodes=1, nproc_per_node=1)
)

# Test serialization with Fix 2
settings = SerializationSettings(image_config=None)
entity_mapping = OrderedDict()

# Serialize the workflow - this should trigger Fix 2 logic
wf_spec = get_serializable_workflow(entity_mapping, settings, test_workflow)

# Verify that the original task config was not permanently modified
assert multi_node_task._task_config.nnodes == 2, "Original task config should remain unchanged"

# Verify that the node still has the override applied
node = list(test_workflow.nodes)[0]
assert node.run_entity._task_config.nnodes == 1, "Node should still have override applied"

# Verify that serialization produced the correct result
# The serialized workflow should contain a task with single-node config
assert wf_spec is not None, "Workflow serialization should succeed"
83 changes: 83 additions & 0 deletions test_elastic_override_fix.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,83 @@
#!/usr/bin/env python3
"""
Test to verify that the Elastic PyTorch task override fix works correctly.
This test demonstrates the issue and verifies the fix.
"""

import logging
from flytekitplugins.kfpytorch.task import Elastic
from flytekit import task, workflow
from flytekit.configuration import SerializationSettings

# Set up logging to see our debug messages
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)


@task(task_config=Elastic(nnodes=2, nproc_per_node=1))
def train_task():
"""A training task with default nnodes=2 (multi-node)."""
print("Training...")
return "done"


@workflow
def training_workflow():
"""Workflow that overrides the task to use nnodes=1 (single-node)."""
# This should override the task to use single-node execution
single_node_task = train_task().with_overrides(
task_config=Elastic(nnodes=1, nproc_per_node=1)
)
return single_node_task


def test_override_fix():
"""Test that demonstrates the fix for Elastic PyTorch task overrides."""
print("=== Testing Elastic PyTorch Task Override Fix ===\n")

# Get the workflow and its node
wf = training_workflow
node = list(wf.nodes)[0] # First (and only) node

print(f"Original task nnodes: {train_task._task_config.nnodes}")
print(f"Overridden node task nnodes: {node.run_entity._task_config.nnodes}")

# Test serialization - this is where the bug manifests
settings = SerializationSettings(image_config=None)

print("\n=== Testing Serialization ===")
print("Getting custom config for original task:")
original_custom = train_task.get_custom(settings)

print("\nGetting custom config for overridden task:")
overridden_custom = node.run_entity.get_custom(settings)

print(f"\nOriginal custom config: {original_custom}")
print(f"Overridden custom config: {overridden_custom}")

# Analyze the results
print("\n=== Results Analysis ===")

# Original task should use PyTorchJob (multi-node)
if original_custom is None:
print("✅ Original task correctly uses single-node (None custom config)")
else:
print("✅ Original task correctly uses multi-node (PyTorchJob custom config)")

# Overridden task should use normal pod (single-node)
if overridden_custom is None:
print("✅ SUCCESS: Overridden task correctly uses single-node execution!")
print(" The override from nnodes=2 to nnodes=1 is working properly.")
else:
print("❌ ISSUE: Overridden task still uses multi-node execution")
print(" The override is not taking effect properly.")

# Additional verification
if original_custom != overridden_custom:
print("✅ Override is taking effect - configurations are different")
else:
print("❌ Override not working - configurations are identical")


if __name__ == "__main__":
test_override_fix()
57 changes: 57 additions & 0 deletions test_elastic_override_issue.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,57 @@
#!/usr/bin/env python3
"""
Test script to reproduce the Elastic PyTorch task override issue.
"""

from flytekitplugins.kfpytorch.task import Elastic
from flytekit import task, workflow
from flytekit.configuration import SerializationSettings


# Create a task with default nnodes=2
@task(task_config=Elastic(nnodes=2, nproc_per_node=1))
def train_task():
print("Training...")
return "done"


@workflow
def training_workflow():
# Override the task to use nnodes=1
single_node_task = train_task().with_overrides(task_config=Elastic(nnodes=1, nproc_per_node=1))
return single_node_task


def test_serialization():
"""Test what happens during serialization."""
print("=== Original Task Config ===")
print(f"Original task nnodes: {train_task._task_config.nnodes}")

# Get the workflow
wf = training_workflow

# Get the node from the workflow
node = list(wf.nodes)[0] # First node
print(f"Node task config nnodes: {node.run_entity._task_config.nnodes}")

# Test serialization
settings = SerializationSettings(image_config=None)

print("\n=== Serialization Test ===")
print("Original task get_custom:")
original_custom = train_task.get_custom(settings)
print(f"Original custom config: {original_custom}")

print("\nOverridden task get_custom:")
overridden_custom = node.run_entity.get_custom(settings)
print(f"Overridden custom config: {overridden_custom}")

# Check if the issue exists
if original_custom == overridden_custom:
print("\n❌ ISSUE CONFIRMED: Override not taking effect in serialization!")
else:
print("\n✅ Override working correctly in serialization")


if __name__ == "__main__":
test_serialization()