()
| 191 | |
| 192 | |
| 193 | def torch_version_float(): |
| 194 | version_str = torch.__version__ |
| 195 | version_re = re.search(r'^([0-9]+\.[0-9]+)', version_str) |
| 196 | if version_re: |
| 197 | version = float(version_re.group(1)) |
| 198 | logx.msg(f'Torch version: {version}, {version_str}') |
| 199 | else: |
| 200 | version = 1.0 |
| 201 | logx.msg(f'Can\'t parse torch version ({version}), assuming {version}') |
| 202 | return version |
| 203 | |
| 204 | |
| 205 | def assert_and_infer_cfg(args, make_immutable=True, train_mode=True): |
no outgoing calls
no test coverage detected