Skip to content

Commit db08ade

Browse files
[bugfix] reset train metric when restoring a fine-tune checkpoint (#651)
Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
1 parent c8082ae commit db08ade

7 files changed

Lines changed: 53 additions & 7 deletions

File tree

docs/source/models/evaluation_metrics.md

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -60,6 +60,8 @@ model_config {
6060
- decay_rate: 默认为0.9,历史指标的衰减率。
6161
- decay_step: 默认为100,历史指标每训练多少step会进行一次衰减,配置需要可以被train_config.log_step_count_steps整除。
6262

63+
训练时指标会保存在checkpoint中,使用--continue_train续跑时会一并恢复。从fine_tune_checkpoint恢复时则会重置,避免继承上一个模型的指标,重置后在第一个decay_step之前打印为0。
64+
6365
______________________________________________________________________
6466

6567
## 指标详情

tzrec/main.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -348,6 +348,7 @@ def _train_and_evaluate(
348348
eval_result_filename: str = TRAIN_EVAL_RESULT_FILENAME,
349349
check_all_workers_data_status: bool = False,
350350
ignore_restore_optimizer: bool = False,
351+
restore_from_model_dir: bool = False,
351352
dataloader_state: Optional[Dict[str, Any]] = None,
352353
delta_embedding_dumper: Optional[DeltaEmbeddingDumper] = None,
353354
pipeline_config_path: Optional[str] = None,
@@ -526,6 +527,10 @@ def run_eval(step: int, epoch: int) -> None:
526527
train_config.fine_tune_ckpt_param_map,
527528
dense_ema=dense_ema,
528529
)
530+
if not restore_from_model_dir:
531+
# a fine-tune checkpoint's train metric describes the source
532+
# model, not the model this job is training.
533+
_model.reset_train_metric()
529534
if delta_embedding_dumper is not None:
530535
delta_embedding_dumper.clear()
531536

@@ -933,6 +938,7 @@ def train_and_evaluate(
933938
ckpt_path=ckpt_path,
934939
check_all_workers_data_status=check_all_workers_data_status,
935940
ignore_restore_optimizer=ignore_restore_optimizer,
941+
restore_from_model_dir=restore_from_model_dir,
936942
dataloader_state=dataloader_state,
937943
delta_embedding_dumper=delta_embedding_dumper,
938944
pipeline_config_path=os.path.join(pipeline_config.model_dir, "pipeline.config"),

tzrec/metrics/train_metric_wrapper.py

Lines changed: 11 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -26,18 +26,18 @@ class TrainMetricWrapper(nn.Module):
2626
decay_step (int): decay step for decay,
2727
"""
2828

29+
_value: Tensor
30+
_step_cnt: Tensor
31+
2932
def __init__(
3033
self, metric_module: Metric, decay_rate: float = 0.5, decay_step: int = 100
3134
) -> None:
3235
super().__init__()
3336
self._decay_rate = decay_rate
3437
self._decay_step = decay_step
3538
self._metric_module = metric_module
36-
self._value = nn.Parameter(torch.tensor(0.0), requires_grad=False)
37-
self._step_total_value = nn.Parameter(torch.tensor(0.0), requires_grad=False)
38-
self._step_cnt = nn.Parameter(
39-
torch.tensor(0, dtype=torch.int), requires_grad=False
40-
)
39+
self.register_buffer("_value", torch.tensor(0.0))
40+
self.register_buffer("_step_cnt", torch.tensor(0, dtype=torch.int))
4141

4242
def update(self, preds: Tensor, target: Tensor) -> None:
4343
"""Update metric module."""
@@ -60,3 +60,9 @@ def update(self, preds: Tensor, target: Tensor) -> None:
6060
def compute(self) -> Tensor:
6161
"""Get metric value."""
6262
return self._value.data
63+
64+
def reset(self) -> None:
65+
"""Reset metric state."""
66+
self._metric_module.reset()
67+
self._value.zero_()
68+
self._step_cnt.zero_()

tzrec/metrics/train_metric_wrapper_test.py

Lines changed: 27 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -48,6 +48,33 @@ def test_module_is_mean_absolute_error(self):
4848
value = metric_wrapper.compute()
4949
torch.testing.assert_close(value, torch.tensor(0.11))
5050

51+
def test_state_restore_and_reset(self):
52+
metric_wrapper = TrainMetricWrapper(
53+
torchmetrics.MeanAbsoluteError(), decay_rate=0.9, decay_step=1
54+
)
55+
preds = torch.tensor([0.1, 0.2])
56+
metric_wrapper.update(preds, torch.tensor([0.6, 0.7]))
57+
torch.testing.assert_close(metric_wrapper.compute(), torch.tensor(0.5))
58+
self.assertEqual(
59+
list(metric_wrapper.state_dict().keys()), ["_value", "_step_cnt"]
60+
)
61+
self.assertEqual(list(metric_wrapper.parameters()), [])
62+
63+
# continue train restores the running metric of the interrupted job.
64+
restored = TrainMetricWrapper(
65+
torchmetrics.MeanAbsoluteError(), decay_rate=0.9, decay_step=1
66+
)
67+
restored.load_state_dict(metric_wrapper.state_dict())
68+
torch.testing.assert_close(restored.compute(), torch.tensor(0.5))
69+
restored.update(preds, torch.tensor([0.2, 0.3]))
70+
torch.testing.assert_close(restored.compute(), torch.tensor(0.46))
71+
72+
# fine tune resets it, so the next value is not blended with the old one.
73+
restored.reset()
74+
torch.testing.assert_close(restored.compute(), torch.tensor(0.0))
75+
restored.update(preds, torch.tensor([0.2, 0.3]))
76+
torch.testing.assert_close(restored.compute(), torch.tensor(0.1))
77+
5178

5279
if __name__ == "__main__":
5380
unittest.main()

tzrec/models/model.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -156,6 +156,11 @@ def compute_train_metric(self) -> Dict[str, torch.Tensor]:
156156
metric_results[metric_name] = metric.compute()
157157
return metric_results
158158

159+
def reset_train_metric(self) -> None:
160+
"""Reset train metric state."""
161+
for metric in self._train_metric_modules.values():
162+
metric.reset()
163+
159164
def on_train_end(self) -> None:
160165
"""Hook fired once after the train_eval loop exits.
161166

tzrec/utils/export_util.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -639,7 +639,7 @@ def _rewrite_dense_serving_graph(
639639
dense_graph_config["sequence__ec"] = []
640640
for node in list(graph.nodes):
641641
if node.op == "call_function" and node.target == fx_mark_keyed_tensor:
642-
name = node.args[0]
642+
name = cast(str, node.args[0])
643643
if node.kwargs.get("is_dense", False):
644644
continue
645645
node_kt = node.args[1]

tzrec/version.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,4 +9,4 @@
99
# See the License for the specific language governing permissions and
1010
# limitations under the License.
1111

12-
__version__ = "1.4.0"
12+
__version__ = "1.4.1"

0 commit comments

Comments
 (0)