SEAGAN: an edge-aware graph attention network for node classification on A-Ci curve graphs.
Project description
SEAGAN
SEAGAN is an edge-aware Graph Attention Network (GAT) for node classification on A-Ci curve graphs. It represents each curve as a graph, uses edge features in the attention layers, and supports weighted focal loss. For the sample SEAGAN model, the focal lambda/gamma is 0.0.
Install
Install the published package from PyPI:
pip install seagan
Model Overview
SEAGAN represents each curve as a graph:
- Node features:
Ci,Anet, computedAc, and computedAj - Edge index: k-nearest-neighbor graph in the
[Ci, Anet]space - Edge features: pairwise differences
[dAc, dAj] - GAT layer 1: input dimension
4, output dimension64,5attention heads, dropout0.2 - GAT layer 2: hidden representation from the first layer, output dimension
64,5attention heads, dropout0.2 - MLP head: maps
64hidden features to3node classes
Quick Start
Build graphs from your data:
from seagan import build_graphs_from_df
# df_points needs curve_id, Ci, Anet, and ID columns.
# df_params needs curve_id and Tleaf columns.
graphs, class_map = build_graphs_from_df(df_points, df_params)
Load the sample checkpoint included with the package:
from seagan import load_pretrained_seagan, standardize_graphs_from_checkpoint
model, checkpoint = load_pretrained_seagan()
graphs = standardize_graphs_from_checkpoint(graphs, checkpoint)
Or load your own checkpoint:
from seagan import load_pretrained_seagan
model, checkpoint = load_pretrained_seagan("path/to/your_checkpoint.pt")
Run inference on one graph:
from seagan import predict_graph
predicted_classes = predict_graph(model, graphs[0], one_indexed=True)
print(predicted_classes)
Data Format
df_points should contain one row per point in a curve:
| column | meaning |
|---|---|
curve_id |
curve identifier |
Ci |
intercellular CO2 concentration |
Anet |
net assimilation |
ID |
node class label, commonly 1, 2, or 3 |
df_params should contain one row per curve:
| column | meaning |
|---|---|
curve_id |
curve identifier |
Tleaf |
leaf temperature in Celsius |
If curve_id is already the DataFrame index, SEAGAN will use it directly.
API
Common imports:
from seagan import (
SEAGAN,
build_graphs_from_df,
build_graphs_single,
compute_edge_Ac_Aj,
focal_loss_multiclass,
load_checkpoint,
load_pretrained_seagan,
predict_graph,
standardize_graphs_from_checkpoint,
)
Build and Publish
For maintainers:
python -m pip install -U build twine
python -m build
python -m twine check dist/*
Upload to PyPI:
python -m twine upload dist/*
Project details
Download files
Download the file for your platform. If you're not sure which to choose, learn more about installing packages.
Source Distribution
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
File details
Details for the file seagan-0.1.5.tar.gz.
File metadata
- Download URL: seagan-0.1.5.tar.gz
- Upload date:
- Size: 1.1 MB
- Tags: Source
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/6.2.0 CPython/3.12.4
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
1152c9b31eb264ed94a4297a88b286c412e89f946a48fc7c446451f90190bbdf
|
|
| MD5 |
6541aa4562b09af436aad76f5459f069
|
|
| BLAKE2b-256 |
c0b8eb03490f2e30b089212b46f932bc2780d7055482638156defb0fac30b0a5
|
File details
Details for the file seagan-0.1.5-py3-none-any.whl.
File metadata
- Download URL: seagan-0.1.5-py3-none-any.whl
- Upload date:
- Size: 420.6 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/6.2.0 CPython/3.12.4
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
1552a41db4b29008066b218b9804a148d9bf142075f33e6e0ea3aa5bcd5b770c
|
|
| MD5 |
37a946f601bb78f813c599770c2f8c12
|
|
| BLAKE2b-256 |
5a792ba0d11a6893cc029a1649ee83d5e0b092fd6bb2dadfaa2144bec82c30e4
|