55from __future__ import annotations
66
77import weakref
8- from collections .abc import Sequence
8+ from collections .abc import Callable , Sequence
99from numbers import Number
1010
1111import numpy as np
@@ -76,7 +76,7 @@ class IndependentNormal(D.Independent):
7676 def __init__ (
7777 self ,
7878 loc : torch .Tensor ,
79- scale : torch .Tensor ,
79+ scale : torch .Tensor | float | Callable [[ torch . Tensor ], torch . Tensor ] ,
8080 upscale : float = 5.0 ,
8181 tanh_loc : bool = False ,
8282 event_dim : int = 1 ,
@@ -86,11 +86,25 @@ def __init__(
8686 self .upscale = upscale
8787 self ._event_dim = event_dim
8888 self ._kwargs = kwargs
89+ # Support callable scale (e.g., torch.ones_like) for compile-friendliness
90+ if callable (scale ) and not isinstance (scale , torch .Tensor ):
91+ scale = scale (loc )
92+ elif not isinstance (scale , torch .Tensor ):
93+ scale = torch .as_tensor (scale , device = loc .device , dtype = loc .dtype )
94+ elif scale .device != loc .device :
95+ scale = scale .to (loc .device , non_blocking = loc .device .type == "cuda" )
8996 super ().__init__ (D .Normal (loc , scale , ** kwargs ), event_dim )
9097
9198 def update (self , loc , scale ):
9299 if self .tanh_loc :
93100 loc = self .upscale * (loc / self .upscale ).tanh ()
101+ # Support callable scale (e.g., torch.ones_like) for compile-friendliness
102+ if callable (scale ) and not isinstance (scale , torch .Tensor ):
103+ scale = scale (loc )
104+ elif not isinstance (scale , torch .Tensor ):
105+ scale = torch .as_tensor (scale , device = loc .device , dtype = loc .dtype )
106+ elif scale .device != loc .device :
107+ scale = scale .to (loc .device , non_blocking = loc .device .type == "cuda" )
94108 super ().__init__ (D .Normal (loc , scale , ** self ._kwargs ), self ._event_dim )
95109
96110 @property
@@ -343,7 +357,7 @@ class TanhNormal(FasterTransformedDistribution):
343357 def __init__ (
344358 self ,
345359 loc : torch .Tensor ,
346- scale : torch .Tensor ,
360+ scale : torch .Tensor | float | Callable [[ torch . Tensor ], torch . Tensor ] ,
347361 upscale : torch .Tensor | Number = 5.0 ,
348362 low : torch .Tensor | Number = - 1.0 ,
349363 high : torch .Tensor | Number = 1.0 ,
@@ -353,8 +367,14 @@ def __init__(
353367 ):
354368 if not isinstance (loc , torch .Tensor ):
355369 loc = torch .as_tensor (loc , dtype = torch .get_default_dtype ())
356- if not isinstance (scale , torch .Tensor ):
357- scale = torch .as_tensor (scale , dtype = torch .get_default_dtype ())
370+ _non_blocking = loc .device .type == "cuda"
371+ # Support callable scale (e.g., torch.ones_like) for compile-friendliness
372+ if callable (scale ) and not isinstance (scale , torch .Tensor ):
373+ scale = scale (loc )
374+ elif not isinstance (scale , torch .Tensor ):
375+ scale = torch .as_tensor (scale , device = loc .device , dtype = loc .dtype )
376+ elif scale .device != loc .device :
377+ scale = scale .to (loc .device , non_blocking = _non_blocking )
358378 if event_dims is None :
359379 event_dims = min (1 , loc .ndim )
360380
@@ -373,11 +393,11 @@ def __init__(
373393 if not isinstance (high , torch .Tensor ):
374394 high = torch .as_tensor (high , device = loc .device )
375395 elif high .device != loc .device :
376- high = high .to (loc .device )
396+ high = high .to (loc .device , non_blocking = _non_blocking )
377397 if not isinstance (low , torch .Tensor ):
378398 low = torch .as_tensor (low , device = loc .device )
379399 elif low .device != loc .device :
380- low = low .to (loc .device )
400+ low = low .to (loc .device , non_blocking = _non_blocking )
381401 if not is_compiling () and not safe_is_current_stream_capturing ():
382402 self .non_trivial_max = (high != 1.0 ).any ()
383403 self .non_trivial_min = (low != - 1.0 ).any ()
@@ -391,10 +411,10 @@ def __init__(
391411 self .upscale = (
392412 upscale
393413 if not isinstance (upscale , torch .Tensor )
394- else upscale .to (self .device )
414+ else upscale .to (self .device , non_blocking = _non_blocking )
395415 )
396416
397- low = low .to (loc .device )
417+ low = low .to (loc .device , non_blocking = _non_blocking )
398418 self .low = low
399419 self .high = high
400420
@@ -434,6 +454,13 @@ def update(self, loc: torch.Tensor, scale: torch.Tensor) -> None:
434454 # loc must be rescaled if tanh_loc
435455 if is_compiling () or (self .non_trivial_max or self .non_trivial_min ):
436456 loc = loc + (self .high - self .low ) / 2 + self .low
457+ # Support callable scale (e.g., torch.ones_like) for compile-friendliness
458+ if callable (scale ) and not isinstance (scale , torch .Tensor ):
459+ scale = scale (loc )
460+ elif not isinstance (scale , torch .Tensor ):
461+ scale = torch .as_tensor (scale , device = loc .device , dtype = loc .dtype )
462+ elif scale .device != loc .device :
463+ scale = scale .to (loc .device , non_blocking = loc .device .type == "cuda" )
437464 self .loc = loc
438465 self .scale = scale
439466
0 commit comments