การพยากรณ์เส้นเชื่อมและการจำแนกกราฟ
งานพยากรณ์เส้นเชื่อม การสุ่มตัวอย่างเชิงลบ การรวมระดับกราฟ และ GINConv สำหรับการจำแนกกราฟ
การพยากรณ์เส้นเชื่อมและการจำแนกกราฟ เป็นบทเรียน Learn AI with Python ฟรีบน CoddyKit นี่คือบทเรียนที่ 4 จากทั้งหมด 4 บทเรียน คุณสามารถอ่านบทเรียนทั้งหมดด้านล่างฟรี — จากนั้นลองปฏิบัติด้วยตัวคุณเองในเบราว์เซอร์พร้อมตัวแก้ไขโค้ดในตัวและติวเตอร์ AI ตลอด 24/7 บทเรียนนี้เป็นส่วนหนึ่งของเส้นทางการเรียน Learn AI with Python และความก้าวหน้าของคุณจะซิงค์ข้ามเว็บและแอป CoddyKit คอร์ส Learn AI with Python มีบทเรียนทั้งหมด 4 บทเรียน
งานกราฟใหม่สองประเภท
นอกเหนือจากการจำแนกโหนดแล้ว GNN ยังรองรับงานต่อไปนี้:
- การทำนายเส้นเชื่อม: จะมีเส้นเชื่อมระหว่างโหนดสองโหนดหรือไม่ (การแนะนำเพื่อน ปฏิสัมพันธ์ระหว่างยา)
- การจำแนกกราฟ: กำหนดป้ายกำกับให้กราฟทั้งกราฟ (โมเลกุลนี้มีพิษหรือไม่)
การตั้งค่าสำหรับการทำนายเส้นเชื่อม
ในการทำนายเส้นเชื่อม ขั้นแรกเราจะคำนวณเวกเตอร์ตัวแทนของโหนดด้วย GNN จากนั้นจึงให้คะแนนคู่โหนดที่เป็นตัวเลือก คะแนนสูงหมายความว่าโมเดลเชื่อว่าควรมีเส้นเชื่อมระหว่างโหนดคู่นั้น
การให้คะแนนเส้นเชื่อม
คะแนนเส้นเชื่อมที่ใช้กันทั่วไปคือผลคูณจุดของเวกเตอร์ตัวแทนโหนดทั้งสอง: score = dot(h_u, h_v) เวกเตอร์ตัวแทนที่คล้ายกันจะให้ผลคูณจุดสูง จึงทำนายได้ว่ามีแนวโน้มจะเป็นเส้นเชื่อม
h = gnn(data.x, data.edge_index) # node embeddings
score = (h[u] * h[v]).sum(dim=-1) # dot product per pairการสุ่มตัวอย่างเชิงลบ
กราฟจะแสดงเฉพาะเส้นเชื่อมที่มีอยู่ (ตัวอย่างบวก) หากต้องการฝึกตัวจำแนก เรายังต้องมีคู่โหนดที่ไม่มีเส้นเชื่อมด้วย การสุ่มตัวอย่างเชิงลบจะเลือกคู่โหนดที่ไม่เชื่อมต่อกันแบบสุ่มเพื่อใช้เป็นตัวอย่างลบ ทำให้ชุดฝึกมีความสมดุล
from torch_geometric.utils import negative_sampling
neg_edge_index = negative_sampling(
edge_index=data.edge_index,
num_nodes=data.num_nodes,
num_neg_samples=data.edge_index.size(1),
)BCEWithLogitsLoss
การทำนายเส้นเชื่อมเป็นปัญหาแบบทวิภาค (มีเส้นเชื่อมหรือไม่มีเส้นเชื่อม) เราให้คะแนนคู่ตัวอย่างบวกและลบ กำหนดป้ายกำกับเป็น 1 และ 0 แล้วฝึกด้วย BCEWithLogitsLoss ซึ่งรวมซิกมอยด์เข้ากับ binary cross-entropy อย่างมีเสถียรภาพทางตัวเลข
import torch
pos = (h[pos_u] * h[pos_v]).sum(-1)
neg = (h[neg_u] * h[neg_v]).sum(-1)
scores = torch.cat([pos, neg])
labels = torch.cat([torch.ones_like(pos), torch.zeros_like(neg)])
loss = torch.nn.functional.binary_cross_entropy_with_logits(scores, labels)การเปลี่ยนไปใช้การจำแนกกราฟ
สำหรับการจำแนกกราฟ เราต้องการเวกเตอร์เดียวต่อกราฟ ไม่ใช่เวกเตอร์ต่อโหนด หลังจากชั้น GNN สร้างเวกเตอร์ตัวแทนของโหนดแล้ว เราจะรวมเวกเตอร์เหล่านั้นเป็นตัวแทนระดับกราฟหนึ่งชุด
global_mean_pool
global_mean_poolคำนวณค่าเฉลี่ยของเวกเตอร์ตัวแทนโหนดทั้งหมดในกราฟเพื่อสร้างเวกเตอร์ขนาดคงที่หนึ่งชุด ไม่ว่ากราฟจะมีขนาดเท่าใด ดัชนี batch จะระบุว่าโหนดใดอยู่ในกราฟใดเมื่อรวมกราฟหลายกราฟเป็นชุดเดียวกัน
from torch_geometric.nn import global_mean_pool
h = gnn(x, edge_index) # [num_nodes, dim]
hg = global_mean_pool(h, batch) # [num_graphs, dim]
logits = classifier(hg)เหตุใดการรวมข้อมูลจึงสำคัญ
การรวมข้อมูลทำให้โมเดลไม่ขึ้นกับลำดับของโหนดและขนาดของกราฟ ดังนั้นกราฟที่มีโครงสร้างเหมือนกันสองกราฟจึงให้เวกเตอร์หลังการรวมเหมือนกัน การรวมด้วยค่าเฉลี่ยเป็นวิธีที่เรียบง่าย ส่วนการรวมด้วยผลรวมและค่าสูงสุดเป็นทางเลือกที่มีความไวต่อคุณสมบัติต่างกัน
GINConv
GINConv (เครือข่ายไอโซมอร์ฟิซึมของกราฟ) คือคอนโวลูชันที่มีความสามารถในการแสดงออกสูงกว่า โดยใช้ MLP และการรวมด้วยผลรวม ซึ่งออกแบบมาโดยเฉพาะเพื่อเพิ่มความสามารถในการจำแนกของการส่งข้อความสำหรับงานระดับกราฟให้สูงสุด
from torch_geometric.nn import GINConv
import torch
mlp = torch.nn.Sequential(
torch.nn.Linear(in_dim, hid),
torch.nn.ReLU(),
torch.nn.Linear(hid, hid),
)
conv = GINConv(mlp)ความเชื่อมโยงกับไวส์ไฟเลอร์–เลมัน
GIN ได้รับการออกแบบให้มีพลังเทียบเท่ากับการทดสอบไวส์ไฟเลอร์–เลมัน (WL) ซึ่งเป็นอัลกอริทึมดั้งเดิมสำหรับแยกแยะกราฟที่ไม่เป็นไอโซมอร์ฟิก GNN ที่เรียบง่ายกว่าหลายแบบไม่สามารถแยกกราฟบางคู่ได้ แต่ GIN ทำได้ภายใต้ขีดจำกัดของการทดสอบ WL จึงเหมาะสำหรับการจำแนกกราฟ
การเลือกเครื่องมือที่เหมาะสม
จับคู่สถาปัตยกรรมให้เหมาะกับงาน:
- การทำนายลิงก์: เวกเตอร์ฝังจาก GNN + การให้คะแนนด้วยผลคูณจุด + การสุ่มตัวอย่างเชิงลบ + ฟังก์ชันสูญเสีย BCE
- การจำแนกกราฟ: คอนโวลูชันที่มีความสามารถในการแสดงออกสูง เช่น GINConv + การพูลรวมทั่วทั้งกราฟ + ตัวจำแนก
ตรวจสอบความเข้าใจอย่างรวดเร็ว
ทดสอบความรู้ของคุณ
ทบทวน
คุณได้เรียนรู้เรื่องการทำนายลิงก์และการจำแนกกราฟ:
- คะแนนของเส้นเชื่อม =
dot(h_u, h_v)โดยฝึกด้วยการสุ่มตัวอย่างเชิงลบและฟังก์ชันสูญเสีย BCE พร้อมลอจิต - global_mean_pool แปลงเวกเตอร์ฝังของโหนดให้เป็นเวกเตอร์ระดับกราฟ
- GINConv มีความสามารถในการแสดงออกสูงมาก โดยสอดคล้องกับการทดสอบไวส์ไฟเลอร์–เลมัน
คำถามที่พบบ่อย
บทเรียน “การพยากรณ์เส้นเชื่อมและการจำแนกกราฟ” ฟรีหรือไม่
ใช่ — ข้อความเต็มของ “การพยากรณ์เส้นเชื่อมและการจำแนกกราฟ” ฟรีให้อ่านที่นี่บนเว็บ เพื่อปฏิบัติแบบโต้ตอบ (ตัวแก้ไขโค้ดในตัวและติวเตอร์ AI ตลอด 24/7) และปลดล็อคส่วนที่เหลือของคอร์ส Learn AI with Python ให้อัปเกรดเป็น CoddyKit PRO คอร์ส Learn AI with Python มีบทเรียนทั้งหมด 4 บทเรียน
คุณจะเรียนรู้อะไรในบทเรียน “การพยากรณ์เส้นเชื่อมและการจำแนกกราฟ”
งานพยากรณ์เส้นเชื่อม การสุ่มตัวอย่างเชิงลบ การรวมระดับกราฟ และ GINConv สำหรับการจำแนกกราฟ คุณปฏิบัติ Learn AI with Python ด้วยโค้ดที่ใช้งานได้จริงที่คุณเรียกใช้โดยตรงในเบราว์เซอร์ และติวเตอร์ AI ตลอด 24/7 ตอบคำถามของคุณขณะที่คุณไปผ่านบทเรียน
คุณต้องมีประสบการณ์ก่อนที่จะเริ่มเรียน Learn AI with Python หรือไม่
ไม่จำเป็นต้องมีประสบการณ์มาก่อน Learn AI with Python บน CoddyKit ออกแบบมาสำหรับผู้เริ่มต้นไปจนถึงผู้เรียนขั้นสูง คุณสามารถเริ่มต้นที่นี่หรือเริ่มจากตัวแรกและเรียนด้วยความเร็วของคุณเอง นี่คือบทเรียนที่ 4 จากทั้งหมด 4 บทเรียน
บทเรียน “การพยากรณ์เส้นเชื่อมและการจำแนกกราฟ” ใช้เวลานานแค่ไหน
บทเรียน CoddyKit ส่วนใหญ่ใช้เวลาประมาณ 5–10 นาที แต่ละบทเรียนจึงสั้นและเป็นแบบโต้ตอบ คุณสามารถก้าวหน้าอย่างต่อเนื่องและกลับมาเรียนต่อจากตรงที่เพิ่งหยุดบนเว็บและแอปได้เลย
ฉันเขียนและรันโค้ดในบทเรียน Learn AI with Python นี้ได้ไหม
ได้ บทเรียน Learn AI with Python ทุกบทมีตัวแก้ไขโค้ดในตัว คุณจึงเขียนและรันโค้ดจริงได้เลยในเบราว์เซอร์ และได้รับข้อเสนอแนะจาก AI ในทันที — ไม่ต้องติดตั้งในเครื่องของคุณ
บทเรียนทั้งหมดในหลักสูตรนี้
- ทฤษฎีกราฟสำหรับการเรียนรู้ของเครื่อง
- โครงข่ายกราฟคอนโวลูชัน (GCN)
- การจำแนกโหนดด้วย GNN
- การพยากรณ์เส้นเชื่อมและการจำแนกกราฟ