| |
| |
|
|
| |
|
|
|
|
| |
| |
| import os |
| import sys |
| import json |
| import yaml |
| import numpy as np |
| import math |
| import time |
| import datetime |
| import random |
| from tqdm import tqdm |
| import webdataset as wds |
| import matplotlib.pyplot as plt |
| import pandas as pd |
| import torch |
| import torch.nn as nn |
| from torchvision import transforms |
| import utils |
| from mae_utils import flat_models |
| from elbow.sinks import BufferedParquetWriter |
|
|
| |
| if utils.is_interactive(): |
| model_name = "NSDflat_large_gsrFalse_" |
| else: |
| model_name = sys.argv[1] |
| outdir = os.path.abspath(f'checkpoints/{model_name}') |
| print("outdir", outdir) |
|
|
| |
| assert os.path.exists(f"{outdir}/config.yaml") |
| config = yaml.load(open(f"{outdir}/config.yaml", 'r'), Loader=yaml.FullLoader) |
| print(f"Loaded config.yaml from ckpt folder {outdir}") |
| |
| print("\n__CONFIG__") |
| for attribute_name in config.keys(): |
| print(f"{attribute_name} = {config[attribute_name]}") |
| globals()[attribute_name] = config[f'{attribute_name}'] |
| print("\n") |
|
|
| if utils.is_interactive(): |
| |
| |
| get_ipython().run_line_magic('load_ext', 'autoreload') |
| get_ipython().run_line_magic('autoreload', '2') |
|
|
| device = torch.device('cuda') |
|
|
| print("PID of this process =",os.getpid()) |
|
|
| |
| utils.seed_everything(seed) |
|
|
|
|
| |
|
|
|
|
| os.environ['HCP_FLAT_ROOT'] = hcp_flat_path |
|
|
|
|
| |
|
|
|
|
| if os.getenv('global_pool') == "False": |
| global_pool = False |
| else: |
| global_pool = True |
| print(f"global_pool = {global_pool}") |
|
|
| try: |
| gsr |
| except: |
| gsr = True |
| print("set gsr to True") |
| print(f"gsr = {gsr}") |
|
|
|
|
| |
|
|
| |
|
|
|
|
| from mae_utils.flat import load_hcp_flat_mask |
| from mae_utils.flat import create_hcp_flat |
| from mae_utils.flat import batch_unmask |
| import mae_utils.visualize as vis |
|
|
| flat_mask = load_hcp_flat_mask(hcp_flat_path) |
|
|
| model = flat_models.mae_vit_large_fmri( |
| patch_size=patch_size, |
| decoder_embed_dim=decoder_embed_dim, |
| t_patch_size=t_patch_size, |
| pred_t_dim=pred_t_dim, |
| decoder_depth=4, |
| cls_embed=cls_embed, |
| norm_pix_loss=norm_pix_loss, |
| no_qkv_bias=no_qkv_bias, |
| sep_pos_embed=sep_pos_embed, |
| trunc_init=trunc_init, |
| pct_masks_to_decode=pct_masks_to_decode, |
| img_mask=flat_mask, |
| ) |
|
|
|
|
| |
|
|
| |
|
|
|
|
| checkpoint_files = [f for f in os.listdir(outdir) if f.endswith('.pth')] |
|
|
| if utils.is_interactive(): |
| latest_checkpoint = "epoch99.pth" |
| else: |
| latest_checkpoint = sys.argv[2] |
| print(f"latest_checkpoint: {latest_checkpoint}") |
|
|
| |
| checkpoint_path = os.path.join(outdir, latest_checkpoint) |
|
|
| state = torch.load(checkpoint_path) |
| model.load_state_dict(state["model_state_dict"], strict=False) |
| model.to(device) |
| model.eval() |
|
|
| print(f"\nLoaded checkpoint {latest_checkpoint} from {outdir}\n") |
|
|
|
|
| |
|
|
| |
|
|
|
|
| from torch.utils.data import default_collate |
| batch_size = 1 |
| print(f"changed batch_size to {batch_size}") |
|
|
| |
| datasets_to_include = "HCP" |
| assert "HCP" in datasets_to_include |
| test_dataset = create_hcp_flat(root=hcp_flat_path, |
| clip_mode="event", frames=num_frames, shuffle=False, gsr=gsr, sub_list = 'test') |
| test_dl = wds.WebLoader( |
| test_dataset.batched(batch_size, partial=False, collation_fn=default_collate), |
| batch_size=None, |
| shuffle=False, |
| num_workers=num_workers, |
| pin_memory=True, |
| ) |
|
|
| |
| assert "HCP" in datasets_to_include |
| train_dataset = create_hcp_flat(root=hcp_flat_path, |
| clip_mode="event", frames=num_frames, shuffle=False, gsr=gsr, sub_list = 'train') |
| train_dl = wds.WebLoader( |
| train_dataset.batched(batch_size, partial=False, collation_fn=default_collate), |
| batch_size=None, |
| shuffle=False, |
| num_workers=num_workers, |
| pin_memory=True, |
| ) |
|
|
|
|
| |
|
|
| |
|
|
|
|
| cnt = 9999 |
|
|
|
|
| |
|
|
|
|
| @torch.no_grad() |
| def extract_features(dl, global_pool=True): |
| for samples in tqdm(dl,total=cnt): |
| samples_meta = samples['meta'] |
| features = model(samples['image'].to(device),global_pool=global_pool, forward_features = True) |
| features = features.flatten(1) |
| features = features.cpu().numpy() |
| meta_dict = {} |
| for key, value in samples_meta.items(): |
| if type(value) == torch.Tensor: |
| value = value.cpu().numpy() |
| meta_dict[key] = value |
| for feat, meta in zip(features, samples_meta): |
| yield {"feature": feat, **meta_dict} |
|
|
|
|
| |
|
|
|
|
| out_folder = f'{outdir}_gp{global_pool}/{latest_checkpoint[:-4]}/HCP' |
| print(out_folder) |
| os.makedirs(out_folder,exist_ok=True) |
|
|
|
|
| |
|
|
|
|
| |
| outdir_parquet = os.path.join(f'{outdir}_gp{global_pool}/{latest_checkpoint[:-4]}', 'HCP') |
| os.makedirs(outdir_parquet, exist_ok=True) |
|
|
| utils.seed_everything(seed) |
|
|
| print("Start extract") |
| start_time = time.time() |
|
|
| with BufferedParquetWriter(f"{outdir_parquet}/test.parquet", blocking=True) as writer: |
| for sample in extract_features(test_dl, global_pool): |
| writer.write(sample) |
|
|
| with BufferedParquetWriter(f"{outdir_parquet}/train.parquet", blocking=True) as writer: |
| for sample in extract_features(train_dl, global_pool): |
| writer.write(sample) |
|
|
| total_time = time.time() - start_time |
| total_time_str = str(datetime.timedelta(seconds=int(total_time))) |
| print("Extract time {}".format(total_time_str)) |
| print(torch.cuda.memory_allocated()) |
|
|
|
|
| |
|
|
|
|
|
|
|
|
|
|