MCPcopy Create free account
hub / github.com/openai/point-e / __init__

Method __init__

point_e/diffusion/sampler.py:26–84  ·  view source on GitHub ↗
(
        self,
        device: torch.device,
        models: Sequence[nn.Module],
        diffusions: Sequence[GaussianDiffusion],
        num_points: Sequence[int],
        aux_channels: Sequence[str],
        model_kwargs_key_filter: Sequence[str] = ("*",),
        guidance_scale: Sequence[float] = (3.0, 3.0),
        clip_denoised: bool = True,
        use_karras: Sequence[bool] = (True, True),
        karras_steps: Sequence[int] = (64, 64),
        sigma_min: Sequence[float] = (1e-3, 1e-3),
        sigma_max: Sequence[float] = (120, 160),
        s_churn: Sequence[float] = (3, 0),
    )

Source from the content-addressed store, hash-verified

24 """
25
26 def __init__(
27 self,
28 device: torch.device,
29 models: Sequence[nn.Module],
30 diffusions: Sequence[GaussianDiffusion],
31 num_points: Sequence[int],
32 aux_channels: Sequence[str],
33 model_kwargs_key_filter: Sequence[str] = ("*",),
34 guidance_scale: Sequence[float] = (3.0, 3.0),
35 clip_denoised: bool = True,
36 use_karras: Sequence[bool] = (True, True),
37 karras_steps: Sequence[int] = (64, 64),
38 sigma_min: Sequence[float] = (1e-3, 1e-3),
39 sigma_max: Sequence[float] = (120, 160),
40 s_churn: Sequence[float] = (3, 0),
41 ):
42 n = len(models)
43 assert n > 0
44
45 if n > 1:
46 if len(guidance_scale) == 1:
47 # Don't guide the upsamplers by default.
48 guidance_scale = list(guidance_scale) + [1.0] * (n - 1)
49 if len(use_karras) == 1:
50 use_karras = use_karras * n
51 if len(karras_steps) == 1:
52 karras_steps = karras_steps * n
53 if len(sigma_min) == 1:
54 sigma_min = sigma_min * n
55 if len(sigma_max) == 1:
56 sigma_max = sigma_max * n
57 if len(s_churn) == 1:
58 s_churn = s_churn * n
59 if len(model_kwargs_key_filter) == 1:
60 model_kwargs_key_filter = model_kwargs_key_filter * n
61 if len(model_kwargs_key_filter) == 0:
62 model_kwargs_key_filter = ["*"] * n
63 assert len(guidance_scale) == n
64 assert len(use_karras) == n
65 assert len(karras_steps) == n
66 assert len(sigma_min) == n
67 assert len(sigma_max) == n
68 assert len(s_churn) == n
69 assert len(model_kwargs_key_filter) == n
70
71 self.device = device
72 self.num_points = num_points
73 self.aux_channels = aux_channels
74 self.model_kwargs_key_filter = model_kwargs_key_filter
75 self.guidance_scale = guidance_scale
76 self.clip_denoised = clip_denoised
77 self.use_karras = use_karras
78 self.karras_steps = karras_steps
79 self.sigma_min = sigma_min
80 self.sigma_max = sigma_max
81 self.s_churn = s_churn
82
83 self.models = models

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected