Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 6 additions & 6 deletions jax3d/projects/generative/nerf/autoencoder/data.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,9 +26,9 @@
@dataclasses.dataclass(frozen=True)
class TFDSImageDatasetReader():
"""Dataset reader wrapping TFDS image datasets."""
dataset_name: str = gin.REQUIRED
resolution: int = gin.REQUIRED
batch_size: int = gin.REQUIRED
dataset_name: str = gin.REQUIRED # pyrefly: ignore[bad-assignment]
resolution: int = gin.REQUIRED # pyrefly: ignore[bad-assignment]
batch_size: int = gin.REQUIRED # pyrefly: ignore[bad-assignment]

split: str = "train"
eval_fraction: float = 0.05
Expand Down Expand Up @@ -82,9 +82,9 @@ def extract_image(data):
@dataclasses.dataclass(frozen=True)
class TiledMNISTDatasetReader():
"""Dataset reader wrapping TFDS image datasets."""
resolution: int = gin.REQUIRED
batch_size: int = gin.REQUIRED
tile_factor: int = gin.REQUIRED
resolution: int = gin.REQUIRED # pyrefly: ignore[bad-assignment]
batch_size: int = gin.REQUIRED # pyrefly: ignore[bad-assignment]
tile_factor: int = gin.REQUIRED # pyrefly: ignore[bad-assignment]

