Wei Liu Claude Sonnet 4.6 commited on
Commit
2aa6ae4
·
1 Parent(s): 6dad5cc

Patch Genesis from_torch to fix PyTorch 2.5 tensor subclass error

Browse files

In PyTorch 2.5+, torch.Tensor.__new__(SubClass, existing_tensor) raises:
"raw Tensor object is already associated to a python object of type
Tensor which is not a subclass of the requested type"
when the existing TensorImpl's Python wrapper is a plain torch.Tensor
(not a subclass of genesis.grad.Tensor).

Genesis's from_torch() does exactly this: torch.zeros() → plain tensor,
then Tensor(plain_tensor) → fails because plain torch.Tensor is not a
subclass of genesis.grad.Tensor.

Fix: use torch.Tensor._make_subclass(Tensor, t) which is PyTorch's
official API for creating tensor subclass views, bypassing the
association check. Patch is applied before startup() runs, so all
genesis gs.zeros()/gs.ones() etc. calls use the fixed path.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>

Files changed (1) hide show
  1. app.py +49 -0
app.py CHANGED
@@ -27,6 +27,54 @@ try:
27
  except Exception:
28
  pass
29
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
30
  import base64
31
  import io
32
  import threading
@@ -911,6 +959,7 @@ def build_demo():
911
  # files are already on disk so snapshot_download() is a fast no-op. By doing
912
  # this here we avoid holding a ZeroGPU allocation while waiting on downloads.
913
  _ensure_models_downloaded()
 
914
  startup() # Load all models and scenes to CPU at module level
915
  demo = build_demo()
916
 
 
27
  except Exception:
28
  pass
29
 
30
+ # Patch Genesis from_torch: in PyTorch 2.5+, Tensor(existing_plain_tensor) raises
31
+ # "raw Tensor object is already associated to a python object of type Tensor
32
+ # which is not a subclass of the requested type"
33
+ # because torch.Tensor.__new__(SubClass, existing_tensor) checks that the existing
34
+ # TensorImpl's Python wrapper is a subclass of SubClass. torch.Tensor is the parent,
35
+ # not a subclass of genesis.grad.Tensor, so the check fails.
36
+ # Fix: use torch.Tensor._make_subclass(cls, t) which is the proper PyTorch API for
37
+ # creating a subclass view of an existing tensor regardless of the wrapper type.
38
+ def _patch_genesis_from_torch():
39
+ try:
40
+ import genesis
41
+ import genesis.grad.creation_ops as _gc_ops
42
+ import genesis.grad.tensor as _gt_mod
43
+ _Tensor = _gt_mod.Tensor
44
+ _gs = genesis
45
+
46
+ def _patched_from_torch(torch_tensor, dtype=None, requires_grad=False, detach=True, scene=None):
47
+ if dtype is None:
48
+ dtype = torch_tensor.dtype
49
+ if dtype in (float, torch.float32, torch.float64):
50
+ dtype = _gs.tc_float
51
+ elif dtype in (int, torch.int32, torch.int64):
52
+ dtype = _gs.tc_int
53
+ elif dtype in (bool, torch.bool):
54
+ dtype = torch.bool
55
+ else:
56
+ _gs.raise_exception(f"Unsupported dtype: {dtype}")
57
+ if torch_tensor.requires_grad and (not detach) and (not requires_grad):
58
+ requires_grad = True
59
+ t = torch_tensor.to(device=_gs.device, dtype=dtype)
60
+ # _make_subclass creates a SubClass view without triggering the PyTorch
61
+ # "already associated" check that Tensor(existing_tensor) hits.
62
+ gs_tensor = torch.Tensor._make_subclass(_Tensor, t)
63
+ gs_tensor.scene = scene
64
+ gs_tensor.uid = _gs.UID()
65
+ gs_tensor.parents = []
66
+ gs_tensor = gs_tensor.clone()
67
+ if detach:
68
+ gs_tensor = gs_tensor.detach(sceneless=False)
69
+ if requires_grad:
70
+ gs_tensor = gs_tensor.requires_grad_()
71
+ return gs_tensor
72
+
73
+ _gc_ops.from_torch = _patched_from_torch
74
+ print("[patch] Genesis from_torch patched (_make_subclass fix for PyTorch 2.5+)")
75
+ except Exception as e:
76
+ print(f"[patch] Genesis from_torch patch skipped: {e}")
77
+
78
  import base64
79
  import io
80
  import threading
 
959
  # files are already on disk so snapshot_download() is a fast no-op. By doing
960
  # this here we avoid holding a ZeroGPU allocation while waiting on downloads.
961
  _ensure_models_downloaded()
962
+ _patch_genesis_from_torch() # Fix Genesis from_torch for PyTorch 2.5 compatibility
963
  startup() # Load all models and scenes to CPU at module level
964
  demo = build_demo()
965