-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathhyena_learned_context_encoder.py
More file actions
102 lines (94 loc) · 5.49 KB
/
Copy pathhyena_learned_context_encoder.py
File metadata and controls
102 lines (94 loc) · 5.49 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
from trainer import *
import os
import h5py
from standalone_hyenadna import HyenaDNAModel
from hyena_encoder import get_model
def main(args):
ddp_setup()
model = get_model(args)
dataset, mini_test_set, optimizer, warmup_scheduler, main_scheduler = load_train_objs(args, model)
seq_len = len(dataset[0][1])
use_devices = [int(item) for item in args.devices.split(',')]
print(f"using devices: {use_devices}")
if os.environ["RANK"] == "0":
num_params = sum(p.numel() for p in model.parameters())
print(f"model has {num_params} parameters")
train_data = prepare_dataloader(dataset, args.batch_size)
mini_test_data = prepare_dataloader(mini_test_set, args.batch_size)
verify_data = prepare_dataloader(verify_set, args.batch_size)
trainer = HyenaLearnedContextEncoderTrainer(
args.view_size,
args.context_size,
args.step_progress,
seq_len,
args.context_lr,
args.context_cycle_frac,
args.batches_per_cycle,
model,
train_data,
mini_test_data,
verify_data,
optimizer,
warmup_scheduler,
main_scheduler,
args.save_every,
args.snapshot_path,
use_devices,
args,
args.limit_steps,
args.gradual_length,
args.mask_frac,
)
self = trainer #TODO: remove
# with sdp_kernel(**backend_map[SDPBackend.FLASH_ATTENTION]) if (has_sdp and not args.no_gradscaler) else nullcontext():
# with torch.autocast(device_type='cuda', dtype=torch.bfloat16):
trainer.train(args.total_epochs)
destroy_process_group()
def interactive_setup():
os.environ["LOCAL_RANK"] = "0"
os.environ["RANK"] = "0"
os.environ["WORLD_SIZE"] = "1"
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "12355"
if __name__ == "__main__":
parser = argparse.ArgumentParser(description='simple distributed training job')
parser.add_argument('--total-epochs', default=2, type=int, help='Total epochs to train the model')
parser.add_argument('--save-every', type=int, default=1, help='How often to save a snapshot')
parser.add_argument('--batch-size', default=32, type=int, help='Input batch size on each device (default: 32)')
# parser.add_argument("--embed-dim", type=int, default="64")
parser.add_argument("--num-layers", type=int, default="2") # 2 - 8 used in paper
parser.add_argument("--num-heads", type=int, default="8")
parser.add_argument("--seq-len", type=int, default="320") # must be updated to match input file
parser.add_argument("--test-frac", type=float, default="0.3")
# parser.add_argument("--h5-file", default="/data/ukbb/net_input/gwas_ldprune_320.h5")
parser.add_argument("--h5-file", default="/data/ukbb/net_input/all_gwas_v6.h5")
parser.add_argument('-d', '--devices', help='comma delimited list of gpus to use', type=str, default="0")
parser.add_argument("--limit-steps", type=int, default="-1")
# parser.add_argument("--encoder-snapshot", default="/data/ukbb/v2_snapshots/hyena_encoder.pt")
parser.add_argument("--augment-data", action='store_true')
parser.add_argument("--warmup-min-lr", type=float, default="1e-5")
parser.add_argument("--warmup-max-lr", type=float, default="1e-2")
parser.add_argument("--warmup-batches", type=int, default="500")
parser.add_argument("--main-scheduler-step-size", type=int, default="100")
parser.add_argument("--main-scheduler-gamma", type=float, default="0.9")
parser.add_argument("--verbose-scheduler", action='store_true')
parser.add_argument("--report-on-batch", type=int, default="50")
parser.add_argument("--no-gradscaler", action='store_true')
parser.add_argument("--dropout", type=float, default="0.2")
parser.add_argument("--augment-frac", type=float, default="0.15", help="fraction of snv sequence to mask/change when augmenting gout data")
parser.add_argument("--augment-mult", type=float, default="2.0", help="ratio of augmented to actual gout data")
# parser.add_argument("--encoder-type", default="classic", help="\'classic\' or \'linformer\'")
# parser.add_argument("--output-transform-layers", type=int, default="5")
parser.add_argument("--snapshot-path", default="/data/ukbb/v2_snapshots/test_hyena_learned_context_encoder.pt")
parser.add_argument("--d-model", type=int, default = "128") # 128 and 256 used in HyenaDNA paper
parser.add_argument("--pretrained-hyena", action='store_const', const=None)
parser.add_argument("--gradual-length", action='store_true')
parser.add_argument("--mask-frac", type=float, default = "0.5")
parser.add_argument("--view-size", type=int, default=128)
parser.add_argument("--context-size", type=int, default=512)
parser.add_argument("--step-progress", action='store_true')
parser.add_argument("--context-lr", type=float, default = "0.1")
parser.add_argument("--context-cycle-frac", type=float, default = "0.02")
parser.add_argument("--batches-per-cycle", type=int, default=128)
args = parser.parse_args()
main(args)