Skip to content
Merged
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
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,16 @@ class SubmissionContext:
reservation_id: Optional[str] = None
quota_warnings: List[str] = field(default_factory=list)

@property
def tracker_metadata(self) -> Dict[str, str]:
metadata = {}
project = self.config.get("--tracker_project_name") or self.config.get("tracker_project_name")
if project:
metadata["tracker_project_name"] = project
if self.tracker_run_name:
metadata["tracker_run_name"] = self.tracker_run_name
return metadata


@dataclass
class SubmissionResult:
Expand Down Expand Up @@ -206,8 +216,7 @@ async def submit(
unified_job.output_url = f"https://huggingface.co/{hub_model_id}"
if snapshot_metadata:
unified_job.metadata["snapshot"] = snapshot_metadata
if ctx.tracker_run_name:
unified_job.metadata["tracker_run_name"] = ctx.tracker_run_name
unified_job.metadata.update(ctx.tracker_metadata)
if ctx.hardware_profile:
unified_job.metadata["hardware_profile"] = ctx.hardware_profile

Expand Down Expand Up @@ -296,8 +305,7 @@ async def _create_upload_job(
metadata: Dict[str, Any] = {"upload_id": job_id}
if snapshot_metadata:
metadata["snapshot"] = snapshot_metadata
if ctx.tracker_run_name:
metadata["tracker_run_name"] = ctx.tracker_run_name
metadata.update(ctx.tracker_metadata)
if ctx.hardware_profile:
metadata["hardware_profile"] = ctx.hardware_profile

Expand Down Expand Up @@ -327,8 +335,7 @@ async def _update_upload_job_from_provider(
metadata.setdefault("prediction_id", cloud_job.job_id)
if snapshot_metadata:
metadata["snapshot"] = snapshot_metadata
if ctx.tracker_run_name:
metadata["tracker_run_name"] = ctx.tracker_run_name
metadata.update(ctx.tracker_metadata)
if ctx.upload_id:
metadata["upload_id"] = ctx.upload_id
if ctx.hardware_profile:
Expand Down
2 changes: 1 addition & 1 deletion simpletuner/static/js/modules/cloud/index.js
Original file line number Diff line number Diff line change
Expand Up @@ -612,7 +612,7 @@ if (!window.cloudDashboardComponent) {
filtered = filtered.filter(j => {
const jobId = (j.job_id || '').toLowerCase();
const configName = (j.config_name || '').toLowerCase();
const trackerName = (j.metadata?.tracker_run_name || '').toLowerCase();
const trackerName = this.jobDisplayName(j).toLowerCase();
const status = (j.status || '').toLowerCase();
const provider = (j.provider || '').toLowerCase();
return jobId.includes(query) || configName.includes(query) ||
Expand Down
19 changes: 19 additions & 0 deletions simpletuner/static/js/modules/cloud/jobs.js
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,25 @@
*/

window.cloudJobMethods = {
async continueJob(job) {
const trainer = Alpine.store('trainer');
if (trainer.continuingJob || job.job_type !== 'local' || job.status !== 'failed') return;
trainer.continuingJob = job.job_id;
try {
const switched = await trainer.switchEnvironment(job.metadata?.env_name || job.config_name);
if (!switched) return;
trainer.switchTab('basic');
await htmx.ajax('GET', '/web/trainer/tabs/basic', { target: '#tab-content', swap: 'innerHTML' });
await Alpine.nextTick();
document.getElementById('runBtn').click();
} catch (error) {
console.error('Failed to continue job:', error);
window.showToast('Failed to continue job', 'error');
} finally {
trainer.continuingJob = null;
}
},

async loadJobs(syncActive = false) {
if (!syncActive) {
this.jobsLoading = true;
Expand Down
9 changes: 7 additions & 2 deletions simpletuner/static/js/modules/cloud/utilities.js
Original file line number Diff line number Diff line change
Expand Up @@ -46,7 +46,12 @@ window.cloudUtilityMethods = {
},

jobDisplayName(job) {
return job.metadata?.tracker_run_name || job.config_name || job.job_id.substring(0, 8);
const metadata = job.metadata || {};
const config = metadata.runtime_config || {};
const project = metadata.tracker_project_name || config['--tracker_project_name'] || config.tracker_project_name;
const run = metadata.tracker_run_name || config['--tracker_run_name'] || config.tracker_run_name ||
metadata.run_name || job.config_name || job.job_id.substring(0, 8);
return project ? `${project} / ${run}` : run;
},

formatDuration(seconds) {
Expand Down Expand Up @@ -204,7 +209,7 @@ window.cloudComputedProperties = {
filtered = filtered.filter(j => {
const jobId = (j.job_id || '').toLowerCase();
const configName = (j.config_name || '').toLowerCase();
const trackerName = (j.metadata?.tracker_run_name || '').toLowerCase();
const trackerName = window.cloudUtilityMethods.jobDisplayName(j).toLowerCase();
const status = (j.status || '').toLowerCase();
const provider = (j.provider || '').toLowerCase();

Expand Down
10 changes: 9 additions & 1 deletion simpletuner/templates/partials/cloud_job_card.html
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@
<i class="fas" :class="jobTypeIcon(job.job_type)" style="font-size: 0.7rem;"></i>
</span>
<span class="text-truncate fw-medium job-name" style="max-width: 140px;"
x-text="jobDisplayName(job)"></span>
:title="jobDisplayName(job)" x-text="jobDisplayName(job)"></span>
</div>
<!-- Timestamp -->
<div class="small text-muted" style="font-size: 0.7rem;" x-text="formatDate(job.created_at)"></div>
Expand Down Expand Up @@ -74,6 +74,14 @@
<i class="fas me-1" :class="statusIcon(job.status)" style="font-size: 0.55rem;"></i>
<span x-text="job.status"></span>
</span>
<template x-if="jobIsFailed(job) && job.job_type === 'local' && (job.metadata?.env_name || job.config_name)">
<button type="button" class="btn btn-sm btn-outline-success job-continue-btn"
:disabled="$store.trainer.continuingJob"
@click.stop="continueJob(job)"
@keydown.enter.stop @keydown.space.stop>
<i class="fas fa-play me-1"></i>Continue
</button>
</template>
<!-- Cost (if available) -->
<template x-if="job.cost_usd">
<span class="text-muted" style="font-size: 0.6rem;">
Expand Down
8 changes: 7 additions & 1 deletion simpletuner/templates/trainer_htmx.html
Original file line number Diff line number Diff line change
Expand Up @@ -780,6 +780,8 @@
this.environmentsLoading = false;
}
},
continuingJob: null,

async switchEnvironment(configName) {
try {
const sanitizedName = sanitizeConfigName(configName);
Expand Down Expand Up @@ -819,15 +821,19 @@
}
}));

await this.fetchActiveEnvironmentConfig();
if (!await this.fetchActiveEnvironmentConfig()) {
throw new Error('Failed to load activated configuration');
}
await this.loadDatasetsAfterEnvironmentChange();

window.showToast(`Switched to ${configName} configuration`, 'success');
// Refresh current tab content instead of reloading
await this.refreshAfterEnvironmentSwitch();
return true;
} catch (error) {
console.error('Error switching config:', error);
window.showToast('Failed to switch configuration', 'error');
return false;
}
},
async refreshAfterEnvironmentSwitch() {
Expand Down
20 changes: 20 additions & 0 deletions tests/js/cloud_utilities.test.js
Original file line number Diff line number Diff line change
Expand Up @@ -466,3 +466,23 @@ describe('cloudComputedProperties', () => {
});
});
});

describe('job project identity', () => {
test.each([
[{ tracker_project_name: 'portraits', tracker_run_name: 'run' }, 'portraits / run'],
[{ runtime_config: { '--tracker_project_name': 'portraits', '--tracker_run_name': 'run' } }, 'portraits / run'],
[{ runtime_config: { tracker_project_name: 'portraits', tracker_run_name: 'run' } }, 'portraits / run'],
[{ run_name: 'legacy-run' }, 'legacy-run'],
])('uses persisted tracker identity %j', (metadata, expected) => {
expect(window.cloudUtilityMethods.jobDisplayName({ job_id: '123456789', config_name: 'config', metadata })).toBe(expected);
});
});

test('project name search distinguishes jobs with the same run name', () => {
const jobs = ['portraits', 'landscapes'].map(project => ({
job_id: project, config_name: 'config', status: 'failed',
metadata: { runtime_config: { tracker_project_name: project, tracker_run_name: 'run' } },
}));
const getter = Object.getOwnPropertyDescriptor(window.cloudComputedProperties, 'filteredJobs').get;
expect(getter.call({ jobs, jobSearchQuery: 'LANDSCAPES' }).map(job => job.job_id)).toEqual(['landscapes']);
});
38 changes: 38 additions & 0 deletions tests/test_cloud_job_tracker_metadata.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,38 @@
"""Tracker identity survives each cloud job persistence path."""

import unittest
from unittest.mock import AsyncMock

from simpletuner.simpletuner_sdk.server.services.cloud.job_submission import JobSubmissionService, SubmissionContext


class CloudJobTrackerMetadataTestCase(unittest.IsolatedAsyncioTestCase):
async def test_upload_job_records_project_and_run(self):
for prefix in ("", "--"):
with self.subTest(prefix=prefix):
store = AsyncMock()
service = JobSubmissionService(store)
context = SubmissionContext(
config={f"{prefix}tracker_project_name": "portraits"},
dataloader_config=[],
tracker_run_name="shared-run",
)
await service._create_upload_job(context, "upload-1", {})
job = store.add_job.call_args.args[0]
self.assertEqual(job.metadata["tracker_project_name"], "portraits")
self.assertEqual(job.metadata["tracker_run_name"], "shared-run")

async def test_provider_update_preserves_project_identity(self):
from simpletuner.simpletuner_sdk.server.services.cloud.base import CloudJobInfo, CloudJobStatus

store = AsyncMock()
service = JobSubmissionService(store)
context = SubmissionContext(
config={"tracker_project_name": "portraits"}, dataloader_config=[], tracker_run_name="shared-run"
)
job = CloudJobInfo(job_id="provider-1", provider="replicate", status=CloudJobStatus.RUNNING, created_at="now")
await service._update_upload_job_from_provider("upload-1", job, context, {}, None, None, None)
metadata = store.update_job.call_args.args[1]["metadata"]
self.assertEqual(metadata["tracker_project_name"], "portraits")
self.assertEqual(metadata["tracker_run_name"], "shared-run")
self.assertEqual(metadata["prediction_id"], "provider-1")
104 changes: 104 additions & 0 deletions tests/test_webui_e2e.py
Original file line number Diff line number Diff line change
Expand Up @@ -3915,5 +3915,109 @@ def _config_created(_driver):
self.for_each_browser("test_save_as_preserves_config_and_repoints_identity", scenario)


class FailedLocalJobContinueTestCase(_TrainerPageMixin, WebUITestCase):
MAX_BROWSERS = 1

def _continue_scenario(self, driver, *, missing_config=False, config_load_failure=False):
driver.get(f"{self.base_url}/web/trainer#cloud")
self.dismiss_onboarding(driver)
self._show_configured_cloud_dashboard(driver)
driver.execute_script(
"""
window.__continueRequests = [];
window.__configLoadFailures = 0;
if (arguments[1]) {
const originalFetch = window.fetch;
window.fetch = (url, options) => {
if (new URL(url, window.location.origin).pathname === '/api/configs/test-config') {
window.__configLoadFailures++;
return Promise.resolve(new Response('{}', {status: 500}));
}
return originalFetch(url, options);
};
}
document.body.addEventListener('htmx:beforeRequest', event => {
if (event.detail.requestConfig.path === '/api/training/start') {
window.__continueRequests.push({
environment: Alpine.store('trainer').activeEnvironment,
parameters: {...event.detail.requestConfig.parameters},
});
event.preventDefault();
}
});
const comp = Alpine.$data(document.querySelector('#cloud-tab-content'));
comp.stopPolling();
comp.jobs = [{job_id: 'failed-job', job_type: 'local', provider: 'local',
status: 'failed', config_name: arguments[0], created_at: new Date().toISOString(),
metadata: {env_name: arguments[0], runtime_config: {
'--tracker_project_name': 'portraits', '--tracker_run_name': 'shared-run'
}}}];
comp.selectedJob = null;
window.__selectedJobs = 0;
comp.selectJob = () => { window.__selectedJobs++; };
""",
"missing-config" if missing_config else "test-config",
config_load_failure,
)
button = WebDriverWait(driver, 15).until(EC.element_to_be_clickable((By.CSS_SELECTOR, ".job-continue-btn")))
name = driver.find_element(By.CSS_SELECTOR, ".job-name")
self.assertEqual(name.text, "portraits / shared-run")
self.assertEqual(name.get_attribute("title"), "portraits / shared-run")
driver.execute_script("Alpine.$data(document.querySelector('#cloud-tab-content')).jobSearchQuery = 'PORTRAITS';")
self.assertEqual(len(driver.find_elements(By.CSS_SELECTOR, ".job-continue-btn")), 1)
self.dismiss_onboarding(driver)
driver.execute_script("arguments[0].scrollIntoView({block: 'center'});", button)
if missing_config:
button.send_keys(Keys.ENTER)
else:
button.click()
if config_load_failure:
WebDriverWait(driver, 15).until(
lambda d: d.execute_script(
"return window.__configLoadFailures > 0 && Alpine.store('trainer').continuingJob === null;"
)
)
self.assertEqual(driver.execute_script("return window.__continueRequests"), [])
self.assertEqual(driver.execute_script("return Alpine.store('trainer').activeEnvironment"), "test-config")
self.assertIsNone(driver.execute_script("return Alpine.store('trainer').activeEnvironmentConfig"))
self.assertFalse(driver.find_element(By.CSS_SELECTOR, ".job-continue-btn").get_property("disabled"))
elif missing_config:
WebDriverWait(driver, 15).until(
lambda d: not d.find_element(By.CSS_SELECTOR, ".job-continue-btn").get_property("disabled")
)
self.assertEqual(driver.execute_script("return window.__continueRequests"), [])
self.assertEqual(driver.execute_script("return Alpine.store('trainer').activeEnvironment"), "default")
else:
WebDriverWait(driver, 20).until(lambda d: d.execute_script("return window.__continueRequests.length === 1"))
request = driver.execute_script("return window.__continueRequests[0]")
self.assertEqual(request["environment"], "test-config")
self.assertEqual(request["parameters"]["--resume_from_checkpoint"], "checkpoint-100")
self.assertEqual(driver.execute_script("return Alpine.store('trainer').activeTab"), "basic")
self.assertEqual(driver.execute_script("return window.__selectedJobs"), 0)

def test_continue_selects_failed_job_config_and_uses_run_flow(self):
self.with_sample_environment()
config_path = self.config_dir / "test-config" / "config.json"
config = json.loads(config_path.read_text())
config["--resume_from_checkpoint"] = "checkpoint-100"
config_path.write_text(json.dumps(config))
self.seed_defaults(active_config="default")
self.for_each_browser("continue_failed_local_job", lambda driver, _: self._continue_scenario(driver))

def test_continue_does_not_run_when_config_activation_fails(self):
self.seed_defaults(active_config="default")
self.for_each_browser(
"continue_missing_config", lambda driver, _: self._continue_scenario(driver, missing_config=True)
)

def test_continue_does_not_run_when_activated_config_cannot_be_loaded(self):
self.with_sample_environment()
self.seed_defaults(active_config="default")
self.for_each_browser(
"continue_config_load_failure",
lambda driver, _: self._continue_scenario(driver, config_load_failure=True),
)


if __name__ == "__main__":
unittest.main()
Loading