| 396 | converted_count = [0] # 使用列表以便在嵌套函数中修改 |
| 397 | |
| 398 | def _convert_recursive(obj): |
| 399 | # 如果是torch tensor且在CUDA上 |
| 400 | if torch.is_tensor(obj): |
| 401 | if cuda_device: |
| 402 | return obj.to(cuda_device) if obj.is_cpu else obj |
| 403 | if obj.is_cuda: |
| 404 | return obj.cpu() |
| 405 | return obj |
| 406 | |
| 407 | # 处理各种容器类型 |
| 408 | elif isinstance(obj, dict): |
| 409 | return {k: _convert_recursive(v) for k, v in obj.items()} |
| 410 | |
| 411 | elif isinstance(obj, list): |
| 412 | return [_convert_recursive(item) for item in obj] |
| 413 | |
| 414 | elif isinstance(obj, tuple): |
| 415 | # 元组不可变,总是创建新的 |
| 416 | return tuple(_convert_recursive(item) for item in obj) |
| 417 | |
| 418 | elif isinstance(obj, set): |
| 419 | return {_convert_recursive(item) for item in obj} |
| 420 | |
| 421 | # 其他数据类型直接返回 |
| 422 | else: |
| 423 | return obj |
| 424 | |
| 425 | return _convert_recursive(data) |
| 426 | |