split: str = "train"
eval_fraction: float = 0.05
Expand Down
6 changes: 3 additions & 3 deletions jax3d/projects/generative/nerf/autoencoder/loss.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,10 +30,10 @@
def transformer_loss_fn(
model_parameters: models.ModelParameters,
data: Dict[str, Any],
rng: PRNGKey,
rng: PRNGKey, # pyrefly: ignore[not-a-type]
step: jnp.ndarray,
pixel_batch_size=512,
) -> Tuple[FloatArray, Dict[str, FloatArray]]:
) -> Tuple[FloatArray, Dict[str, FloatArray]]: # pyrefly: ignore[not-a-type]
"""The main autoencoder loss function.

Args:
Expand Down Expand Up @@ -92,7 +92,7 @@ def take(pixels, inds):
total_loss += reconstruction_loss * reconstruction_loss_weight

# Always compute PSNR on gamma values for consistency
psnr = metrics.psnr(predicted_rgb, gt_rgb)
psnr = metrics.psnr(predicted_rgb, gt_rgb) # pyrefly: ignore[bad-argument-type]
loss_terms["Training PSNR"] = psnr

return total_loss, loss_terms
2 changes: 1 addition & 1 deletion jax3d/projects/generative/nerf/autoencoder/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -246,7 +246,7 @@ def __call__(self, image, pixels=None, step=None, is_training=False):

return rgb

def initialize_parameters(self, rng_key: PRNGKey,
def initialize_parameters(self, rng_key: PRNGKey, # pyrefly: ignore[not-a-type]
image_size) -> ModelParameters:
batch_size = 7
pixel_batch_size = 11
Expand Down
14 changes: 7 additions & 7 deletions jax3d/projects/generative/nerf/autoencoder/trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,19 +39,19 @@ class TransformerTrainState(trainer.TrainState):
"""Training state for the NeRF Transformer Model."""
model_parameters: models.ModelParameters
optimizer_state: optax.OptState
rng: PRNGKey
rng: PRNGKey # pyrefly: ignore[not-a-type]

def to_serializable(self) -> "TransformerTrainState":
"""Transforms the state values into a form suitable for serialization."""
return self.replace(
return self.replace( # pyrefly: ignore[missing-attribute]
model_parameters=jax_utils.unreplicate(self.model_parameters),
optimizer_state=jax_utils.unreplicate(self.optimizer_state),
)

def from_serializable(self) -> "TransformerTrainState":
"""Transforms deserialized values into a form suitable for training."""
state = self
state = state.replace(
state = state.replace( # pyrefly: ignore[missing-attribute]
model_parameters=jax_utils.replicate(state.model_parameters),
optimizer_state=jax_utils.replicate(state.optimizer_state),
)
Expand Down Expand Up @@ -89,7 +89,7 @@ def data_reader(self) -> "dataset_reader_class":

def init_state(self) -> TransformerTrainState:
"""Initializes training state."""
rng = jax.random.PRNGKey(self.random_seed)
rng = jax.random.PRNGKey(self.random_seed) # pyrefly: ignore[bad-argument-type]
train_rng, model_init_rng = jax.random.split(rng)
res = self.data_reader.resolution
model_parameters = self.model.initialize_parameters(model_init_rng,
Expand Down Expand Up @@ -135,7 +135,7 @@ def _loss_fn(model_parameters, data, rng, step):

return jax.pmap(per_device_train_step, "replicas")

def train_step(
def train_step( # pyrefly: ignore[bad-override]
self, train_state: TransformerTrainState, inputs: Dict[str, Any],
scratch: Optional[Dict[str, Any]]
) -> Tuple[TransformerTrainState, Dict[str, Any], Dict[str, Any]]:
Expand Down Expand Up @@ -173,15 +173,15 @@ def train_step(
loss_terms = jax.tree.map(np.array, loss_terms)
loss_terms["param_count"] = scratch["param_count"]

new_train_state = train_state.replace(
new_train_state = train_state.replace( # pyrefly: ignore[missing-attribute]
model_parameters=new_model_parameters,
optimizer_state=new_optimizer_state,
rng=next_rng,
step=train_state.step + 1)

return new_train_state, loss_terms, scratch

def eval_step(self, train_state: TransformerTrainState,
def eval_step(self, train_state: TransformerTrainState, # pyrefly: ignore[bad-override]
summary_writer: tensorboard.SummaryWriter,
scratch: Optional[Any]) -> Any:
"""Performs evaluation on a model checkpoint.
Expand Down
6 changes: 3 additions & 3 deletions jax3d/projects/generative/nerf/camera.py
Original file line number Diff line number Diff line change
Expand Up @@ -150,8 +150,8 @@ def make_camera( # pytype: disable=annotation-type-mismatch # jax-ndarray
focal_length: Union[jnp.ndarray, float],
principal_point: jnp.ndarray,
image_size: jnp.ndarray,
skew: Union[jnp.ndarray, float] = None,
pixel_aspect_ratio: Union[jnp.ndarray, float] = None,
skew: Union[jnp.ndarray, float] = None, # pyrefly: ignore[bad-function-definition]
pixel_aspect_ratio: Union[jnp.ndarray, float] = None, # pyrefly: ignore[bad-function-definition]
radial_distortion: Optional[jnp.ndarray] = None,
tangential_distortion: Optional[jnp.ndarray] = None) -> CameraType:
"""Create a dictionary containing standard values for representing a camera.
Expand Down Expand Up @@ -222,7 +222,7 @@ def look_at(eye_position: jnp.ndarray,
# +Z = forward
camera_forward = target - eye_position
camera_forward /= jnp.linalg.norm(camera_forward)
camera_right = jnp.cross(camera_forward, global_up)
camera_right = jnp.cross(camera_forward, global_up) # pyrefly: ignore[bad-argument-type]
camera_right /= jnp.linalg.norm(camera_right)
camera_up = jnp.cross(camera_right, camera_forward)
orientation = jnp.stack([-camera_right, camera_up, camera_forward], axis=0)
Expand Down
2 changes: 1 addition & 1 deletion jax3d/projects/generative/nerf/glo_nerf/eval.py
Original file line number Diff line number Diff line change
Expand Up @@ -179,7 +179,7 @@ def flatten_views(t):
render = models.Model().apply(
model_parameters, inputs, rays, rng=rng, step=step)

pred = render["gamma_rgb"]
pred = render["gamma_rgb"] # pyrefly: ignore[bad-index]
gt = data_flat["gamma_rgb"]
if apply_mask:
pred *= data_flat["weight"]
Expand Down
16 changes: 8 additions & 8 deletions jax3d/projects/generative/nerf/glo_nerf/loss.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,11 +32,11 @@ def transformer_nerf_loss_fn(
model_parameters: models.ModelParameters,
inputs: models.ModelInputs,
data: Dict[str, Any],
rng: PRNGKey,
rng: PRNGKey, # pyrefly: ignore[not-a-type]
step: jnp.ndarray,
mask_mode: str = "alpha_supervision",
reconstruct_gamma_rgb: bool = True
) -> Tuple[FloatArray, Dict[str, FloatArray]]:
) -> Tuple[FloatArray, Dict[str, FloatArray]]: # pyrefly: ignore[not-a-type]
"""The main GLO NeRF loss function including all terms.

