Get all [test] classes in a model test file with attribute `all_model_classes` that are non-empty. These are usually the (model) test classes containing the (non-slow) tests to run and are subclasses of one of the classes `ModelTesterMixin`, `TFModelTesterMixin` or `FlaxModelTesterMixin`, a
(test_file)
| 71 | |
| 72 | |
| 73 | def get_test_classes(test_file): |
| 74 | """Get all [test] classes in a model test file with attribute `all_model_classes` that are non-empty. |
| 75 | |
| 76 | These are usually the (model) test classes containing the (non-slow) tests to run and are subclasses of one of the |
| 77 | classes `ModelTesterMixin`, `TFModelTesterMixin` or `FlaxModelTesterMixin`, as well as a subclass of |
| 78 | `unittest.TestCase`. Exceptions include `RagTestMixin` (and its subclasses). |
| 79 | """ |
| 80 | test_classes = [] |
| 81 | test_module = get_test_module(test_file) |
| 82 | for attr in dir(test_module): |
| 83 | attr_value = getattr(test_module, attr) |
| 84 | # (TF/Flax)ModelTesterMixin is also an attribute in specific model test module. Let's exclude them by checking |
| 85 | # `all_model_classes` is not empty (which also excludes other special classes). |
| 86 | model_classes = getattr(attr_value, "all_model_classes", []) |
| 87 | if len(model_classes) > 0: |
| 88 | test_classes.append(attr_value) |
| 89 | |
| 90 | # sort with class names |
| 91 | return sorted(test_classes, key=lambda x: x.__name__) |
| 92 | |
| 93 | |
| 94 | def get_model_classes(test_file): |
no test coverage detected