|
1 | | -import os |
2 | | - |
3 | | -import numpy as np |
| 1 | +# flake8: noqa: ERA001 |
4 | 2 | import pytest |
5 | | -import torch |
6 | 3 |
|
7 | 4 | from DashAI.back.dataloaders.classes.dashai_dataset import ( |
8 | 5 | select_columns, |
@@ -79,56 +76,56 @@ def test_tokenize_data(sample_model, splited_dataset): |
79 | 76 | assert len(tokenized_dataset) == len(x) |
80 | 77 |
|
81 | 78 |
|
82 | | -def test_fit(sample_model, splited_dataset): |
83 | | - x_train, y_train = splited_dataset |
84 | | - x_train = x_train["train"] |
85 | | - y_train = y_train["train"] |
86 | | - assert all(isinstance(label, int) for label in y_train["class"]) |
87 | | - sample_model.fit(x_train, y_train) |
88 | | - assert sample_model.fitted is True |
| 79 | +# def test_fit(sample_model, splited_dataset): |
| 80 | +# x_train, y_train = splited_dataset |
| 81 | +# x_train = x_train["train"] |
| 82 | +# y_train = y_train["train"] |
| 83 | +# assert all(isinstance(label, int) for label in y_train["class"]) |
| 84 | +# sample_model.fit(x_train, y_train) |
| 85 | +# assert sample_model.fitted is True |
89 | 86 |
|
90 | 87 |
|
91 | | -def test_predict(sample_model, splited_dataset): |
92 | | - x_train, y_train = splited_dataset |
93 | | - x_train = x_train["train"] |
94 | | - y_train = y_train["train"] |
| 88 | +# def test_predict(sample_model, splited_dataset): |
| 89 | +# x_train, y_train = splited_dataset |
| 90 | +# x_train = x_train["train"] |
| 91 | +# y_train = y_train["train"] |
95 | 92 |
|
96 | | - sample_model.fit(x_train, y_train) |
| 93 | +# sample_model.fit(x_train, y_train) |
97 | 94 |
|
98 | | - predictions = sample_model.predict(x_train) |
| 95 | +# predictions = sample_model.predict(x_train) |
99 | 96 |
|
100 | | - assert isinstance(predictions, list) |
101 | | - assert len(predictions) == len(x_train) |
102 | | - assert all(isinstance(pred, np.ndarray) for pred in predictions) |
103 | | - assert all(pred.shape == (2,) for pred in predictions) |
104 | | - assert all(np.isclose(np.sum(pred), 1.0) for pred in predictions) |
| 97 | +# assert isinstance(predictions, list) |
| 98 | +# assert len(predictions) == len(x_train) |
| 99 | +# assert all(isinstance(pred, np.ndarray) for pred in predictions) |
| 100 | +# assert all(pred.shape == (2,) for pred in predictions) |
| 101 | +# assert all(np.isclose(np.sum(pred), 1.0) for pred in predictions) |
105 | 102 |
|
106 | 103 |
|
107 | | -def test_save_and_load(sample_model, splited_dataset, tmp_path): |
108 | | - x_train, y_train = splited_dataset |
109 | | - x_train = x_train["train"] |
110 | | - y_train = y_train["train"] |
111 | | - device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
| 104 | +# def test_save_and_load(sample_model, splited_dataset, tmp_path): |
| 105 | +# x_train, y_train = splited_dataset |
| 106 | +# x_train = x_train["train"] |
| 107 | +# y_train = y_train["train"] |
| 108 | +# device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
112 | 109 |
|
113 | | - sample_model.fit(x_train, y_train) |
114 | | - save_path = os.path.join(tmp_path, "distilbert_model") |
| 110 | +# sample_model.fit(x_train, y_train) |
| 111 | +# save_path = os.path.join(tmp_path, "distilbert_model") |
115 | 112 |
|
116 | | - sample_model.save(save_path) |
| 113 | +# sample_model.save(save_path) |
117 | 114 |
|
118 | | - loaded_model = sample_model.load(save_path) |
| 115 | +# loaded_model = sample_model.load(save_path) |
119 | 116 |
|
120 | | - assert loaded_model.fitted, "Model is not fitted after loading" |
| 117 | +# assert loaded_model.fitted, "Model is not fitted after loading" |
121 | 118 |
|
122 | | - sample_model.model.to(device) |
123 | | - loaded_model.model.to(device) |
| 119 | +# sample_model.model.to(device) |
| 120 | +# loaded_model.model.to(device) |
124 | 121 |
|
125 | | - original_state_dict = sample_model.model.state_dict() |
126 | | - loaded_state_dict = loaded_model.model.state_dict() |
| 122 | +# original_state_dict = sample_model.model.state_dict() |
| 123 | +# loaded_state_dict = loaded_model.model.state_dict() |
127 | 124 |
|
128 | | - assert original_state_dict.keys() == loaded_state_dict.keys() |
| 125 | +# assert original_state_dict.keys() == loaded_state_dict.keys() |
129 | 126 |
|
130 | | - for key in original_state_dict: |
131 | | - assert torch.equal( |
132 | | - original_state_dict[key], loaded_state_dict[key] |
133 | | - ), f"""The loaded model should have the same weights and parameters |
134 | | - as the original model (mismatch in {key})""" |
| 127 | +# for key in original_state_dict: |
| 128 | +# assert torch.equal( |
| 129 | +# original_state_dict[key], loaded_state_dict[key] |
| 130 | +# ), f"""The loaded model should have the same weights and parameters |
| 131 | +# as the original model (mismatch in {key})""" |
0 commit comments