Args:
Expand Down Expand Up @@ -65,7 +65,7 @@ def flatten_views(t):
data = jax.tree.map(flatten_views, data)
latent_tokens = inputs.latent_tokens
latent_tokens = jax.tree.map(flatten_views, latent_tokens)
inputs = inputs.replace(latent_tokens=latent_tokens)
inputs = inputs.replace(latent_tokens=latent_tokens) # pyrefly: ignore[missing-attribute]

origins, directions = jax.vmap(jax_camera.pixels_to_rays)(
data["camera"], data["pixel_coordinates"])
Expand Down Expand Up @@ -96,10 +96,10 @@ def flatten_views(t):
# gamma decoding their images this is effectively the loss they are using.
# (https://en.wikipedia.org/wiki/Gamma_correction).
gt_rgb = data["gamma_rgb"]
predicted_rgb = render["gamma_rgb"]
predicted_rgb = render["gamma_rgb"] # pyrefly: ignore[bad-index]
else:
gt_rgb = image_utility.srgb_gamma_to_linear(data["gamma_rgb"])
predicted_rgb = render["linear_rgb"]
predicted_rgb = render["linear_rgb"] # pyrefly: ignore[bad-index]

if mask_mode == "multiply":
gt_rgb *= foreground_mask
Expand All @@ -115,19 +115,19 @@ def flatten_views(t):
gt_rgb = data["gamma_rgb"]
if mask_mode == "multiply":
gt_rgb *= foreground_mask
psnr = metrics.psnr(render["gamma_rgb"], gt_rgb)
psnr = metrics.psnr(render["gamma_rgb"], gt_rgb) # pyrefly: ignore[bad-index]
loss_terms["Training PSNR"] = psnr

if mask_mode == "alpha_supervision":
with gin.config_scope("alpha"):
alpha_loss, alpha_loss_weight = losses.reconstruction(
foreground_mask, render["alpha"])
foreground_mask, render["alpha"]) # pyrefly: ignore[bad-index]
if alpha_loss_weight != 0.0:
loss_terms["Alpha"] = alpha_loss
total_loss += alpha_loss_weight * alpha_loss

hard_surface_loss, hard_surface_loss_weight = losses.hard_surface(
render["sample_weights"])
render["sample_weights"]) # pyrefly: ignore[bad-index]
loss_terms["Hard Surface"] = hard_surface_loss
total_loss += hard_surface_loss_weight * hard_surface_loss

Expand Down
28 changes: 14 additions & 14 deletions jax3d/projects/generative/nerf/glo_nerf/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -313,20 +313,20 @@ def setup(self):

def evaluate_nerf(
self,
points: FloatArray["N", ..., 3],
directions: FloatArray["N", ..., 3],
points: FloatArray["N", ..., 3], # pyrefly: ignore[not-a-type, unknown-name]
directions: FloatArray["N", ..., 3], # pyrefly: ignore[not-a-type, unknown-name]
inputs: ModelInputs,
step: int,
noise_rng: Optional[PRNGKey] = None) -> Dict[str, FloatArray]:
noise_rng: Optional[PRNGKey] = None) -> Dict[str, FloatArray]: # pyrefly: ignore[not-a-type]
"""NeRF decoder forward pass."""
density, rgb = self.decoder(points, inputs.latent_tokens)

