File size: 5,067 Bytes
4a4df15
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
"""Fresh-interpreter CPU workers: never execute embedding/ASR inference in ZeroGPU's parent."""
from __future__ import annotations

import atexit
import json
import os
from pathlib import Path
import queue
import subprocess
import sys
import threading
import time
import uuid
from typing import Callable


class WorkerError(RuntimeError):
    pass


class CPUWorker:
    def __init__(self, kind: str):
        self.kind = kind
        self.proc: subprocess.Popen | None = None
        self.inbox: queue.Queue | None = None
        self.lock = threading.Lock()
        atexit.register(self.close)

    def _start(self) -> None:
        env = os.environ.copy()
        env.update(CUDA_VISIBLE_DEVICES="", NVIDIA_VISIBLE_DEVICES="", TOKENIZERS_PARALLELISM="false",
                   OMP_NUM_THREADS=os.getenv("LOCAL_CPU_THREADS", "2"), MKL_NUM_THREADS=os.getenv("LOCAL_CPU_THREADS", "2"), HF_HUB_DISABLE_PROGRESS_BARS="1",
                   HF_HUB_DISABLE_TELEMETRY="1", TRANSFORMERS_VERBOSITY="error", DO_NOT_TRACK="1")
        for name in ("HF_TOKEN", "HUGGING_FACE_HUB_TOKEN", "HF_TOKEN_PATH"):
            env.pop(name, None)
        env["HF_HUB_DISABLE_IMPLICIT_TOKEN"] = "1"
        if env.get("TRAINER_OFFLINE_MODELS") == "1":
            env["HF_HUB_OFFLINE"] = "1"
            env["TRANSFORMERS_OFFLINE"] = "1"
        # No model credentials are required for the default public checkpoints.
        # subprocess performs exec into a fresh interpreter, not a multiprocessing fork worker.
        self.proc = subprocess.Popen([sys.executable, "-u", "-m", "engine.cpu_worker", self.kind],
                                     cwd=str(Path(__file__).resolve().parents[1]), env=env,
                                     stdin=subprocess.PIPE, stdout=subprocess.PIPE, stderr=None,
                                     text=True, encoding="utf-8", bufsize=1)
        inbox: queue.Queue = queue.Queue()
        self.inbox = inbox
        proc = self.proc

        def read_stdout() -> None:
            assert proc.stdout is not None
            try:
                for line in proc.stdout:
                    if len(line) > 35_000_000:
                        inbox.put({"fatal": "CPU worker response exceeded its size limit."})
                        break
                    try:
                        inbox.put(json.loads(line))
                    except json.JSONDecodeError:
                        continue
            finally:
                inbox.put({"fatal": "CPU worker stopped unexpectedly."})
        threading.Thread(target=read_stdout, name=f"{self.kind}-reader", daemon=True).start()

    def close(self) -> None:
        proc, self.proc = self.proc, None
        self.inbox = None
        if proc is not None:
            try:
                proc.terminate()
                proc.wait(timeout=3)
            except Exception:
                try:
                    proc.kill()
                    proc.wait(timeout=2)
                except Exception:
                    pass
            for stream in (proc.stdin, proc.stdout):
                try:
                    if stream:
                        stream.close()
                except Exception:
                    pass

    def request(self, operation: str, payload: dict, *, timeout: float = 300,
                progress: Callable[[str], None] | None = None) -> dict:
        if not self.lock.acquire(timeout=timeout):
            raise WorkerError("The CPU processing queue is busy. Please try again shortly.")
        try:
            if self.proc is None or self.proc.poll() is not None:
                self.close()
                self._start()
            assert self.proc is not None and self.proc.stdin is not None and self.inbox is not None
            request_id = uuid.uuid4().hex
            self.proc.stdin.write(json.dumps({"id": request_id, "op": operation, "data": payload}, ensure_ascii=False) + "\n")
            self.proc.stdin.flush()
            deadline = time.monotonic() + timeout
            while True:
                remaining = deadline - time.monotonic()
                if remaining <= 0:
                    raise TimeoutError("CPU processing timed out. Try fewer documents or a shorter recording.")
                try:
                    result = self.inbox.get(timeout=remaining)
                except queue.Empty as exc:
                    raise TimeoutError("CPU processing timed out. Try again with a smaller input.") from exc
                if result.get("fatal"):
                    raise WorkerError(result["fatal"])
                if result.get("id") != request_id:
                    continue
                if "stage" in result:
                    if progress:
                        progress(result["stage"])
                    continue
                if "error" in result:
                    raise WorkerError(result["error"])
                return result["result"]
        except (TimeoutError, BrokenPipeError, OSError):
            self.close()
            raise
        finally:
            self.lock.release()