From 7f072cd6129f4c2b9d3b7bbce682cb731c534988 Mon Sep 17 00:00:00 2001 From: Vivek Kalyan Date: Tue, 28 Jul 2026 15:41:08 -0700 Subject: [PATCH] perf: Submit Tinker training requests in the same cycle --- src/art/tinker_native/backend.py | 17 ++++++++--------- 1 file changed, 8 insertions(+), 9 deletions(-) diff --git a/src/art/tinker_native/backend.py b/src/art/tinker_native/backend.py index 2aba95947..8641c3318 100644 --- a/src/art/tinker_native/backend.py +++ b/src/art/tinker_native/backend.py @@ -424,16 +424,15 @@ def remove_mask(datum: tinker.Datum) -> tinker.Datum: model_input=datum.model_input, loss_fn_inputs=loss_fn_inputs ) - forward_output = await self._tinker_train_call( - "forward_backward", - state.training_client.forward_backward( - [remove_mask(datum) for datum in datums], - loss_fn=loss_fn, - loss_fn_config=loss_fn_config, - ), + forward_future = state.training_client.forward_backward( + [remove_mask(datum) for datum in datums], + loss_fn=loss_fn, + loss_fn_config=loss_fn_config, ) - optim_output = await self._tinker_train_call( - "optim_step", state.training_client.optim_step(adam_params) + optim_future = state.training_client.optim_step(adam_params) + forward_output, optim_output = await asyncio.gather( + self._tinker_train_call("forward_backward", forward_future), + self._tinker_train_call("optim_step", optim_future), ) if forward_output.metrics: