From ed70db1e048b4d1178859a7fd585dd8f983db52d Mon Sep 17 00:00:00 2001 From: symoon9 Date: Fri, 21 Apr 2023 16:14:57 -0700 Subject: [PATCH] Remove loss_list from train code --- STAGATE_pyG/Train_STAGATE.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/STAGATE_pyG/Train_STAGATE.py b/STAGATE_pyG/Train_STAGATE.py index 09f995f..17adc65 100644 --- a/STAGATE_pyG/Train_STAGATE.py +++ b/STAGATE_pyG/Train_STAGATE.py @@ -75,13 +75,11 @@ def train_STAGATE(adata, hidden_dims=[512, 30], n_epochs=1000, lr=0.001, key_add optimizer = torch.optim.Adam(model.parameters(), lr=lr, weight_decay=weight_decay) - loss_list = [] for epoch in tqdm(range(1, n_epochs+1)): model.train() optimizer.zero_grad() z, out = model(data.x, data.edge_index) loss = F.mse_loss(data.x, out) #F.nll_loss(out[data.train_mask], data.y[data.train_mask]) - loss_list.append(loss) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), gradient_clipping) optimizer.step()