Compare commits
No commits in common. "6e062e75d69471e7479071311ff0ef4f1f59a41d" and "46c273e0ddd2589eee6111b38d499a5908e2005a" have entirely different histories.
6e062e75d6
...
46c273e0dd
10
.gitignore
vendored
10
.gitignore
vendored
@ -1,10 +0,0 @@
|
|||||||
__pycache__/
|
|
||||||
*.pyc
|
|
||||||
*.pyo
|
|
||||||
logs/
|
|
||||||
*.log
|
|
||||||
nohup.out
|
|
||||||
py3/
|
|
||||||
*.egg-info/
|
|
||||||
*.pid
|
|
||||||
venv/
|
|
||||||
177
app.py
177
app.py
@ -1,177 +0,0 @@
|
|||||||
#!/usr/bin/env python3
|
|
||||||
"""
|
|
||||||
GLiNER NER Service - Zero-shot entity extraction
|
|
||||||
Port: 9093, GPU 5
|
|
||||||
"""
|
|
||||||
import os
|
|
||||||
import re
|
|
||||||
import logging
|
|
||||||
from fastapi import FastAPI, HTTPException
|
|
||||||
from fastapi.middleware.cors import CORSMiddleware
|
|
||||||
from pydantic import BaseModel
|
|
||||||
from typing import List
|
|
||||||
import uvicorn
|
|
||||||
|
|
||||||
logging.basicConfig(level=logging.INFO)
|
|
||||||
logger = logging.getLogger("ner-service")
|
|
||||||
|
|
||||||
app = FastAPI(title="GLiNER NER Service", version="1.0.0")
|
|
||||||
app.add_middleware(
|
|
||||||
CORSMiddleware,
|
|
||||||
allow_origins=["*"],
|
|
||||||
allow_methods=["*"],
|
|
||||||
allow_headers=["*"],
|
|
||||||
)
|
|
||||||
|
|
||||||
# Global model
|
|
||||||
model = None
|
|
||||||
MODEL_PATH = "/data/ymq/models/gliner-multitask-large-v0.5"
|
|
||||||
|
|
||||||
|
|
||||||
def load_model():
|
|
||||||
global model
|
|
||||||
try:
|
|
||||||
from gliner import GLiNER
|
|
||||||
logger.info(f"Loading GLiNER from {MODEL_PATH}...")
|
|
||||||
model = GLiNER.from_pretrained(MODEL_PATH, local_files_only=True)
|
|
||||||
logger.info("Model loaded successfully")
|
|
||||||
return True
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Failed to load model: {e}")
|
|
||||||
return False
|
|
||||||
|
|
||||||
|
|
||||||
class ExtractRequest(BaseModel):
|
|
||||||
text: str
|
|
||||||
entities: List[str] = ["person", "company", "product", "location", "concept", "event"]
|
|
||||||
threshold: float = 0.5
|
|
||||||
|
|
||||||
|
|
||||||
class Entity(BaseModel):
|
|
||||||
text: str
|
|
||||||
label: str
|
|
||||||
start: int
|
|
||||||
end: int
|
|
||||||
score: float
|
|
||||||
|
|
||||||
|
|
||||||
class Relation(BaseModel):
|
|
||||||
source: str
|
|
||||||
target: str
|
|
||||||
relation: str
|
|
||||||
confidence: float
|
|
||||||
|
|
||||||
|
|
||||||
class ExtractResponse(BaseModel):
|
|
||||||
entities: List[Entity]
|
|
||||||
relations: List[Relation]
|
|
||||||
|
|
||||||
|
|
||||||
@app.on_event("startup")
|
|
||||||
async def startup():
|
|
||||||
success = load_model()
|
|
||||||
if not success:
|
|
||||||
logger.warning("Model not loaded, /api/extract will return 503")
|
|
||||||
|
|
||||||
|
|
||||||
@app.get("/api/status")
|
|
||||||
async def status():
|
|
||||||
return {
|
|
||||||
"service": "ner-service",
|
|
||||||
"status": "ready" if model else "no_model",
|
|
||||||
"model": "gliner-multitask-large-v0.5",
|
|
||||||
"model_path": MODEL_PATH,
|
|
||||||
"model_loaded": model is not None,
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
@app.post("/api/extract", response_model=ExtractResponse)
|
|
||||||
async def extract(req: ExtractRequest):
|
|
||||||
if not model:
|
|
||||||
raise HTTPException(status_code=503, detail="Model not loaded")
|
|
||||||
|
|
||||||
if not req.text or len(req.text.strip()) == 0:
|
|
||||||
return ExtractResponse(entities=[], relations=[])
|
|
||||||
|
|
||||||
try:
|
|
||||||
entities_raw = model.predict_entities(
|
|
||||||
req.text, req.entities, threshold=req.threshold
|
|
||||||
)
|
|
||||||
|
|
||||||
entities = [
|
|
||||||
Entity(
|
|
||||||
text=e["text"],
|
|
||||||
label=e["label"],
|
|
||||||
start=e["start"],
|
|
||||||
end=e["end"],
|
|
||||||
score=round(e["score"], 4),
|
|
||||||
)
|
|
||||||
for e in entities_raw
|
|
||||||
]
|
|
||||||
|
|
||||||
relations = extract_relations(req.text, entities)
|
|
||||||
return ExtractResponse(entities=entities, relations=relations)
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Extraction failed: {e}", exc_info=True)
|
|
||||||
raise HTTPException(status_code=500, detail=str(e))
|
|
||||||
|
|
||||||
|
|
||||||
def extract_relations(text: str, entities: List[Entity]) -> List[Relation]:
|
|
||||||
"""Rule-based relation extraction using co-occurrence in sentences."""
|
|
||||||
relations = []
|
|
||||||
entity_texts = {e.text: e.label for e in entities}
|
|
||||||
|
|
||||||
sentences = re.split(r'[。!?.!?\n]+', text)
|
|
||||||
|
|
||||||
relation_patterns = [
|
|
||||||
(r'(.+?)在(.+?)(?:担任|工作|任职)', "works_at"),
|
|
||||||
(r'(.+?)是(.+?)的(?:CEO|总裁|董事长|创始人)', "leads"),
|
|
||||||
(r'(.+?)收购(?:了)?(.+?)', "acquired"),
|
|
||||||
(r'(.+?)发布(?:了)?(.+?)', "released"),
|
|
||||||
(r'(.+?)位于(.+?)', "located_in"),
|
|
||||||
(r'(.+?)投资(?:了)?(.+?)', "invested_in"),
|
|
||||||
(r'(.+?)属于(.+?)', "belongs_to"),
|
|
||||||
(r'(.+?)与(.+?)合作', "collaborates"),
|
|
||||||
]
|
|
||||||
|
|
||||||
for sent in sentences:
|
|
||||||
sent = sent.strip()
|
|
||||||
if not sent:
|
|
||||||
continue
|
|
||||||
|
|
||||||
sent_entities = [e for e in entities if e.text in sent]
|
|
||||||
if len(sent_entities) < 2:
|
|
||||||
continue
|
|
||||||
|
|
||||||
# Pattern-based relations
|
|
||||||
for pattern, rel_type in relation_patterns:
|
|
||||||
match = re.search(pattern, sent)
|
|
||||||
if match:
|
|
||||||
groups = match.groups()
|
|
||||||
if len(groups) >= 2:
|
|
||||||
src, tgt = groups[0].strip(), groups[1].strip()
|
|
||||||
if src in entity_texts and tgt in entity_texts:
|
|
||||||
relations.append(Relation(
|
|
||||||
source=src, target=tgt,
|
|
||||||
relation=rel_type, confidence=0.8,
|
|
||||||
))
|
|
||||||
|
|
||||||
# Co-occurrence fallback
|
|
||||||
seen_pairs = set()
|
|
||||||
for i in range(len(sent_entities)):
|
|
||||||
for j in range(i + 1, len(sent_entities)):
|
|
||||||
e1, e2 = sent_entities[i], sent_entities[j]
|
|
||||||
pair = tuple(sorted([e1.text, e2.text]))
|
|
||||||
if pair not in seen_pairs:
|
|
||||||
seen_pairs.add(pair)
|
|
||||||
relations.append(Relation(
|
|
||||||
source=pair[0], target=pair[1],
|
|
||||||
relation="co_occur", confidence=0.5,
|
|
||||||
))
|
|
||||||
|
|
||||||
return relations
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
uvicorn.run(app, host="0.0.0.0", port=9093)
|
|
||||||
Loading…
x
Reference in New Issue
Block a user