results = {"sample_density": density, "sample_values": {"rgb": rgb,}}
return results

def surface_normal_from_density(
self, points: FloatArray["N", ..., 3], inputs, step: int
) -> FloatArray["N", ..., 3]:
self, points: FloatArray["N", ..., 3], inputs, step: int # pyrefly: ignore[not-a-type, unknown-name]
) -> FloatArray["N", ..., 3]: # pyrefly: ignore[not-a-type, unknown-name]
"""Compute the normals of the density field at the given points."""

# Input points are of the form [ N, ..., 3]. Compute the product of the
Expand All @@ -349,7 +349,7 @@ def density_fn(fn_points):

def __call__(self,
inputs: ModelInputs,
rays: FloatArray["K R", 6],
rays: FloatArray["K R", 6], # pyrefly: ignore[not-a-type]
rng=None,
near=None,
far=None,
Expand Down Expand Up @@ -404,7 +404,7 @@ def __call__(self,

initial_nerf_result = self.evaluate_nerf(initial_sample_coordinates,
initial_sample_directions, inputs,
step, initial_noise_rng)
step, initial_noise_rng) # pyrefly: ignore[bad-argument-type]

initial_render_results = volume_rendering.volume_rendering(
sample_values=(), # We ignore the accumulated RGB from initial samples
Expand Down Expand Up @@ -448,7 +448,7 @@ def __call__(self,

combined_nerf_result = self.evaluate_nerf(combined_sample_coordinates,
combined_sample_directions,
inputs, step,
inputs, step, # pyrefly: ignore[bad-argument-type]
importance_noise_rng)

else:
Expand All @@ -473,7 +473,7 @@ def __call__(self,

importance_nerf_result = self.evaluate_nerf(importance_sample_coordinates,
importance_sample_directions,
inputs, step,
inputs, step, # pyrefly: ignore[bad-argument-type]
importance_noise_rng)

# We first append importance sample depths and values along the sample
Expand Down Expand Up @@ -556,7 +556,7 @@ def __call__(self,
surface_points = origins + directions * expected_depth
return_values["depth"] = expected_depth

normals = self.surface_normal_from_density(surface_points, inputs, step)
normals = self.surface_normal_from_density(surface_points, inputs, step) # pyrefly: ignore[bad-argument-type]
return_values["analytic_normal"] = normals

if return_additional_sample_data:
Expand All @@ -565,7 +565,7 @@ def __call__(self,
alpha = render_results.ray_alpha[..., None]
return_values["alpha"] = alpha

foreground_rgb = render_results.ray_values["rgb"]
foreground_rgb = render_results.ray_values["rgb"] # pyrefly: ignore[bad-index]

if self.use_background_model:
background_latent = inputs.latent_tokens[..., 0, :]
Expand All @@ -574,7 +574,7 @@ def __call__(self,
# need to add it to the attenuated background value.
pixel_rgb = foreground_rgb + (1.0 - alpha) * background_rgb
else:
background_rgb = jnp.zeros_like(foreground_rgb)
background_rgb = jnp.zeros_like(foreground_rgb) # pyrefly: ignore[bad-argument-type]
pixel_rgb = foreground_rgb

# Up until this point RGB values could be interpreted a gamma encoded or
Expand Down Expand Up @@ -726,13 +726,13 @@ def unpad(tensor):

for key in results:
result_i = jnp.concatenate(results[key])
results[key] = result_i.reshape(height, width, *result_i.shape[1:])
results[key] = result_i.reshape(height, width, *result_i.shape[1:]) # pyrefly: ignore[unsupported-operation]

return results

return render_image

def initialize_parameters(self, rng_key: PRNGKey, num_tokens,
def initialize_parameters(self, rng_key: PRNGKey, num_tokens, # pyrefly: ignore[not-a-type]
token_dim) -> ModelParameters:
batch_size = 7
num_rays = 13 # Prime numbers to help catch shape errors.
Expand Down
14 changes: 7 additions & 7 deletions jax3d/projects/generative/nerf/glo_nerf/trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,19 +40,19 @@ class TransformerNeRFTrainState(trainer.TrainState):
model_parameters: models.ModelParameters
latent_table: np.ndarray
optimizer_state: optax.OptState
rng: PRNGKey
rng: PRNGKey # pyrefly: ignore[not-a-type]

def to_serializable(self) -> "TransformerNeRFTrainState":
"""Transforms the state values into a form suitable for serialization."""
return self.replace(
return self.replace( # pyrefly: ignore[missing-attribute]
model_parameters=jax_utils.unreplicate(self.model_parameters),
optimizer_state=jax_utils.unreplicate(self.optimizer_state),
)

def from_serializable(self) -> "TransformerNeRFTrainState":
"""Transforms deserialized values into a form suitable for training."""
state = self
state = state.replace(
state = state.replace( # pyrefly: ignore[missing-attribute]
latent_table=np.copy(state.latent_table),
model_parameters=jax_utils.replicate(state.model_parameters),
optimizer_state=jax_utils.replicate(state.optimizer_state),
Expand Down Expand Up @@ -99,7 +99,7 @@ def init_state(self) -> TransformerNeRFTrainState:
latent_table = np.zeros((self.data_reader.identity_count,
self.num_latent_tokens, self.latent_token_dim))

rng = jax.random.PRNGKey(self.random_seed)
rng = jax.random.PRNGKey(self.random_seed) # pyrefly: ignore[bad-argument-type]
train_rng, model_init_rng = jax.random.split(rng)
model_parameters = self.model.initialize_parameters(model_init_rng,
self.num_latent_tokens,
Expand Down Expand Up @@ -155,7 +155,7 @@ def _loss_fn(params, data, rng, step):

return jax.pmap(per_device_train_step, "replicas")

def train_step(
def train_step( # pyrefly: ignore[bad-override]
self, train_state: TransformerNeRFTrainState, inputs: Dict[str, Any],
scratch: Optional[Dict[str, Any]]
) -> Tuple[TransformerNeRFTrainState, Dict[str, Any], Dict[str, Any]]:
Expand Down Expand Up @@ -199,7 +199,7 @@ def train_step(
steps = -latent_learning_rate * latent_grad
train_state.latent_table[latent_ids] += steps

new_train_state = train_state.replace(
new_train_state = train_state.replace( # pyrefly: ignore[missing-attribute]
model_parameters=new_model_parameters,
optimizer_state=new_optimizer_state,
latent_table=train_state.latent_table,
Expand All @@ -208,7 +208,7 @@ def train_step(

return new_train_state, loss_terms, scratch

def eval_step(self, train_state: TransformerNeRFTrainState,
def eval_step(self, train_state: TransformerNeRFTrainState, # pyrefly: ignore[bad-override]
summary_writer: tensorboard.SummaryWriter,
scratch: Optional[Any]) -> Any:
"""Performs evaluation on a model checkpoint.
Expand Down
2 changes: 1 addition & 1 deletion jax3d/projects/generative/nerf/lightfield/eval.py
Original file line number Diff line number Diff line change
Expand Up @@ -179,7 +179,7 @@ def flatten_views(t):
render = models.Model().apply(
model_parameters, inputs, rays, rng=rng, step=step)

pred = render["gamma_rgb"]
pred = render["gamma_rgb"] # pyrefly: ignore[bad-index]
gt = data_flat["gamma_rgb"]
if apply_mask:
pred *= data_flat["weight"]
Expand Down
Loading
Loading