GRADE / src /Baselines /radarcam-depth /utils /log_utils.py
Bin-0815's picture
Release all GRADE models, checkpoints, and reviewed evaluation code (part 2)
f348660 verified
Raw History Blame Contribute Delete
1.83 kB
import os
import torch
import numpy as np
from matplotlib import pyplot as plt
def log(s, filepath=None, to_console=True):
'''
Logs a string to either file or console
Arg(s):
s : str
string to log
filepath
output filepath for logging
to_console : bool
log to console
'''
if to_console:
print(s)
if filepath is not None:
if not os.path.isdir(os.path.dirname(filepath)):
os.makedirs(os.path.dirname(filepath))
with open(filepath, 'w+') as o:
o.write(s + '\n')
else:
with open(filepath, 'a+') as o:
o.write(s + '\n')
def colorize(T, colormap='magma', return_numpy=False):
'''
Colorizes a 1-channel tensor with matplotlib colormaps
Arg(s):
T : torch.Tensor[float32]
1-channel tensor
colormap : str
matplotlib colormap
'''
cm = plt.cm.get_cmap(colormap)
shape = T.shape
# Convert to numpy array and transpose
if shape[0] > 1:
T = np.squeeze(np.transpose(T.cpu().numpy(), (0, 2, 3, 1)))
else:
T = np.squeeze(np.transpose(T.cpu().numpy(), (0, 2, 3, 1)), axis=-1)
# Colorize using colormap
color = np.concatenate([
np.expand_dims(cm(T[n, ...])[..., 0:3], 0) for n in range(T.shape[0])],
axis=0)
if return_numpy:
return color
else:
# Transpose back to torch format
color = np.transpose(color, (0, 3, 1, 2))
# Convert back to tensor
return torch.from_numpy(color.astype(np.float32))
def log_params(log_path, params_dict):
with open(log_path, 'w') as log_file:
for param_name, param_value in params_dict.items():
log_file.write(f"{param_name}: {param_value}\n")