FazBrowse GitHub Viewer
|
Trending
|
URL:
|
Home
Tools:
[Download Repo ZIP]
[View Raw Code]
[Original HTTPS Page]
diffusers/src/diffusers/utils/outputs.py at main · feisan/diffusers · GitHub
feisan
/
diffusers
Public
forked from
huggingface/diffusers
Notifications
You must be signed in to change notification settings
Fork
0
Star
0
Code
Pull requests
0
Actions
Projects
Security and quality
0
Insights
Additional navigation options
Code
Pull requests
Actions
Projects
Security and quality
Insights
Expand file tree
Breadcrumbs
diffusers
/
src
/
diffusers
/
utils
/
outputs.py
Copy path
More file actions
More file actions
Latest commit
History
History
History
108 lines (84 loc) · 3.57 KB
Breadcrumbs
diffusers
/
src
/
diffusers
/
utils
/
outputs.py
Copy path
File metadata and controls
108 lines (84 loc) · 3.57 KB
Raw
Copy raw file
Download raw file
Open symbols panel
Edit and raw actions
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
# Copyright 2022 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""
Generic utilities
"""
from
collections
import
OrderedDict
from
dataclasses
import
fields
from
typing
import
Any
,
Tuple
import
numpy
as
np
from
.
import_utils
import
is_torch_available
def
is_tensor
(
x
):
"""
Tests if `x` is a `torch.Tensor` or `np.ndarray`.
"""
if
is_torch_available
():
import
torch
if
isinstance
(
x
,
torch
.
Tensor
):
return
True
return
isinstance
(
x
,
np
.
ndarray
)
class
BaseOutput
(
OrderedDict
):
"""
Base class for all model outputs as dataclass. Has a `__getitem__` that allows indexing by integer or slice (like a
tuple) or strings (like a dictionary) that will ignore the `None` attributes. Otherwise behaves like a regular
python dictionary.
<Tip warning={true}>
You can't unpack a `BaseOutput` directly. Use the [`~utils.BaseOutput.to_tuple`] method to convert it to a tuple
before.
</Tip>
"""
def
__post_init__
(
self
):
class_fields
=
fields
(
self
)
# Safety and consistency checks
if
not
len
(
class_fields
):
raise
ValueError
(
f"
{
self
.
__class__
.
__name__
}
has no fields."
)
first_field
=
getattr
(
self
,
class_fields
[
0
].
name
)
other_fields_are_none
=
all
(
getattr
(
self
,
field
.
name
)
is
None
for
field
in
class_fields
[
1
:])
if
other_fields_are_none
and
isinstance
(
first_field
,
dict
):
for
key
,
value
in
first_field
.
items
():
self
[
key
]
=
value
else
:
for
field
in
class_fields
:
v
=
getattr
(
self
,
field
.
name
)
if
v
is
not
None
:
self
[
field
.
name
]
=
v
def
__delitem__
(
self
,
*
args
,
**
kwargs
):
raise
Exception
(
f"You cannot use ``__delitem__`` on a
{
self
.
__class__
.
__name__
}
instance."
)
def
setdefault
(
self
,
*
args
,
**
kwargs
):
raise
Exception
(
f"You cannot use ``setdefault`` on a
{
self
.
__class__
.
__name__
}
instance."
)
def
pop
(
self
,
*
args
,
**
kwargs
):
raise
Exception
(
f"You cannot use ``pop`` on a
{
self
.
__class__
.
__name__
}
instance."
)
def
update
(
self
,
*
args
,
**
kwargs
):
raise
Exception
(
f"You cannot use ``update`` on a
{
self
.
__class__
.
__name__
}
instance."
)
def
__getitem__
(
self
,
k
):
if
isinstance
(
k
,
str
):
inner_dict
=
{
k
:
v
for
(
k
,
v
)
in
self
.
items
()}
return
inner_dict
[
k
]
else
:
return
self
.
to_tuple
()[
k
]
def
__setattr__
(
self
,
name
,
value
):
if
name
in
self
.
keys
()
and
value
is
not
None
:
# Don't call self.__setitem__ to avoid recursion errors
super
().
__setitem__
(
name
,
value
)
super
().
__setattr__
(
name
,
value
)
def
__setitem__
(
self
,
key
,
value
):
# Will raise a KeyException if needed
super
().
__setitem__
(
key
,
value
)
# Don't call self.__setattr__ to avoid recursion errors
super
().
__setattr__
(
key
,
value
)
def
to_tuple
(
self
)
->
Tuple
[
Any
]:
"""
Convert self to a tuple containing all the attributes/keys that are not `None`.
"""
return
tuple
(
self
[
k
]
for
k
in
self
.
keys
())
Back
|
FazBrowse Home
|
New Git URL